From d00d97c9fde78c943ba8e1346738a1ff81a3eb0a Mon Sep 17 00:00:00 2001 From: Kevin Heye Date: Thu, 20 Aug 2026 12:16:13 +0200 Subject: [PATCH 1/7] Add AGENTS.md and re-sync the shared CONTRIBUTING block The team-wide part of CONTRIBUTING.md is a copy of templates/CONTRIBUTING.md in the data_team repository, kept in sync by the sync-style-guide skill. Four repos held byte-identical copies and classification_node had already drifted, so the copies were re-synced from one corrected template and the marker heading convention the skill expects is now documented in the vault. AGENTS.md gives an agent the project structure, the commands that actually work here and the traps specific to this repository. It points at CONTRIBUTING.md for the coding standards instead of repeating them a third time. CLAUDE.md imports it via @AGENTS.md, matching nicegui, zoe-app and data_team. Co-authored-by: Claude Opus 5 (1M context) --- AGENTS.md | 56 +++++++++++++++++++++++++++++++++++++++++ CLAUDE.md | 1 + CONTRIBUTING.md | 66 ++++++++++++++++++++++++++++++++++++------------- 3 files changed, 106 insertions(+), 17 deletions(-) create mode 100644 AGENTS.md create mode 100644 CLAUDE.md diff --git a/AGENTS.md b/AGENTS.md new file mode 100644 index 00000000..1762605a --- /dev/null +++ b/AGENTS.md @@ -0,0 +1,56 @@ +# AI Agent Guidelines for the Learning Loop Node Library + +`learning_loop_node` is the **public Python library** every node uses to talk to the +[Learning Loop](https://learning-loop.ai). Four node types build on it: Trainer, Detector, +Annotator and Converter. It is published to PyPI, so its public surface is an API that +`loop`, `dfine_node`, `yolov5_node` and `classification_node` depend on. + +For coding standards see [CONTRIBUTING.md](CONTRIBUTING.md). [README.md](README.md) documents the +environment variables, the node types and how to write a node against them. + +## Layout + +- `learning_loop_node/` — the library itself: `node.py` and `rest.py` for the shared node + machinery, `trainer/`, `detector/`, `annotation/` for the per-type base logic, `data_classes/` + and `enums/` for the wire types, `loop_communication.py` and `data_exchanger.py` for the + loop-facing HTTP and socket.io traffic. +- `learning_loop_node/tests/` — `annotator`, `detector`, `trainer` and `general` suites. +- `mock_trainer/`, `mock_detector/`, `mock_annotator/` — reference implementations with their own + tests. They are what `loop`'s CI runs against, so they are also the best template for a new node. +- `demo_segmentation_tool/` — a worked annotator example. + +## Running and testing + +The suites talk to a real Learning Loop instance and read their credentials from a local `.env` +(`LOOP_HOST`, `LOOP_USERNAME`, `LOOP_PASSWORD`). Without a reachable loop they cannot pass — do not +treat their failure as a regression you introduced. + +```bash +./run_tests.sh # all suites +./run_tests.sh # passed to pytest as -k +``` + +There is no `.pre-commit-config.yaml` here and no ruff in the project environment, despite what the +shared Linting section says. Lint with: + +```bash +uvx ruff check . +``` + +A clean tree already reports several hundred ruff findings, so a clean run is not a reachable goal. +Compare the count on the files you touched, before and after. + +`.github/workflows/pytest.yml` runs the suites, `publish.yml` releases to PyPI on a tagged release. + +## Working in this repository + +- **This is a library — declaring a dependency is part of its API.** Before removing or loosening + one, grep the consumers (`../loop`, `../dfine_node`, `../yolov5_node`, + `../classification_node`) for the package: a consumer that imports it without declaring it + inherits it from here and breaks when it goes away. +- **Renaming or reshaping anything exported** breaks those four repositories. Say so in the pull + request and check whether a companion change is needed there. +- `../loop` checks this repository out as its `nodes` symlink, so a local change is visible to a + local loop immediately — but only a released version reaches CI and production. +- Bump `version` in `pyproject.toml` for a release; the trainer nodes pin the library version in + their image tags (`A.B.C-nlvX.Y.Z`). diff --git a/CLAUDE.md b/CLAUDE.md new file mode 100644 index 00000000..43c994c2 --- /dev/null +++ b/CLAUDE.md @@ -0,0 +1 @@ +@AGENTS.md diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 5e7f64ed..4ccaa835 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -1,21 +1,35 @@ +# Contributing to Learning Loop Node + ## Data Team Standards ## Linting -We use ruff and [pre-commit](https://github.com/pre-commit/pre-commit) to make sure the coding style is enforced. -You first need to install pre-commit and the corresponding git commit hooks by running the following commands: +We use [ruff](https://docs.astral.sh/ruff/) to enforce the coding style. +Its rules live in the `pyproject.toml` of the respective (sub-)project. + +Repositories that ship a `.pre-commit-config.yaml` run ruff through +[pre-commit](https://github.com/pre-commit/pre-commit). +Install the git hooks once: ```bash python3 -m pip install pre-commit pre-commit install ``` -After that you can make sure your code satisfies the coding style by running the following command: +After that you can check the whole repository with: ```bash pre-commit run --all-files ``` +Repositories without a `.pre-commit-config.yaml` run ruff directly: + +```bash +uvx ruff check . +``` + +The repository-specific part of this file names the exact command wherever it differs. + ## Style Guide ### 1. Single or double quotes @@ -26,8 +40,8 @@ Double quotes may also be used for strings if the string contains a single quote ### 2. String formatting We use f-strings if possible (due to their readability). -Logs shell use the lazy formatting provided by the logging (performance optimization). -When diverging from above rules, please provide a short comment with `# NOTE: ...` to explaining the reason. +Logs shall use the lazy formatting provided by logging (performance optimization). +When diverging from above rules, please provide a short comment with `# NOTE: ...` to explain the reason. ### 3. Line continuation @@ -72,7 +86,7 @@ def parse_config(path: Path) -> Config | None: ... ``` -## 6. Ordering: Important things first +### 6. Ordering: Important things first We want to have main classes and functions at the top of the file, while helper functions and classes should be placed below. This allows to quickly understand the main purpose of the file without having to scroll through a lot of code. @@ -101,23 +115,23 @@ class ReportGenerator: ... ``` -## 7. Docstrings and type hints +### 7. Docstrings and type hints We use the reStructuredText (reST) or Sphinx-style docstring format. We don't declare types in the docstring but use type hints. While type hints are mandatory, parameter descriptions and docstrings should only be used if the function and parameter names are not self-explanatory. -Examples are listetd below: +Examples are listed below: -### One-line Docstring without parameters +#### One-line Docstring without parameters ```python def greet() -> None: """Print a friendly greeting.""" - print("Hello!") + print('Hello!') ``` -### One-line Docstring with parameters +#### One-line Docstring with parameters ```python def square(x: int | float) -> int | float: @@ -129,7 +143,7 @@ def square(x: int | float) -> int | float: return x * x ``` -### Multi-line Docstring without parameters +#### Multi-line Docstring without parameters ```python def get_timestamp() -> str: @@ -140,10 +154,10 @@ def get_timestamp() -> str: as "YYYY-MM-DD HH:MM:SS". It can be useful for logging or displaying time information in user interfaces. """ - return datetime.now().strftime("%Y-%m-%d %H:%M:%S") + return datetime.now().strftime('%Y-%m-%d %H:%M:%S') ``` -### Multi-line Docstring with parameters +#### Multi-line Docstring with parameters ```python def save_data(data: str | bytes, filename: str, overwrite: bool = False) -> None: @@ -157,12 +171,30 @@ def save_data(data: str | bytes, filename: str, overwrite: bool = False) -> None :param data: The content to be saved. Can be text or bytes. :param filename: Path to the destination file. :param overwrite: Whether to overwrite an existing file. - :return: True if the data was saved successfully, False otherwise. :raises FileExistsError: If the file exists and overwrite is False. """ if os.path.exists(filename) and not overwrite: - raise FileExistsError(f"{filename} already exists.") - mode = "wb" if isinstance(data, bytes) else "w" + raise FileExistsError(f'{filename} already exists.') + mode = 'wb' if isinstance(data, bytes) else 'w' with open(filename, mode) as f: f.write(data) ``` + +### 8. Relative imports within a package + +Inside a package we import other modules of the same package with relative imports, not with the package name. +This way the package can be renamed or moved without touching its imports. +We use the shortest path: siblings with a single dot, no going up and back down (`from ..ui.numpad import ...` in a module that lives in `ui/` itself). +Code outside the package — scripts, tests — has no choice and uses absolute imports. + +```python +# preferred (in app/ui/scan_view.py) +from .. import config +from ..system import System +from .numpad import Numpad + +# avoid +from app import config +from app.system import System +from app.ui.numpad import Numpad +``` From 2456d631d065b646b69c9f0c71f7cad9f6c166a9 Mon Sep 17 00:00:00 2001 From: Kevin Heye Date: Thu, 20 Aug 2026 12:40:14 +0200 Subject: [PATCH 2/7] Improve AGENTS.md --- AGENTS.md | 52 ++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 52 insertions(+) diff --git a/AGENTS.md b/AGENTS.md index 1762605a..54866a31 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -19,6 +19,43 @@ environment variables, the node types and how to write a node against them. tests. They are what `loop`'s CI runs against, so they are also the best template for a new node. - `demo_segmentation_tool/` — a worked annotator example. +## Architecture + +Every node is a `FastAPI` subclass (`node.py`): its lifespan connects to the loop and starts a +`repeat_loop` that calls the subclass' `on_repeat` every `repeat_loop_cycle_sec` (5 s) — that loop, +not an event handler, is what drives status reporting, model updates and training continuation. +Subclasses implement `on_startup`, `on_shutdown`, `on_repeat` and `register_sio_events`. + +Two channels lead to the loop and both are needed: `LoopCommunicator` (httpx, login cookies, retry +on 401/429) for the REST API, and a socket.io *client* for status updates and loop-issued commands. +`DataExchanger` sits on top of the communicator and moves images and model zips. + +- **Trainer** — `TrainerLogicGeneric._training_loop` is a state machine over `TrainerState` + (`enums/trainer.py`): download data → download base model → train → sync confusion matrix → + upload model → detect → upload detections → cleanup. `_perform_state` wraps each step: an + ordinary exception records the error and rewinds to the previous state (retried on the next + cycle), a `CriticalError` jumps to `ReadyForCleanup`. Every transition is persisted through + `LastTrainingIO`, so `try_continue_run_if_incomplete` resumes an interrupted training after a + restart. The loop starts a training via the `begin_training` sio event. A concrete trainer + implements `_train`, `_do_detections`, `_get_new_best_training_state`, `_on_metrics_published`, + `_get_latest_model_files` and `_clear_training_data`; `TrainerLogic` adds an `Executor` for + trainers that shell out to a training process. +- **Detector** — the exception to the pattern: constructed with `needs_login=False, needs_sio=False`, + so it has no sio client to the loop. It *hosts* a socket.io server for its own clients and polls + `/{org}/projects/{project}/deployment/target` over REST in `on_repeat` instead. `_DetectorState` + (`_Initializing` / `_Updating` / `_ActiveDetector`) models the model swap: download to + `models/`, build a `DetectorLogic` through the factory, then swap atomically so the old + model keeps serving until the new one is ready (unless `EXCLUSIVE_MODEL_BUILD` frees VRAM first). + `OperationMode` gates whether updates may happen at all. Detections flow through + `RelevanceFilter`, which writes selected images to the `Outbox` on disk; a separate upload process + drains it. +- **Annotator** — thin: it forwards the loop frontend's `handle_user_input` events into + `AnnotatorLogic` and keeps a per-frontend history. + +All node state lives under `GLOBALS.data_folder` (`DATA_FOLDER`, default `/data`): `uuids.json` +(the node uuid is derived from its name and reused across restarts), `models/` plus the +`current_model` symlink, `outbox/`, and the per-project training folders. + ## Running and testing The suites talk to a real Learning Loop instance and read their credentials from a local `.env` @@ -30,6 +67,21 @@ treat their failure as a regression you introduced. ./run_tests.sh # passed to pytest as -k ``` +Each suite carries its own `pytest.ini` (that is where `asyncio_mode = auto` comes from), so always +run pytest with a path inside one suite — a bare `pytest` from the repository root picks up no +config and the async tests error out: + +```bash +python -m pytest learning_loop_node/tests/trainer -v # one suite +python -m pytest learning_loop_node/tests/trainer/test_errors.py -v # one file +python -m pytest learning_loop_node/tests/trainer -v -k # one test +``` + +An autouse fixture repoints `GLOBALS.data_folder` at `/tmp/learning_loop_lib_data` and wipes it +around every test, so tests never touch `/data`. The `general` suite generates and deletes a real +`zauberzeug/pytest_nodelib_general` project on the loop; the detector suite starts the node in a +forked uvicorn process on `GLOBALS.detector_port`. + There is no `.pre-commit-config.yaml` here and no ruff in the project environment, despite what the shared Linting section says. Lint with: From a004efedf42bfc98c10514f782db77acecdc6386 Mon Sep 17 00:00:00 2001 From: Kevin Heye Date: Thu, 20 Aug 2026 15:48:13 +0200 Subject: [PATCH 3/7] Keep AGENTS.md free of private repository names MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit learning_loop_node is public, but AGENTS.md named dfine_node and classification_node, which are not — neither appears anywhere else in the public tree. The dependency rule they illustrated holds without naming them, and "grep ../dfine_node" is advice no external contributor can follow anyway. Also drop the claim that the shared Linting section contradicts this repo: that section now covers repositories without a .pre-commit-config.yaml explicitly. Co-authored-by: Claude Opus 5 (1M context) --- AGENTS.md | 16 +++++++--------- 1 file changed, 7 insertions(+), 9 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index 54866a31..823eeee6 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -2,8 +2,8 @@ `learning_loop_node` is the **public Python library** every node uses to talk to the [Learning Loop](https://learning-loop.ai). Four node types build on it: Trainer, Detector, -Annotator and Converter. It is published to PyPI, so its public surface is an API that -`loop`, `dfine_node`, `yolov5_node` and `classification_node` depend on. +Annotator and Converter. It is published to PyPI, so its public surface is an API that the +Learning Loop backend and every node repository depends on — `../yolov5_node` is the public one. For coding standards see [CONTRIBUTING.md](CONTRIBUTING.md). [README.md](README.md) documents the environment variables, the node types and how to write a node against them. @@ -82,8 +82,7 @@ around every test, so tests never touch `/data`. The `general` suite generates a `zauberzeug/pytest_nodelib_general` project on the loop; the detector suite starts the node in a forked uvicorn process on `GLOBALS.detector_port`. -There is no `.pre-commit-config.yaml` here and no ruff in the project environment, despite what the -shared Linting section says. Lint with: +There is no `.pre-commit-config.yaml` here and no ruff in the project environment. Lint with: ```bash uvx ruff check . @@ -97,11 +96,10 @@ Compare the count on the files you touched, before and after. ## Working in this repository - **This is a library — declaring a dependency is part of its API.** Before removing or loosening - one, grep the consumers (`../loop`, `../dfine_node`, `../yolov5_node`, - `../classification_node`) for the package: a consumer that imports it without declaring it - inherits it from here and breaks when it goes away. -- **Renaming or reshaping anything exported** breaks those four repositories. Say so in the pull - request and check whether a companion change is needed there. + one, grep every consuming repository checked out beside this one for the package: a consumer that + imports it without declaring it inherits it from here and breaks when it goes away. +- **Renaming or reshaping anything exported** breaks those repositories. Say so in the pull request + and check whether a companion change is needed there. - `../loop` checks this repository out as its `nodes` symlink, so a local change is visible to a local loop immediately — but only a released version reaches CI and production. - Bump `version` in `pyproject.toml` for a release; the trainer nodes pin the library version in From 7747813a0387db9a38d083aee158727d23da8d06 Mon Sep 17 00:00:00 2001 From: Kevin Heye Date: Thu, 20 Aug 2026 16:47:01 +0200 Subject: [PATCH 4/7] Share detector postprocessing across nodes Every detector node does the same three things with a model's raw output: drop low-confidence predictions, suppress overlapping boxes, and turn what survives into the loop's dataclasses. None of it depends on the model, yet each node repository had grown its own copy, and the copies had drifted: clip_point was byte-identical in three places, while clip_box existed in the same three plus a fourth that used a centre-based convention under the same name. Move that code here as detector/postprocess.py and detector/geometry.py, and give the centre-based variant its own name (clip_box_centered) so the two conventions can no longer be confused. The library also carried two containers for the same detections: a detector reports ImageMetadata, a trainer's auto-detection pass reports Detections. to_image_metadata and to_detections now build both from one routine, so the two paths clip and filter identically instead of diverging. The scalar arguments of non_max_suppression and post_process are keyword-only: transposing origin_h and origin_w was too easy to do. Add tests/unit, the first suite that runs without a Learning Loop, and gate the credential-dependent jobs on it in CI. Co-Authored-By: Claude Opus 5 (1M context) --- .github/workflows/pytest.yml | 22 ++ AGENTS.md | 28 +- learning_loop_node/detector/geometry.py | 69 +++++ learning_loop_node/detector/postprocess.py | 264 ++++++++++++++++++ learning_loop_node/tests/unit/__init__.py | 0 learning_loop_node/tests/unit/pytest.ini | 10 + .../tests/unit/test_geometry.py | 35 +++ .../tests/unit/test_postprocess.py | 168 +++++++++++ run_tests.sh | 3 + 9 files changed, 594 insertions(+), 5 deletions(-) create mode 100644 learning_loop_node/detector/geometry.py create mode 100644 learning_loop_node/detector/postprocess.py create mode 100644 learning_loop_node/tests/unit/__init__.py create mode 100644 learning_loop_node/tests/unit/pytest.ini create mode 100644 learning_loop_node/tests/unit/test_geometry.py create mode 100644 learning_loop_node/tests/unit/test_postprocess.py diff --git a/.github/workflows/pytest.yml b/.github/workflows/pytest.yml index 3e456ca2..e334d0b4 100644 --- a/.github/workflows/pytest.yml +++ b/.github/workflows/pytest.yml @@ -3,7 +3,28 @@ name: Run Tests on: [push] jobs: + unit: + runs-on: ubuntu-latest + timeout-minutes: 10 + steps: + - uses: actions/checkout@v4 + - name: set up Python + uses: actions/setup-python@v5 + with: + python-version: "3.10" + - name: set up uv + uses: astral-sh/setup-uv@v3 + - name: install dependencies + run: | + uv sync --extra dev --frozen + - name: test_unit + # no Learning Loop and no secrets needed, so this suite gates the slow ones + run: | + uv run python -m pytest learning_loop_node/tests/unit -v + pytest_3_10: + needs: + - unit runs-on: ubuntu-latest timeout-minutes: 30 strategy: @@ -91,6 +112,7 @@ jobs: slack: needs: + - unit - pytest_3_10 - pytest_3_13 if: always() # also execute when pytest fails diff --git a/AGENTS.md b/AGENTS.md index 823eeee6..cf26c812 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -14,7 +14,9 @@ environment variables, the node types and how to write a node against them. machinery, `trainer/`, `detector/`, `annotation/` for the per-type base logic, `data_classes/` and `enums/` for the wire types, `loop_communication.py` and `data_exchanger.py` for the loop-facing HTTP and socket.io traffic. -- `learning_loop_node/tests/` — `annotator`, `detector`, `trainer` and `general` suites. +- `learning_loop_node/tests/` — `unit`, `annotator`, `detector`, `trainer` and `general` suites. + Only `unit` runs without a Learning Loop; it covers the pure helpers such as + `detector/postprocess.py` and `detector/geometry.py`. - `mock_trainer/`, `mock_detector/`, `mock_annotator/` — reference implementations with their own tests. They are what `loop`'s CI runs against, so they are also the best template for a new node. - `demo_segmentation_tool/` — a worked annotator example. @@ -52,15 +54,30 @@ on 401/429) for the REST API, and a socket.io *client* for status updates and lo - **Annotator** — thin: it forwards the loop frontend's `handle_user_input` events into `AnnotatorLogic` and keeps a per-frontend history. +`detector/postprocess.py` and `detector/geometry.py` hold the parts of a detector that do *not* +depend on the model: confidence filtering, per-class NMS, box/point clipping, and turning +predictions into the loop's dataclasses. A node should import them rather than write its own — +every node repository had grown its own drifting copy, which is why they live here. Note the two +containers: `to_image_metadata` builds what a **detector** node reports, `to_detections` what a +**trainer**'s auto-detection pass reports. Both go through one routine, so the two paths cannot +drift apart again. + All node state lives under `GLOBALS.data_folder` (`DATA_FOLDER`, default `/data`): `uuids.json` (the node uuid is derived from its name and reused across restarts), `models/` plus the `current_model` symlink, `outbox/`, and the per-project training folders. ## Running and testing -The suites talk to a real Learning Loop instance and read their credentials from a local `.env` -(`LOOP_HOST`, `LOOP_USERNAME`, `LOOP_PASSWORD`). Without a reachable loop they cannot pass — do not -treat their failure as a regression you introduced. +The `unit` suite is self-contained — no Learning Loop, no credentials, no network — so it is the +one to run while iterating, and the one CI gates the others on: + +```bash +python -m pytest learning_loop_node/tests/unit -v +``` + +Every other suite talks to a real Learning Loop instance and reads its credentials from a local +`.env` (`LOOP_HOST`, `LOOP_USERNAME`, `LOOP_PASSWORD`). Without a reachable loop those cannot pass — +do not treat their failure as a regression you introduced. ```bash ./run_tests.sh # all suites @@ -91,7 +108,8 @@ uvx ruff check . A clean tree already reports several hundred ruff findings, so a clean run is not a reachable goal. Compare the count on the files you touched, before and after. -`.github/workflows/pytest.yml` runs the suites, `publish.yml` releases to PyPI on a tagged release. +`.github/workflows/pytest.yml` runs the suites (the `unit` job first, without secrets), +`publish.yml` releases to PyPI on a tagged release. ## Working in this repository diff --git a/learning_loop_node/detector/geometry.py b/learning_loop_node/detector/geometry.py new file mode 100644 index 00000000..b17ed33b --- /dev/null +++ b/learning_loop_node/detector/geometry.py @@ -0,0 +1,69 @@ +"""Box and point clipping shared by every detector node. + +The loop stores a box as its top-left corner plus a size, so :func:`clip_box` is the form a +node needs when it hands detections to the loop. Model outputs are not always in that form — +:func:`clip_box_centered` keeps the centre-based convention explicit instead of letting two +incompatible functions share one name, which is how the same helper ended up meaning two +different things in different node repositories. +""" + + +def clip_box( + *, + x1: float, + y1: float, + width: float, + height: float, + img_width: int, + img_height: int, +) -> tuple[int, int, int, int]: + """Clip a top-left-anchored box to the image bounds. + + :param x1: Left edge of the box. + :param y1: Top edge of the box. + :return: The clipped ``(x1, y1, width, height)`` as ints; the size is never negative. + """ + x2 = x1 + width + y2 = y1 + height + + clipped_x1 = round(max(0.0, x1)) + clipped_y1 = round(max(0.0, y1)) + clipped_x2 = round(min(float(img_width), x2)) + clipped_y2 = round(min(float(img_height), y2)) + + clipped_width = max(clipped_x2 - clipped_x1, 0) + clipped_height = max(clipped_y2 - clipped_y1, 0) + + return clipped_x1, clipped_y1, clipped_width, clipped_height + + +def clip_box_centered( + *, + x: float, + y: float, + width: float, + height: float, + img_width: int, + img_height: int, +) -> tuple[float, float, float, float]: + """Clip a centre-anchored box to the image bounds, keeping it centre-anchored. + + Clipping moves the centre, because only the part of the box inside the image survives. + + :param x: Horizontal centre of the box. + :param y: Vertical centre of the box. + :return: The clipped ``(x, y, width, height)``, still centre-anchored. + """ + left = max(0.0, x - 0.5 * width) + top = max(0.0, y - 0.5 * height) + right = min(float(img_width), x + 0.5 * width) + bottom = min(float(img_height), y + 0.5 * height) + + return 0.5 * (left + right), 0.5 * (top + bottom), right - left, bottom - top + + +def clip_point(x: float, y: float, img_width: int, img_height: int) -> tuple[float, float]: + """Clamp a point into the image bounds.""" + x = min(max(0, x), img_width) + y = min(max(0, y), img_height) + return x, y diff --git a/learning_loop_node/detector/postprocess.py b/learning_loop_node/detector/postprocess.py new file mode 100644 index 00000000..ae47fbcc --- /dev/null +++ b/learning_loop_node/detector/postprocess.py @@ -0,0 +1,264 @@ +"""Model-agnostic detection postprocessing. + +Every detection node ends up doing the same three things with a model's raw output: drop +low-confidence predictions, suppress overlapping boxes, and turn what survives into the +loop's detection dataclasses. None of that depends on the model, so it lives here rather +than being re-derived — and re-diverging — in each node repository. + +Two containers carry the same detections in this library: a detector node reports +:class:`~learning_loop_node.data_classes.image_metadata.ImageMetadata`, while a trainer's +auto-detection pass reports :class:`~learning_loop_node.data_classes.detections.Detections`. +:func:`to_image_metadata` and :func:`to_detections` build them from the same routine, so both +paths clip and filter identically. +""" + +import logging +from collections import namedtuple + +import numpy as np + +from ..data_classes import ( + BoxDetection, + Category, + Detections, + ImageMetadata, + ModelInformation, + PointDetection, +) +from ..enums import CategoryType +from .geometry import clip_box, clip_point + +MIN_BOX_SIZE: int = 2 +"""Boxes this small are dropped: they carry no usable information and clutter the loop.""" + +Detection = namedtuple('Detection', 'x y w h category probability') +"""One surviving prediction. ``x``/``y`` are the top-left corner, ``category`` is an index +into :attr:`ModelInformation.categories`.""" + + +def category_by_index(model_information: ModelInformation, index: int) -> Category: + """Resolve the category a model's class index refers to. + + Models emit class indices, and the order of ``model_information.categories`` is what + gives them meaning — so an out-of-range index is a model/metadata mismatch, not a + detection to skip quietly. + + :raises ValueError: If the index is outside the model's category list. + """ + categories = model_information.categories + if not 0 <= index < len(categories): + raise ValueError( + f'category index {index} is out of range for a model with {len(categories)} categories') + return categories[index] + + +def category_by_name(model_information: ModelInformation, name: str) -> Category: + """Resolve a category by name, for models whose outputs are named rather than indexed. + + :raises ValueError: If no category of that name exists. + """ + for category in model_information.categories: + if category.name == name: + return category + known = ', '.join(category.name for category in model_information.categories) + raise ValueError(f'unknown category name {name!r}; the model knows: {known}') + + +def bbox_iou( + box1: np.ndarray, + box2: np.ndarray, +) -> np.ndarray: + """Compute IoU between box1 (1x4) and box2 (Nx4), both in x1y1x2y2 format.""" + b1_x1, b1_y1, b1_x2, b1_y2 = box1[:, 0], box1[:, 1], box1[:, 2], box1[:, 3] + b2_x1, b2_y1, b2_x2, b2_y2 = box2[:, 0], box2[:, 1], box2[:, 2], box2[:, 3] + + inter_x1 = np.maximum(b1_x1, b2_x1) + inter_y1 = np.maximum(b1_y1, b2_y1) + inter_x2 = np.minimum(b1_x2, b2_x2) + inter_y2 = np.minimum(b1_y2, b2_y2) + + inter_area = np.clip(inter_x2 - inter_x1 + 1, 0, None) * np.clip(inter_y2 - inter_y1 + 1, 0, None) + b1_area = (b1_x2 - b1_x1 + 1) * (b1_y2 - b1_y1 + 1) + b2_area = (b2_x2 - b2_x1 + 1) * (b2_y2 - b2_y1 + 1) + + return inter_area / (b1_area + b2_area - inter_area + 1e-16) + + +def non_max_suppression( + boxes: np.ndarray, + scores: np.ndarray, + classes: np.ndarray, + *, + iou_threshold: float, + origin_h: int, + origin_w: int, +) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + """Clip to image bounds, sort by score descending, apply per-class NMS. + + :return: ``(boxes, scores, classes)`` — filtered arrays in the same order. + """ + boxes = boxes.copy() + boxes[:, 0] = np.clip(boxes[:, 0], 0, origin_w - 1) + boxes[:, 2] = np.clip(boxes[:, 2], 0, origin_w - 1) + boxes[:, 1] = np.clip(boxes[:, 1], 0, origin_h - 1) + boxes[:, 3] = np.clip(boxes[:, 3], 0, origin_h - 1) + + order = np.argsort(-scores) + boxes = boxes[order] + scores = scores[order] + classes = classes[order] + + keep_indices: list[int] = [] + for cls in np.unique(classes): + cls_mask = classes == cls + cls_indices = np.where(cls_mask)[0] + cls_boxes = boxes[cls_mask] + + while len(cls_indices) > 0: + keep_indices.append(cls_indices[0]) + if len(cls_indices) == 1: + break + ious = bbox_iou(cls_boxes[0:1], cls_boxes[1:]) + keep_mask = ious.flatten() <= iou_threshold + cls_indices = cls_indices[1:][keep_mask] + cls_boxes = cls_boxes[1:][keep_mask] + + keep_indices = sorted(keep_indices) + return boxes[keep_indices], scores[keep_indices], classes[keep_indices] + + +def post_process( + boxes: np.ndarray, + scores: np.ndarray, + classes: np.ndarray, + *, + conf_threshold: float, + iou_threshold: float, + origin_h: int, + origin_w: int, +) -> list[Detection]: + """Filter by confidence, run NMS, return a :class:`Detection` list in x/y/w/h form.""" + mask = scores > conf_threshold + boxes = boxes[mask].copy() + scores = scores[mask] + classes = classes[mask] + + if len(scores) == 0: + return [] + + boxes, scores, classes = non_max_suppression( + boxes, scores, classes, + iou_threshold=iou_threshold, origin_h=origin_h, origin_w=origin_w) + + result = [] + for j, box in enumerate(boxes): + x1, y1, x2, y2 = box + w = x2 - x1 + h = y2 - y1 + result.append(Detection(int(x1), int(y1), int(w), int(h), + int(classes[j]), round(float(scores[j]), 2))) + return result + + +def detections_from_xyxy( + *, + labels: list[float], + boxes: list[list[float]], + scores: list[float], +) -> list[Detection]: + """Convert already-suppressed model output into :class:`Detection` values. + + For nodes whose model (or a torch/ONNX op) has done the suppression already, so only the + coordinate conversion is left. Corners are rounded rather than truncated, which is half a + pixel more faithful than :func:`post_process` — that one keeps truncating so its output + stays bit-identical to what detectors reported before this module existed. + """ + result = [] + for label, box, score in zip(labels, boxes, scores, strict=True): + x1, y1, x2, y2 = (round(value) for value in box) + result.append(Detection(x1, y1, x2 - x1, y2 - y1, int(label), score)) + return result + + +def _append_detections( + target: ImageMetadata | Detections, + detections: list[Detection], + model_information: ModelInformation, + im_height: int, + im_width: int, +) -> None: + """Resolve each detection's category and append it to ``target``, clipped to the image.""" + skipped_detections = [] + + for detection in detections: + x, y, w, h, category_idx, probability = detection + category = category_by_index(model_information, category_idx) + if w <= MIN_BOX_SIZE or h <= MIN_BOX_SIZE: + skipped_detections.append((category.name, detection)) + continue + if category.type == CategoryType.Box: + clipped_x1, clipped_y1, clipped_w, clipped_h = clip_box( + x1=x, + y1=y, + width=w, + height=h, + img_width=im_width, + img_height=im_height, + ) + target.box_detections.append( + BoxDetection( + category_name=category.name, + x=clipped_x1, + y=clipped_y1, + width=clipped_w, + height=clipped_h, + category_id=category.id, + model_name=model_information.version, + confidence=probability, + ) + ) + elif category.type == CategoryType.Point: + cx, cy = x + w / 2, y + h / 2 + cx, cy = clip_point(cx, cy, im_width, im_height) + target.point_detections.append( + PointDetection( + category_name=category.name, + x=cx, + y=cy, + category_id=category.id, + model_name=model_information.version, + confidence=probability, + ) + ) + else: + logging.warning('Unsupported category type %s for category %s', category.type, category.name) + + if skipped_detections: + log_msg = '\n'.join([str(d) for d in skipped_detections]) + logging.warning('Removed %d small detections from result: \n%s', len(skipped_detections), log_msg) + + +def to_image_metadata( + detections: list[Detection], + model_information: ModelInformation, + im_height: int, + im_width: int, +) -> ImageMetadata: + """Build the container a *detector* node reports from a list of detections.""" + image_metadata = ImageMetadata() + _append_detections(image_metadata, detections, model_information, im_height, im_width) + return image_metadata + + +def to_detections( + detections: list[Detection], + model_information: ModelInformation, + im_height: int, + im_width: int, + *, + image_id: str | None = None, +) -> Detections: + """Build the container a *trainer*'s auto-detection pass reports.""" + result = Detections(image_id=image_id) + _append_detections(result, detections, model_information, im_height, im_width) + return result diff --git a/learning_loop_node/tests/unit/__init__.py b/learning_loop_node/tests/unit/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/learning_loop_node/tests/unit/pytest.ini b/learning_loop_node/tests/unit/pytest.ini new file mode 100644 index 00000000..2207d1ac --- /dev/null +++ b/learning_loop_node/tests/unit/pytest.ini @@ -0,0 +1,10 @@ +[pytest] +python_files = test_*.py +asyncio_mode = auto + +cache_dir = /tmp/pytest_cache + +# for debbuging tests: +; log_cli_level = INFO +; log_cli_format = %(asctime)s [%(levelname)8s] %(message)s (%(filename)s:%(lineno)s) +; log_cli_date_format=%Y-%m-%d %H:%M:%S \ No newline at end of file diff --git a/learning_loop_node/tests/unit/test_geometry.py b/learning_loop_node/tests/unit/test_geometry.py new file mode 100644 index 00000000..f056f79d --- /dev/null +++ b/learning_loop_node/tests/unit/test_geometry.py @@ -0,0 +1,35 @@ +from ...detector.geometry import clip_box, clip_box_centered, clip_point + + +def test_box_inside_the_image_is_unchanged(): + assert clip_box(x1=10, y1=20, width=30, height=40, img_width=100, img_height=100) == (10, 20, 30, 40) + + +def test_box_is_clipped_to_the_image_bounds(): + assert clip_box(x1=-20, y1=-20, width=60, height=60, img_width=100, img_height=100) == (0, 0, 40, 40) + assert clip_box(x1=80, y1=80, width=40, height=40, img_width=100, img_height=100) == (80, 80, 20, 20) + + +def test_box_fully_outside_the_image_collapses_to_zero_size(): + # The corner is only clamped at the lower bound, so it stays at 200 — the zero size is + # what marks the box as empty. Detector output cannot reach here, because + # non_max_suppression already clips every box into the image. + assert clip_box(x1=200, y1=200, width=10, height=10, img_width=100, img_height=100) == (200, 200, 0, 0) + + +def test_box_corners_are_rounded(): + assert clip_box(x1=10.4, y1=10.6, width=20.0, height=20.0, img_width=100, img_height=100) == (10, 11, 20, 20) + + +def test_centered_box_keeps_its_centre_when_it_fits(): + assert clip_box_centered(x=50, y=50, width=20, height=20, img_width=100, img_height=100) == (50, 50, 20, 20) + + +def test_clipping_a_centered_box_moves_its_centre(): + # only the right half of the box is inside the image, so the centre moves right + assert clip_box_centered(x=0, y=50, width=20, height=20, img_width=100, img_height=100) == (5, 50, 10, 20) + + +def test_point_is_clamped_into_the_image(): + assert clip_point(50, 50, 100, 100) == (50, 50) + assert clip_point(-10, 150, 100, 100) == (0, 100) diff --git a/learning_loop_node/tests/unit/test_postprocess.py b/learning_loop_node/tests/unit/test_postprocess.py new file mode 100644 index 00000000..14dbcd6a --- /dev/null +++ b/learning_loop_node/tests/unit/test_postprocess.py @@ -0,0 +1,168 @@ +import numpy as np +import pytest + +from ...data_classes import Category, ModelInformation +from ...detector.postprocess import ( + Detection, + bbox_iou, + category_by_index, + category_by_name, + detections_from_xyxy, + non_max_suppression, + post_process, + to_detections, + to_image_metadata, +) +from ...enums import CategoryType + +BOX = Category(id='uuid-box', name='car', type=CategoryType.Box) +POINT = Category(id='uuid-point', name='weed', type=CategoryType.Point) + + +def model_information(*categories: Category) -> ModelInformation: + return ModelInformation(id='model-uuid', host='localhost', organization='zauberzeug', + project='pytest', version='1.2', categories=list(categories or (BOX, POINT))) + + +# ---------------------------------------------------------------- category resolution + +def test_category_is_resolved_by_index(): + assert category_by_index(model_information(), 1) is POINT + + +@pytest.mark.parametrize('index', [-1, 2, 99]) +def test_an_index_outside_the_model_categories_is_an_error(index: int): + # a mismatch between model and metadata must not be silently skipped + with pytest.raises(ValueError, match='out of range'): + category_by_index(model_information(), index) + + +def test_category_is_resolved_by_name(): + assert category_by_name(model_information(), 'weed') is POINT + + +def test_an_unknown_category_name_lists_the_known_ones(): + with pytest.raises(ValueError, match='car, weed'): + category_by_name(model_information(), 'tractor') + + +# ---------------------------------------------------------------- iou and suppression + +def test_identical_boxes_have_an_iou_of_one(): + box = np.array([[0, 0, 10, 10]], dtype=np.float32) + assert bbox_iou(box, box)[0] == pytest.approx(1.0) + + +def test_disjoint_boxes_have_an_iou_of_zero(): + a = np.array([[0, 0, 10, 10]], dtype=np.float32) + b = np.array([[100, 100, 110, 110]], dtype=np.float32) + assert bbox_iou(a, b)[0] == pytest.approx(0.0) + + +def test_overlapping_boxes_of_one_class_are_suppressed_keeping_the_best(): + boxes = np.array([[10, 10, 60, 60], [12, 12, 62, 62]], dtype=np.float32) + scores = np.array([0.9, 0.8], dtype=np.float32) + kept_boxes, kept_scores, _ = non_max_suppression( + boxes, scores, np.array([0, 0]), iou_threshold=0.45, origin_h=200, origin_w=200) + assert len(kept_boxes) == 1 + assert kept_scores[0] == pytest.approx(0.9) + + +def test_overlapping_boxes_of_different_classes_both_survive(): + boxes = np.array([[10, 10, 60, 60], [12, 12, 62, 62]], dtype=np.float32) + kept_boxes, _, _ = non_max_suppression( + boxes, np.array([0.9, 0.8], dtype=np.float32), np.array([0, 1]), + iou_threshold=0.45, origin_h=200, origin_w=200) + assert len(kept_boxes) == 2 + + +def test_suppression_clips_boxes_into_the_image(): + boxes = np.array([[-10, -10, 300, 300]], dtype=np.float32) + kept_boxes, _, _ = non_max_suppression( + boxes, np.array([0.9], dtype=np.float32), np.array([0]), + iou_threshold=0.45, origin_h=100, origin_w=100) + assert list(kept_boxes[0]) == [0, 0, 99, 99] + + +# ---------------------------------------------------------------- post_process + +def test_post_process_drops_predictions_below_the_confidence_threshold(): + boxes = np.array([[10, 10, 60, 60], [100, 100, 150, 150]], dtype=np.float32) + result = post_process(boxes, np.array([0.9, 0.1], dtype=np.float32), np.array([0, 0]), + conf_threshold=0.5, iou_threshold=0.45, origin_h=200, origin_w=200) + assert result == [Detection(10, 10, 50, 50, 0, 0.9)] + + +def test_post_process_on_an_empty_prediction_returns_nothing(): + empty_boxes = np.zeros((0, 4), dtype=np.float32) + assert post_process(empty_boxes, np.zeros(0, dtype=np.float32), np.zeros(0, dtype=int), + conf_threshold=0.5, iou_threshold=0.45, origin_h=10, origin_w=10) == [] + + +def test_already_suppressed_output_is_converted_with_rounded_corners(): + assert detections_from_xyxy(labels=[1.0], boxes=[[10.4, 10.6, 60.4, 60.6]], scores=[0.55]) == \ + [Detection(10, 11, 50, 50, 1, 0.55)] + + +def test_converting_already_suppressed_output_requires_matching_lengths(): + with pytest.raises(ValueError): + detections_from_xyxy(labels=[1.0, 2.0], boxes=[[0.0, 0.0, 1.0, 1.0]], scores=[0.5]) + + +# ---------------------------------------------------------------- building the containers + +def test_a_box_category_becomes_a_box_detection(): + metadata = to_image_metadata([Detection(10, 20, 30, 40, 0, 0.9)], model_information(), 200, 200) + assert len(metadata.point_detections) == 0 + detection = metadata.box_detections[0] + assert (detection.x, detection.y, detection.width, detection.height) == (10, 20, 30, 40) + assert (detection.category_name, detection.category_id) == ('car', 'uuid-box') + assert detection.model_name == '1.2' + assert detection.confidence == pytest.approx(0.9) + + +def test_a_point_category_becomes_the_centre_of_the_box(): + metadata = to_image_metadata([Detection(100, 100, 40, 40, 1, 0.7)], model_information(), 200, 200) + assert len(metadata.box_detections) == 0 + detection = metadata.point_detections[0] + assert (detection.x, detection.y) == (120, 120) + assert detection.category_id == 'uuid-point' + + +def test_detections_are_clipped_to_the_image(): + metadata = to_image_metadata([Detection(-20, -20, 60, 60, 0, 0.5)], model_information(), 200, 200) + detection = metadata.box_detections[0] + assert (detection.x, detection.y, detection.width, detection.height) == (0, 0, 40, 40) + + +@pytest.mark.parametrize('width,height', [(2, 30), (30, 2), (1, 1)]) +def test_boxes_too_small_to_be_useful_are_dropped(width: int, height: int): + metadata = to_image_metadata([Detection(5, 5, width, height, 0, 0.5)], model_information(), 200, 200) + assert len(metadata) == 0 + + +def test_a_category_type_the_node_cannot_report_is_skipped(): + classification = Category(id='uuid-cls', name='ripe', type=CategoryType.Classification) + metadata = to_image_metadata([Detection(10, 10, 30, 30, 0, 0.5)], + model_information(classification), 200, 200) + assert len(metadata) == 0 + + +def test_the_trainer_container_carries_the_image_id(): + result = to_detections([Detection(10, 20, 30, 40, 0, 0.9)], model_information(), 200, 200, + image_id='image-uuid') + assert result.image_id == 'image-uuid' + assert len(result.box_detections) == 1 + + +def test_trainer_and_detector_paths_agree_on_the_same_detections(): + """The whole point of sharing this code: auto-detections and live detections must match.""" + detections = [Detection(-5, -5, 60, 60, 0, 0.9), Detection(100, 100, 40, 40, 1, 0.7), + Detection(5, 5, 1, 1, 0, 0.5)] + metadata = to_image_metadata(detections, model_information(), 200, 200) + result = to_detections(detections, model_information(), 200, 200, image_id='image-uuid') + + assert [(d.x, d.y, d.width, d.height, d.category_id) for d in result.box_detections] == \ + [(d.x, d.y, d.width, d.height, d.category_id) for d in metadata.box_detections] + assert [(d.x, d.y, d.category_id) for d in result.point_detections] == \ + [(d.x, d.y, d.category_id) for d in metadata.point_detections] diff --git a/run_tests.sh b/run_tests.sh index e83d3852..e8ae7b66 100755 --- a/run_tests.sh +++ b/run_tests.sh @@ -6,6 +6,7 @@ set -o allexport; source .env; set +o allexport # Check if argument is provided if [ $# -eq 1 ]; then # Run tests with filter + python -m pytest learning_loop_node/tests/unit -v -s -k "$1" python -m pytest learning_loop_node/tests/annotator -v -s -k "$1" python -m pytest learning_loop_node/tests/detector -v -s -k "$1" python -m pytest learning_loop_node/tests/trainer -v -s -k "$1" @@ -17,6 +18,8 @@ fi # Run the tests +# unit runs first: it is the only suite that needs no Learning Loop +python -m pytest learning_loop_node/tests/unit -v python -m pytest learning_loop_node/tests/annotator -v python -m pytest learning_loop_node/tests/detector -v python -m pytest learning_loop_node/tests/trainer -v From fa8505c1d76170ec20e50aa8e522c1615dfee3b3 Mon Sep 17 00:00:00 2001 From: Kevin Heye Date: Thu, 20 Aug 2026 19:18:53 +0200 Subject: [PATCH 5/7] Share the node entry point boilerplate Every main.py in every node repository repeats the same twenty lines: read UVICORN_RELOAD, build the settings, construct the node, call uvicorn.run. Eight copies across three repositories, and they had drifted - dfine parses settings with configargparse so each one is both a flag and an environment variable, while the others hand-roll os.getenv plus asserts and hardcode the host and port. node_parser and run_node put that in one place. A setting is named after its flag: --conf-threshold reads CONF_THRESHOLD. That is what yolov5_node and classification_node already use, so adopting this renames nothing in the repositories that are actually deployed. --host and --port are the exception, reading NODE_HOST and NODE_PORT. The bare HOST already means the loop's address - every deployment sets it - and a node adopting that would hand it to uvicorn and fail to bind. dfine_node required a DFINE_TRAINER_ / DFINE_DETECTOR_ prefix, so it passes legacy_env_prefix and its prefixed names keep working with a warning naming the replacement. The logging.basicConfig call the trainers carry is not reproduced here - the library configures the root logger on import, so it has been a no-op. Co-Authored-By: Claude Opus 5 (1M context) --- AGENTS.md | 8 ++ learning_loop_node/helpers/entrypoint.py | 75 +++++++++++++++++++ .../tests/unit/test_entrypoint.py | 71 ++++++++++++++++++ pyproject.toml | 1 + uv.lock | 11 +++ 5 files changed, 166 insertions(+) create mode 100644 learning_loop_node/helpers/entrypoint.py create mode 100644 learning_loop_node/tests/unit/test_entrypoint.py diff --git a/AGENTS.md b/AGENTS.md index cf26c812..c7877d5d 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -54,6 +54,14 @@ on 401/429) for the REST API, and a socket.io *client* for status updates and lo - **Annotator** — thin: it forwards the loop frontend's `handle_user_input` events into `AnnotatorLogic` and keeps a per-frontend history. +`helpers/entrypoint.py` holds what every node's `main.py` repeats: `node_parser` builds a +configargparse parser with `--host`/`--port`, `run_node` starts uvicorn. A setting is a flag +*and* an environment variable from one declaration — `--conf-threshold` reads +`CONF_THRESHOLD`. The exception is `--host`/`--port`, which read `NODE_HOST`/`NODE_PORT`: the +bare `HOST` already means the loop's address, and a node adopting it would hand it to uvicorn +and fail to bind. A node that used to require a prefix passes `legacy_env_prefix`, and the +prefixed names keep working with a warning. + `detector/postprocess.py` and `detector/geometry.py` hold the parts of a detector that do *not* depend on the model: confidence filtering, per-class NMS, box/point clipping, and turning predictions into the loop's dataclasses. A node should import them rather than write its own — diff --git a/learning_loop_node/helpers/entrypoint.py b/learning_loop_node/helpers/entrypoint.py new file mode 100644 index 00000000..830c89ec --- /dev/null +++ b/learning_loop_node/helpers/entrypoint.py @@ -0,0 +1,75 @@ +"""The boilerplate every node's ``main.py`` repeats. + +A node entry point always does the same four things: read a handful of settings, build the +logic object, construct the node, and hand it to uvicorn. Only the middle two are the node's +own, so :func:`node_parser` and :func:`run_node` cover the rest. + +Settings come from a flag *or* an environment variable, because a node is configured on the +command line while developing and through the container environment in deployment. One +declaration gives both: ``--conf-threshold`` reads ``CONF_THRESHOLD``. + +``--host`` and ``--port`` are the exception — they read ``NODE_HOST`` and ``NODE_PORT`` rather +than the names their flags imply, because the bare ``HOST`` already means *the loop's address* +and a node adopting it would hand it to uvicorn and fail to bind. +""" + +import logging +import os +from argparse import Action, Namespace + +import configargparse +import uvicorn + + +def node_parser(*, description: str, legacy_env_prefix: str = '') -> configargparse.ArgumentParser: + """Build the parser for a node, pre-loaded with the settings every node has. + + :param legacy_env_prefix: A prefix an earlier version of this node required, e.g. + ``'DFINE_DETECTOR_'``. Prefixed names are still honoured, with a warning, so a + deployment keeps working until it is updated. Leave empty for a node that has always + read unprefixed names. + """ + parser = _NodeArgumentParser(description=description, legacy_env_prefix=legacy_env_prefix) + parser.add_argument('--host', default='0.0.0.0', env_var='NODE_HOST', + help='Host interface to bind to') + parser.add_argument('--port', type=int, default=80, env_var='NODE_PORT', help='Port to bind to') + return parser + + +def run_node(app: str, args: Namespace) -> None: + """Serve the node. + + :param app: Import string of the node object, conventionally ``'main:node'``. + """ + reload = os.getenv('UVICORN_RELOAD', 'FALSE').lower() in ('true', '1') + logging.info('Uvicorn reload is set to: %s', reload) + uvicorn.run(app, host=args.host, port=args.port, lifespan='on', reload=reload) + + +class _NodeArgumentParser(configargparse.ArgumentParser): + """Parser whose every setting is also an environment variable named after its flag.""" + + def __init__(self, *, description: str, legacy_env_prefix: str) -> None: + super().__init__(description=description) + self.legacy_env_prefix = legacy_env_prefix + + def add_argument(self, *args, **kwargs) -> Action: # type: ignore[override] + action = super().add_argument(*args, **kwargs) + if getattr(action, 'env_var', None) is None and action.dest != 'help': + action.env_var = action.dest.upper() + return action + + def parse_args(self, *args, **kwargs) -> Namespace: # type: ignore[override] + self._adopt_legacy_env_vars() + return super().parse_args(*args, **kwargs) + + def _adopt_legacy_env_vars(self) -> None: + if not self.legacy_env_prefix: + return + for action in self._actions: + name = getattr(action, 'env_var', None) + legacy = self.legacy_env_prefix + name if name else None + if not legacy or name in os.environ or legacy not in os.environ: + continue + os.environ[name] = os.environ[legacy] + logging.warning('%s is deprecated and will stop being read; set %s instead', legacy, name) diff --git a/learning_loop_node/tests/unit/test_entrypoint.py b/learning_loop_node/tests/unit/test_entrypoint.py new file mode 100644 index 00000000..a478563b --- /dev/null +++ b/learning_loop_node/tests/unit/test_entrypoint.py @@ -0,0 +1,71 @@ +import pytest + +from ...helpers.entrypoint import node_parser + +MANAGED = ('WEIGHT_TYPE', 'DFINE_DETECTOR_WEIGHT_TYPE', 'HOST', 'NODE_HOST', 'NODE_PORT', 'PORT') + + +@pytest.fixture(autouse=True) +def clean_env(monkeypatch: pytest.MonkeyPatch): + """Every test starts without the variables it is about to set.""" + for name in MANAGED: + monkeypatch.delenv(name, raising=False) + + +def _parser(**kwargs): + parser = node_parser(description='a node', **kwargs) + parser.add_argument('--weight-type', default='FP16') + return parser + + +def test_every_node_gets_a_host_and_a_port(): + args = _parser().parse_args([]) + assert (args.host, args.port) == ('0.0.0.0', 80) + + +def test_a_flag_beats_everything(): + args = _parser().parse_args(['--weight-type', 'FP32']) + assert args.weight_type == 'FP32' + + +def test_a_setting_is_read_from_the_variable_named_after_its_flag(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv('WEIGHT_TYPE', 'FP32') + assert _parser().parse_args([]).weight_type == 'FP32' + + +def test_the_loop_own_host_is_never_mistaken_for_the_bind_address(monkeypatch: pytest.MonkeyPatch): + """HOST is the loop's address. Binding uvicorn to it would leave the node unreachable.""" + monkeypatch.setenv('HOST', 'preview.learning-loop.ai') + assert _parser().parse_args([]).host == '0.0.0.0' + + +def test_the_bind_address_has_a_name_of_its_own(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv('NODE_HOST', '127.0.0.1') + monkeypatch.setenv('NODE_PORT', '8080') + args = _parser().parse_args([]) + assert (args.host, args.port) == ('127.0.0.1', 8080) + + +def test_a_node_that_used_a_prefix_still_reads_it(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv('DFINE_DETECTOR_WEIGHT_TYPE', 'FP32') + parser = _parser(legacy_env_prefix='DFINE_DETECTOR_') + assert parser.parse_args([]).weight_type == 'FP32' + + +def test_the_prefixed_name_warns_which_one_to_use_instead(monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture): + monkeypatch.setenv('DFINE_DETECTOR_WEIGHT_TYPE', 'FP32') + _parser(legacy_env_prefix='DFINE_DETECTOR_').parse_args([]) + assert 'DFINE_DETECTOR_WEIGHT_TYPE' in caplog.text + assert 'WEIGHT_TYPE' in caplog.text + + +def test_the_current_name_wins_over_the_prefixed_one(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv('DFINE_DETECTOR_WEIGHT_TYPE', 'FP32') + monkeypatch.setenv('WEIGHT_TYPE', 'FP16') + assert _parser(legacy_env_prefix='DFINE_DETECTOR_').parse_args([]).weight_type == 'FP16' + + +def test_a_node_without_a_legacy_prefix_ignores_prefixed_names(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv('DFINE_DETECTOR_WEIGHT_TYPE', 'FP32') + assert _parser().parse_args([]).weight_type == 'FP16' diff --git a/pyproject.toml b/pyproject.toml index 47409dff..d05073f5 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -13,6 +13,7 @@ dependencies = [ "python-socketio>=5.16.2,<6.0.0", "aiofiles>=0.7.0", "python-multipart>=0.0.31", + "configargparse>=1.7.1", "psutil>=5.9.0,<8.0.0", "numpy>=2.0,<3.0", "Pillow>=12.3.0,<13.0.0", diff --git a/uv.lock b/uv.lock index da0c0bf7..1b373e19 100644 --- a/uv.lock +++ b/uv.lock @@ -449,6 +449,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/d1/d6/3965ed04c63042e047cb6a3e6ed1a63a35087b6a609aa3a15ed8ac56c221/colorama-0.4.6-py2.py3-none-any.whl", hash = "sha256:4f1d9991f5acc0ca119f9d443620b77f9d6b33703e51011c16baf57afb285fc6", size = 25335, upload-time = "2022-10-25T02:36:20.889Z" }, ] +[[package]] +name = "configargparse" +version = "1.7.5" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/3f/0b/30328302903c55218ffc5199646d0e9d28348ff26c02ba77b2ffc58d294a/configargparse-1.7.5.tar.gz", hash = "sha256:e3f9a7bb6be34d66b2e3c4a2f58e3045f8dfae47b0dc039f87bcfaa0f193fb0f", size = 53548, upload-time = "2026-03-11T02:19:38.144Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/fe/19/3ba5e1b0bcc7b91aeab6c258afd70e4907d220fed3972febe38feb40db30/configargparse-1.7.5-py3-none-any.whl", hash = "sha256:1e63fdffedf94da9cd435fc13a1cd24777e76879dd2343912c1f871d4ac8c592", size = 27692, upload-time = "2026-03-11T02:19:36.442Z" }, +] + [[package]] name = "dacite" version = "1.9.2" @@ -781,6 +790,7 @@ source = { virtual = "." } dependencies = [ { name = "aiofiles" }, { name = "aiohttp" }, + { name = "configargparse" }, { name = "dacite" }, { name = "fastapi" }, { name = "httpx" }, @@ -813,6 +823,7 @@ requires-dist = [ { name = "aiofiles", specifier = ">=0.7.0" }, { name = "aiohttp", specifier = ">=3.14.3,<4.0.0" }, { name = "autopep8", marker = "extra == 'dev'", specifier = ">=2.0.2,<3.0.0" }, + { name = "configargparse", specifier = ">=1.7.1" }, { name = "dacite", specifier = ">=1.8.1,<2.0.0" }, { name = "debugpy", marker = "extra == 'dev'", specifier = ">=1.6.7.post1,<2.0.0" }, { name = "fastapi", specifier = ">=0.135,<1.0" }, From f3b3ca94accd6e84611feff76a2ab7bb022aae97 Mon Sep 17 00:00:00 2001 From: Kevin Heye Date: Thu, 20 Aug 2026 19:52:49 +0200 Subject: [PATCH 6/7] Share the framework-agnostic parts of a trainer Three things every trainer needs, none of which depend on the training framework, all of which existed only inside dfine_node or in worse form elsewhere: iterator_cpu_bound runs a training generator in a spawned process. A trainer that trains in-process needs this - the event loop must stay responsive and CUDA state must stay out of the node - and it is pure stdlib. find_batch_size probes for the largest power-of-two batch that fits, around a fits predicate the trainer supplies. The search is pure arithmetic; only the predicate needs a framework, so nothing here imports torch. Three repositories had three different implementations of this: a real probe, an estimate parsed out of torchinfo's summary text, and a doubling loop catching bare RuntimeError. is_out_of_memory exists because of that last one - catching RuntimeError treats every crash as a full card. macro_f1 scores the confusion matrix the loop stores, which was implemented three times over two different data shapes. _get_executor_error_from_log gets a default implementation, because yolov5_node and classification_node had byte-identical copies of it. The PEP 695 type parameters the generator used are rewritten as TypeVar and ParamSpec: this library supports Python 3.10 and dfine_node targets 3.12. Co-Authored-By: Claude Opus 5 (1M context) --- AGENTS.md | 9 ++ .../tests/unit/test_batch_size.py | 118 ++++++++++++++++++ learning_loop_node/tests/unit/test_metrics.py | 39 ++++++ .../tests/unit/test_subprocess.py | 54 ++++++++ learning_loop_node/trainer/batch_size.py | 86 +++++++++++++ learning_loop_node/trainer/metrics.py | 36 ++++++ learning_loop_node/trainer/subprocess.py | 93 ++++++++++++++ learning_loop_node/trainer/trainer_logic.py | 14 ++- 8 files changed, 447 insertions(+), 2 deletions(-) create mode 100644 learning_loop_node/tests/unit/test_batch_size.py create mode 100644 learning_loop_node/tests/unit/test_metrics.py create mode 100644 learning_loop_node/tests/unit/test_subprocess.py create mode 100644 learning_loop_node/trainer/batch_size.py create mode 100644 learning_loop_node/trainer/metrics.py create mode 100644 learning_loop_node/trainer/subprocess.py diff --git a/AGENTS.md b/AGENTS.md index c7877d5d..b6158d96 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -54,6 +54,15 @@ on 401/429) for the REST API, and a socket.io *client* for status updates and lo - **Annotator** — thin: it forwards the loop frontend's `handle_user_input` events into `AnnotatorLogic` and keeps a per-frontend history. +`trainer/subprocess.py`, `trainer/batch_size.py` and `trainer/metrics.py` hold the parts of a +trainer that are not framework-specific. `iterator_cpu_bound` runs a training generator in a +spawned process and yields its progress through a `maxsize=1` queue, so the event loop stays +responsive, CUDA state stays out of the node process, and the training can never run more than +one item ahead of the bookkeeping. `find_batch_size` probes for the largest power-of-two batch +that fits, around a `fits` predicate the trainer supplies — none of it imports torch, so a node +brings its own way of running a step. `macro_f1` scores the confusion matrix +`_get_new_best_training_state` returns. + `helpers/entrypoint.py` holds what every node's `main.py` repeats: `node_parser` builds a configargparse parser with `--host`/`--port`, `run_node` starts uvicorn. A setting is a flag *and* an environment variable from one declaration — `--conf-threshold` reads diff --git a/learning_loop_node/tests/unit/test_batch_size.py b/learning_loop_node/tests/unit/test_batch_size.py new file mode 100644 index 00000000..6b499a6a --- /dev/null +++ b/learning_loop_node/tests/unit/test_batch_size.py @@ -0,0 +1,118 @@ +from collections.abc import Callable + +import pytest + +from ...trainer.batch_size import ( + MIN_TRAIN_STEPS_PER_EPOCH, + batch_count, + dataset_limit, + find_batch_size, + is_out_of_memory, + no_gpu_batch_size, + smaller_pot, +) + + +def _recording(capacity: int) -> tuple[Callable[[int], bool], list[int]]: + """A `fits` predicate for a machine of `capacity`, plus the sizes it gets asked about.""" + calls: list[int] = [] + + def fits(batch_size: int) -> bool: + calls.append(batch_size) + return batch_size <= capacity + + return fits, calls + + +def _fits_up_to(capacity: int) -> Callable[[int], bool]: + fits, _ = _recording(capacity) + return fits + + +def test_the_search_doubles_up_to_the_limit(): + fits, calls = _recording(1024) + assert find_batch_size(fits, limit=16) == 16 + assert calls == [1, 2, 4, 8, 16] + + +def test_the_search_backs_off_to_the_last_size_that_fit(): + fits, calls = _recording(20) + assert find_batch_size(fits, limit=64) == 16 + assert calls == [1, 2, 4, 8, 16, 32] # probes 32, fails, keeps 16 + + +def test_a_limit_that_is_not_a_power_of_two_is_rounded_down(): + assert find_batch_size(_fits_up_to(1024), limit=48) == 32 + assert find_batch_size(_fits_up_to(1024), limit=1) == 1 + + +def test_a_machine_that_cannot_take_one_sample_is_an_error(): + fits, calls = _recording(0) + with pytest.raises(RuntimeError, match='batch size 1 does not fit'): + find_batch_size(fits, limit=64) + assert calls == [1], 'must give up instead of probing larger sizes' + + +@pytest.mark.parametrize('capacity', range(1, 130)) +def test_the_result_is_always_the_largest_power_of_two_that_fits(capacity: int): + assert find_batch_size(_fits_up_to(capacity), limit=512) == smaller_pot(capacity) + + +def test_equal_hardware_yields_an_equal_recipe(): + """Only powers of two, so two machines of similar size train identically.""" + assert find_batch_size(_fits_up_to(37), limit=512) == find_batch_size(_fits_up_to(39), limit=512) + + +def test_smaller_pot(): + assert [smaller_pot(n) for n in (1, 2, 3, 4, 7, 8, 15, 1293)] == [1, 2, 2, 4, 4, 8, 8, 1024] + with pytest.raises(ValueError, match='n must be >= 1'): + smaller_pot(0) + + +def test_batch_count_covers_the_whole_set_without_overshooting_by_a_batch(): + for sample_count in (1, 7, 8, 900, 5000): + for batch_size in (1, 2, 8, 64, 512): + covered = batch_count(sample_count, batch_size) * batch_size + assert covered >= sample_count + assert covered - sample_count < batch_size + + +def test_the_dataset_limit_keeps_enough_steps_per_epoch(): + for sample_count in (8, 20, 47, 100, 1000, 118_000): + limit = smaller_pot(dataset_limit(sample_count)) + assert sample_count // limit >= MIN_TRAIN_STEPS_PER_EPOCH + + +def test_the_dataset_limit_stays_usable_for_a_tiny_set(): + assert dataset_limit(7) == 1 + with pytest.raises(ValueError, match='sample_count must be >= 1'): + dataset_limit(0) + + +def test_memory_still_decides_below_the_dataset_limit(): + """The bound is a ceiling only: a card that fits just 2 keeps training at 2.""" + assert find_batch_size(_fits_up_to(2), limit=dataset_limit(20)) == 2 + + +def test_without_a_gpu_the_fallback_respects_the_limit(): + assert no_gpu_batch_size(1024, 'probe') == 8 + assert no_gpu_batch_size(2, 'probe') == 2 + + +@pytest.mark.parametrize('message', [ + 'CUDA out of memory. Tried to allocate 20.00 MiB', + 'cuDNN error: CUDNN_STATUS_ALLOC_FAILED', + 'CUDA error: unknown error', +]) +def test_allocation_failures_are_recognised_however_they_surface(message: str): + assert is_out_of_memory(RuntimeError(message)) + + +def test_a_real_bug_is_not_mistaken_for_a_full_card(): + """A trainer catching bare RuntimeError treats every crash as 'too big'; this does not.""" + assert not is_out_of_memory(RuntimeError('shape mismatch in forward pass')) + assert not is_out_of_memory(ValueError('bad config')) + + +def test_the_host_running_out_of_memory_counts_too(): + assert is_out_of_memory(MemoryError()) diff --git a/learning_loop_node/tests/unit/test_metrics.py b/learning_loop_node/tests/unit/test_metrics.py new file mode 100644 index 00000000..7c8537dc --- /dev/null +++ b/learning_loop_node/tests/unit/test_metrics.py @@ -0,0 +1,39 @@ +import pytest + +from ...trainer.metrics import category_f1, confusion_matrix_from_counts, macro_f1 + + +def test_the_score_averages_the_categories_instead_of_pooling_them(): + """Counts from a real epoch: face 248/19/59 and hand 39/54/73 - 74% pooled, 62% averaged.""" + assert macro_f1({'face': {'tp': 248, 'fp': 19, 'fn': 59}, + 'hand': {'tp': 39, 'fp': 54, 'fn': 73}}) == pytest.approx(0.622, abs=0.001) + assert macro_f1({'both': {'tp': 287, 'fp': 73, 'fn': 132}}) == pytest.approx(0.737, abs=0.001) + + +def test_a_rare_category_carries_the_same_weight(): + """A category the model never finds halves the score, however few instances it has.""" + assert macro_f1({'frequent': {'tp': 1000, 'fp': 0, 'fn': 0}, + 'rare': {'tp': 0, 'fp': 0, 'fn': 3}}) == pytest.approx(0.5) + + +def test_a_category_without_any_counts_scores_zero_instead_of_raising(): + assert macro_f1({'unseen': {'tp': 0, 'fp': 0, 'fn': 0}}) == 0.0 + assert category_f1({'tp': 0, 'fp': 0, 'fn': 0}) == 0.0 + + +def test_an_empty_confusion_matrix_scores_zero(): + assert macro_f1({}) == 0.0 + + +def test_a_perfect_category_scores_one(): + assert category_f1({'tp': 10, 'fp': 0, 'fn': 0}) == pytest.approx(1.0) + + +def test_precision_and_recall_are_balanced(): + assert category_f1({'tp': 5, 'fp': 5, 'fn': 0}) == category_f1({'tp': 5, 'fp': 0, 'fn': 5}) + + +def test_counts_can_be_assembled_from_separate_mappings(): + matrix = confusion_matrix_from_counts({'a': 3}, {'a': 1, 'b': 2}, {'b': 4}) + assert matrix == {'a': {'tp': 3, 'fp': 1, 'fn': 0}, + 'b': {'tp': 0, 'fp': 2, 'fn': 4}} diff --git a/learning_loop_node/tests/unit/test_subprocess.py b/learning_loop_node/tests/unit/test_subprocess.py new file mode 100644 index 00000000..f2517ced --- /dev/null +++ b/learning_loop_node/tests/unit/test_subprocess.py @@ -0,0 +1,54 @@ +import pytest + +from ...trainer.subprocess import iterator_cpu_bound + + +def counting(limit: int): + yield from range(limit) + + +def raising_after(limit: int): + yield from range(limit) + raise RuntimeError('the training crashed') + + +def empty(): + return iter(()) + + +async def _collect(iterator) -> list: + return [item async for item in iterator] + + +async def test_everything_the_generator_yields_arrives_in_order(): + async with iterator_cpu_bound(counting, 5) as iterator: + assert await _collect(iterator) == [0, 1, 2, 3, 4] + + +async def test_a_generator_that_yields_nothing_simply_finishes(): + async with iterator_cpu_bound(empty) as iterator: + assert await _collect(iterator) == [] + + +async def test_a_failure_in_the_process_is_raised_in_the_caller(): + """Otherwise a crashed training would look like a training that finished.""" + with pytest.raises(RuntimeError, match='the training crashed'): + async with iterator_cpu_bound(raising_after, 2) as iterator: + await _collect(iterator) + + +async def test_what_the_generator_produced_before_failing_still_arrives(): + received = [] + with pytest.raises(RuntimeError): + async with iterator_cpu_bound(raising_after, 3) as iterator: + async for item in iterator: + received.append(item) + assert received == [0, 1, 2] + + +async def test_leaving_early_does_not_leave_the_process_running(): + async with iterator_cpu_bound(counting, 1000) as iterator: + async for item in iterator: + if item == 2: + break + # the context manager killed and joined the process; reaching here without hanging is the test diff --git a/learning_loop_node/trainer/batch_size.py b/learning_loop_node/trainer/batch_size.py new file mode 100644 index 00000000..707bb8c6 --- /dev/null +++ b/learning_loop_node/trainer/batch_size.py @@ -0,0 +1,86 @@ +"""Choosing a batch size by probing, rather than configuring one. + +The largest batch that fits depends on the model, the image resolution and the card, so a +configured value is either wrong on some machines or too small on all of them. Probing answers +it directly: run a representative step at doubling sizes and keep the last one that survived. + +What "a representative step" means is the trainer's business — it supplies a ``fits`` predicate +that runs a real step and reports whether it ran out of memory. Everything here is that +predicate's scaffolding, and none of it imports a deep-learning framework. + +Only powers of two are visited, so equal hardware yields an equal recipe. That matters when a +trainer scales its learning rate by the batch size: a probe that returned 37 on one machine and +39 on another would make the two trainings quietly different. + +Adapted from PyTorch Lightning's ``BatchSizeFinder`` (power-scaling mode). +Copyright The Lightning AI team. Licensed under the Apache License, Version 2.0. +https://github.com/Lightning-AI/pytorch-lightning +""" + +import logging +from collections.abc import Callable + +logger = logging.getLogger(__name__) + +NO_GPU_BATCH_SIZE = 8 +"""What to fall back on with no GPU to probe: enough to make progress, small enough to fit.""" + +MIN_TRAIN_STEPS_PER_EPOCH = 8 +"""Below this an epoch is one or two optimizer steps, which trains poorly whatever the card.""" + + +def find_batch_size(fits: Callable[[int], bool], *, limit: int) -> int: + """Return the largest power-of-two batch size that fits, never exceeding ``limit``. + + :param fits: Runs a representative probe; ``False`` on out-of-memory. + :raises RuntimeError: If not even a batch size of 1 fits. + """ + limit = smaller_pot(limit) + if not fits(1): + raise RuntimeError('batch size 1 does not fit in memory') + + size = 1 + while size < limit and fits(size * 2): + size *= 2 + + return size + + +def dataset_limit(sample_count: int) -> int: + """The batch size ceiling that still leaves ``MIN_TRAIN_STEPS_PER_EPOCH`` steps per epoch.""" + if sample_count < 1: + raise ValueError(f'sample_count must be >= 1, got {sample_count}') + return max(1, sample_count // MIN_TRAIN_STEPS_PER_EPOCH) + + +def no_gpu_batch_size(limit: int, probe: str) -> int: + """The batch size to fall back on when there is no GPU to probe.""" + batch_size = min(smaller_pot(limit), NO_GPU_BATCH_SIZE) + logger.warning('%s: CUDA is unavailable; using batch size %d without probing', probe, batch_size) + return batch_size + + +def smaller_pot(n: int) -> int: + """The largest power of two that is <= ``n``.""" + if n < 1: + raise ValueError(f'n must be >= 1, got {n}') + return 1 << (n.bit_length() - 1) + + +def batch_count(sample_count: int, batch_size: int) -> int: + """How many batches a loader yields, the last one possibly short.""" + return -(-sample_count // batch_size) + + +def is_out_of_memory(exception: BaseException) -> bool: + """Whether the exception signals exhausted memory, on the GPU or the host. + + Allocation failures do not all surface as a framework's dedicated error type: cuDNN and + cuBLAS workspaces raise a plain ``RuntimeError``. Catching those by message is what keeps a + probe from mistaking a real bug for a full card — a trainer that catches bare + ``RuntimeError`` around its probe silently treats every crash as "too big". + """ + if isinstance(exception, MemoryError): + return True + message = str(exception).lower() + return any(text in message for text in ('out of memory', 'alloc_failed', 'cuda error: unknown error')) diff --git a/learning_loop_node/trainer/metrics.py b/learning_loop_node/trainer/metrics.py new file mode 100644 index 00000000..aee553b8 --- /dev/null +++ b/learning_loop_node/trainer/metrics.py @@ -0,0 +1,36 @@ +"""Scoring a training from the confusion matrix the loop stores. + +The loop keeps one ``{'tp': .., 'fp': .., 'fn': ..}`` per category, which is what +``TrainerLogicGeneric._get_new_best_training_state`` returns. Deciding whether an epoch beat +the previous best means reducing that to one number, and every trainer needs the same one. +""" + +import statistics +from collections.abc import Mapping + + +def macro_f1(confusion_matrix: Mapping[str, Mapping[str, int]]) -> float: + """The unweighted mean of the per-category F1 scores, as the loop's UI shows it by default. + + Unweighted is the point: a rare category the model never finds drags the score down as much + as a frequent one, where pooling the counts first would hide it. + """ + scores = [category_f1(counts) for counts in confusion_matrix.values()] + return statistics.mean(scores) if scores else 0.0 + + +def category_f1(counts: Mapping[str, int]) -> float: + """The F1 score of one category, 0.0 where it is undefined rather than a division error.""" + tp, fp, fn = counts['tp'], counts['fp'], counts['fn'] + precision = tp / (tp + fp) if tp + fp else 0.0 + recall = tp / (tp + fn) if tp + fn else 0.0 + return 2 * precision * recall / (precision + recall) if precision + recall else 0.0 + + +def confusion_matrix_from_counts(true_positives: Mapping[str, int], false_positives: Mapping[str, int], + false_negatives: Mapping[str, int]) -> dict[str, dict[str, int]]: + """Assemble the loop's confusion matrix from three per-category count mappings.""" + return {category: {'tp': true_positives.get(category, 0), + 'fp': false_positives.get(category, 0), + 'fn': false_negatives.get(category, 0)} + for category in {*true_positives, *false_positives, *false_negatives}} diff --git a/learning_loop_node/trainer/subprocess.py b/learning_loop_node/trainer/subprocess.py new file mode 100644 index 00000000..dbfa181b --- /dev/null +++ b/learning_loop_node/trainer/subprocess.py @@ -0,0 +1,93 @@ +"""Run a blocking, CPU-bound generator in its own process without blocking the event loop. + +A trainer that trains in-process has a problem: the training must not stall the node, and CUDA +state must stay out of the node process so a crashed training cannot take the node with it. +:func:`iterator_cpu_bound` runs the generator in a spawned process and yields what it produces +through a ``maxsize=1`` queue, so the producer can never run more than one item ahead of the +bookkeeping that consumes it — which is what lets a trainer alternate between two model files +and know the one it is copying is not being rewritten. + +Exceptions raised inside the process are re-raised in the caller, and the process is killed if +the caller leaves the context early. +""" +from __future__ import annotations + +import asyncio +import multiprocessing +import queue +import traceback +from collections.abc import AsyncGenerator, Callable, Iterator +from contextlib import asynccontextmanager +from multiprocessing.queues import Queue as MPQueue +from typing import Any, ParamSpec, TypeVar + +T = TypeVar('T') +P = ParamSpec('P') + + +class IteratorDone: + pass + + +def _iterator_wrapper( + it: Callable[..., Iterator[T]], + state_queue: MPQueue[T | Exception | IteratorDone], + args: tuple[Any, ...], + kwargs: dict[str, Any], +) -> None: + try: + for data in it(*args, **kwargs): + state_queue.put(data) + except Exception as e: + print(traceback.format_exc()) + state_queue.put(e) + + state_queue.put(IteratorDone()) + + +async def _iterator_cpu_bound_inner( + it: Callable[P, Iterator[T]], + *args: P.args, + **kwargs: P.kwargs, +) -> AsyncGenerator[T, None]: + state_queue: MPQueue[T | Exception | IteratorDone] = multiprocessing.Queue(maxsize=1) + process = multiprocessing.Process( + target=_iterator_wrapper, + args=(it, state_queue, args, kwargs), + name='iterator_cpu_bound', + ) + + process.start() + + try: + while True: + try: + item = await asyncio.to_thread(state_queue.get, True, 0.5) + except queue.Empty: + if not process.is_alive(): + break + continue + match item: + case IteratorDone(): + break + case Exception() as e: + raise e + case _ as other: + yield other + finally: + if process.is_alive(): + process.kill() + process.join() + + +@asynccontextmanager +async def iterator_cpu_bound( + it: Callable[P, Iterator[T]], + *args: P.args, + **kwargs: P.kwargs, +) -> AsyncGenerator[AsyncGenerator[T, None], None]: + iterator = _iterator_cpu_bound_inner(it, *args, **kwargs) + try: + yield iterator + finally: + await asyncio.shield(iterator.aclose()) diff --git a/learning_loop_node/trainer/trainer_logic.py b/learning_loop_node/trainer/trainer_logic.py index c5a9bade..9827b41b 100644 --- a/learning_loop_node/trainer/trainer_logic.py +++ b/learning_loop_node/trainer/trainer_logic.py @@ -178,9 +178,19 @@ async def _resume(self) -> None: '''Is called when self.can_resume() returns True. One may resume the training on a previously trained model stored by self.on_model_published(basic_model).''' - @abstractmethod def _get_executor_error_from_log(self) -> Optional[str]: - '''Should be used to provide error informations to the Learning Loop by extracting data from self.executor.get_log().''' + '''Reports what went wrong to the Learning Loop by reading self.executor's log. + + The default recognises the CUDA failures every trainer hits. Override to add messages a + particular training framework produces, and call super() to keep these.''' + if self._executor is None: + return None + for line in self._executor.get_log_by_lines(tail=50): + if 'CUDA out of memory' in line: + return 'graphics card is out of memory' + if 'CUDA error: invalid device ordinal' in line: + return 'graphics card not found' + return None @abstractmethod async def _detect(self, model_information: ModelInformation, images: List[str], model_folder: str) -> List[Detections]: From e66f05a3048b94b406d32ff9ab16650da19a312e Mon Sep 17 00:00:00 2001 From: Kevin Heye Date: Mon, 24 Aug 2026 17:28:14 +0200 Subject: [PATCH 7/7] Ship the node test helpers with the library learning_loop_node/tests/ is excluded from the wheel, so nothing in it is importable by a node repository. That is why the node repositories have almost no tests: every one of them would have to reinvent TestingTrainerLogic, the data_folder fixture and the condition poller before writing its first assertion. dfine_node/trainer/tests/ is the only offline suite in any of them. The reusable half now lives in learning_loop_node/testing/, a real package that ships: - helpers.py - condition, update_attributes, get_files_in_folder, unzip, get_latest_model_id - detections.py - get_dummy_detections, get_dummy_metadata - fixtures.py - the data_folder and clear_loggers fixtures, as an opt-in pytest plugin - trainer.py - TestingTrainerLogic, create_active_training_file, assert_training_state - detector.py - TestingDetectorLogic, TestingDetectorFactory tests/ keeps only the library's own suites and stays excluded, as before. A node opts into the fixtures with pytest_plugins = ['learning_loop_node.testing.fixtures'] which is deliberately not a pytest11 entry point: data_folder is autouse and wipes a directory, so no project should get it merely by installing the library. For the same reason testing/__init__.py does not import fixtures, and the package stays importable without pytest. The data_folder fixture existed in six copies, and two of them - mock_trainer's and mock_detector's - had drifted to not create the folder they point at. They all become the one that does. Destructive fixtures now refuse to run against production. LoopCommunicator reads environment_reader.host(default='learning-loop.ai') and run_tests.sh sourced .env with no guard, so a missing LOOP_HOST aimed project creation and deletion at real customer data. assert_not_production_loop() is called by every fixture that generates or deletes a project, and run_tests.sh fails before pytest starts. CI is unaffected - the workflows pin preview.learning-loop.ai. assert_training_state loses a debug except-Exception branch that logged '##### was ist das hier?' and re-raised; it is public API now. Verified: 213 unit tests pass (from 197); every other suite still collects; the built wheel contains learning_loop_node/testing/ and no learning_loop_node/ tests/; in a fresh Python 3.10 venv holding only that wheel, the package imports with pytest absent, and with pytest added a throwaway project uses the fixtures plugin and the helpers exactly as documented. Repository-wide ruff goes from 717 findings to 713, with no rule higher than before. This repository has no pre-commit config and no pylint or pyright in its environment, so ruff and pytest are what ran. Co-Authored-By: Claude Opus 5 --- AGENTS.md | 27 ++- README.md | 46 +++++- .../app_code/tests/conftest.py | 2 + learning_loop_node/testing/__init__.py | 38 +++++ learning_loop_node/testing/detections.py | 36 ++++ .../detector.py} | 8 +- learning_loop_node/testing/fixtures.py | 43 +++++ .../test_helper.py => testing/helpers.py} | 59 +++---- .../trainer.py} | 25 ++- .../tests/annotator/conftest.py | 32 +--- learning_loop_node/tests/detector/conftest.py | 30 +--- .../detector/test_client_communication.py | 2 +- .../tests/detector/test_relevance_filter.py | 2 +- learning_loop_node/tests/general/conftest.py | 33 +--- .../tests/general/test_downloader.py | 8 +- learning_loop_node/tests/trainer/conftest.py | 32 +--- .../tests/trainer/state_helper.py | 22 --- .../trainer/states/test_state_cleanup.py | 3 +- .../trainer/states/test_state_detecting.py | 4 +- .../states/test_state_download_train_model.py | 6 +- .../trainer/states/test_state_prepare.py | 3 +- .../test_state_sync_confusion_matrix.py | 3 +- .../tests/trainer/states/test_state_train.py | 4 +- .../states/test_state_upload_detections.py | 4 +- .../trainer/states/test_state_upload_model.py | 3 +- .../tests/trainer/test_errors.py | 3 +- .../tests/trainer/test_trainer_states.py | 2 +- .../tests/unit/test_testing_package.py | 156 ++++++++++++++++++ mock_detector/app_code/tests/conftest.py | 10 +- mock_trainer/app_code/tests/conftest.py | 14 +- .../app_code/tests/test_detections.py | 5 +- pyproject.toml | 5 + run_tests.sh | 19 ++- uv.lock | 8 +- 34 files changed, 451 insertions(+), 246 deletions(-) create mode 100644 learning_loop_node/testing/__init__.py create mode 100644 learning_loop_node/testing/detections.py rename learning_loop_node/{tests/detector/testing_detector.py => testing/detector.py} (76%) create mode 100644 learning_loop_node/testing/fixtures.py rename learning_loop_node/{tests/test_helper.py => testing/helpers.py} (52%) rename learning_loop_node/{tests/trainer/testing_trainer_logic.py => testing/trainer.py} (80%) delete mode 100644 learning_loop_node/tests/trainer/state_helper.py create mode 100644 learning_loop_node/tests/unit/test_testing_package.py diff --git a/AGENTS.md b/AGENTS.md index b6158d96..60086328 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -14,9 +14,13 @@ environment variables, the node types and how to write a node against them. machinery, `trainer/`, `detector/`, `annotation/` for the per-type base logic, `data_classes/` and `enums/` for the wire types, `loop_communication.py` and `data_exchanger.py` for the loop-facing HTTP and socket.io traffic. -- `learning_loop_node/tests/` — `unit`, `annotator`, `detector`, `trainer` and `general` suites. - Only `unit` runs without a Learning Loop; it covers the pure helpers such as - `detector/postprocess.py` and `detector/geometry.py`. +- `learning_loop_node/testing/` — test helpers that **ship in the wheel**, so a node repository + can import them: `TestingTrainerLogic`, `TestingDetectorLogic`, the dummy detections, the + `condition` poller and the `fixtures` pytest plugin. Anything reusable belongs here, not in + `tests/`. +- `learning_loop_node/tests/` — the library's own suites (`unit`, `annotator`, `detector`, + `trainer`, `general`), excluded from the wheel. Only `unit` runs without a Learning Loop; it + covers the pure helpers such as `detector/postprocess.py` and `detector/geometry.py`. - `mock_trainer/`, `mock_detector/`, `mock_annotator/` — reference implementations with their own tests. They are what `loop`'s CI runs against, so they are also the best template for a new node. - `demo_segmentation_tool/` — a worked annotator example. @@ -112,9 +116,20 @@ python -m pytest learning_loop_node/tests/trainer -v -k # one t ``` An autouse fixture repoints `GLOBALS.data_folder` at `/tmp/learning_loop_lib_data` and wipes it -around every test, so tests never touch `/data`. The `general` suite generates and deletes a real -`zauberzeug/pytest_nodelib_general` project on the loop; the detector suite starts the node in a -forked uvicorn process on `GLOBALS.detector_port`. +around every test, so tests never touch `/data`. It lives in `learning_loop_node/testing/fixtures.py` +and every suite pulls it in with + +```python +from ...testing.fixtures import clear_loggers, data_folder # noqa: F401 +``` + +The `general` suite generates and deletes a real `zauberzeug/pytest_nodelib_general` project on the +loop; the detector suite starts the node in a forked uvicorn process on `GLOBALS.detector_port`. + +Every fixture that creates or deletes a project calls `assert_not_production_loop()` first, and +`run_tests.sh` refuses to start without `LOOP_HOST`. `LoopCommunicator` defaults to +`learning-loop.ai` — production — so a forgotten `.env` would otherwise aim those fixtures at real +customer data. CI is unaffected: the workflows pin `preview.learning-loop.ai`. There is no `.pre-commit-config.yaml` here and no ruff in the project environment. Lint with: diff --git a/README.md b/README.md index 6372a5b3..1bb7687f 100644 --- a/README.md +++ b/README.md @@ -36,11 +36,49 @@ You can configure connection to our Learning Loop by specifying the following en Note that organization and project IDs are always lower case and may differ from the names in the Learning Loop which can have uppercase letters. -#### Testing +#### Testing your own node -We use github actions for CI. Tests can also be executed locally by running -`LOOP_HOST=XXXXXXXX LOOP_USERNAME=XXXXXXXX LOOP_PASSWORD=XXXXXXXX python -m pytest -v` -from learning_loop_node/learning_loop_node +`learning_loop_node.testing` ships with the library, so a node repository can test itself without +a Learning Loop: + +```bash +pip install learning_loop_node[testing] +``` + +```python +# tests/conftest.py +pytest_plugins = ['learning_loop_node.testing.fixtures'] +``` + +That gives every test an autouse `data_folder` fixture, which repoints `GLOBALS.data_folder` at +`/tmp/learning_loop_lib_data` and wipes it around each test, so a test never touches `/data`. The +path is shared, so do not run two such suites at once. The module also provides: + +| | | +| --- | --- | +| `TestingTrainerLogic` | a `TrainerLogic` that trains a sleeping subprocess — drive the state machine without a framework | +| `TestingDetectorLogic`, `TestingDetectorFactory` | a detector that returns fixed detections | +| `get_dummy_detections()`, `get_dummy_metadata()` | one detection of every type | +| `condition(...)` | await a predicate with a timeout | +| `assert_training_state(...)`, `create_active_training_file(...)` | drive and assert the trainer state machine | +| `assert_not_production_loop()` | call this first in any fixture that creates or deletes a project | + +`assert_not_production_loop()` matters because `LoopCommunicator` falls back to `learning-loop.ai` +when neither `LOOP_HOST` nor `HOST` is set — a forgotten `.env` would otherwise point a destructive +fixture at production. + +#### Testing this library + +We use github actions for CI. Locally, the `unit` suite needs nothing at all: + +```bash +python -m pytest learning_loop_node/tests/unit -v +``` + +The other suites need a reachable Learning Loop and its credentials in a local `.env` +(`LOOP_HOST`, `LOOP_USERNAME`, `LOOP_PASSWORD`); `./run_tests.sh` runs them all and refuses to +start without `LOOP_HOST`. Each suite carries its own `pytest.ini`, so always pass a path inside +one suite — a bare `pytest` from the repository root picks up no config. ## Detector Node diff --git a/demo_segmentation_tool/app_code/tests/conftest.py b/demo_segmentation_tool/app_code/tests/conftest.py index 83be5afe..daaa2dc4 100644 --- a/demo_segmentation_tool/app_code/tests/conftest.py +++ b/demo_segmentation_tool/app_code/tests/conftest.py @@ -4,6 +4,7 @@ import pytest from learning_loop_node.loop_communication import LoopCommunicator +from learning_loop_node.testing import assert_not_production_loop @pytest.fixture() @@ -15,6 +16,7 @@ async def glc(): @pytest.fixture(autouse=True, scope='function') async def setup_test_project(glc: LoopCommunicator): + assert_not_production_loop() await glc.delete("/zauberzeug/projects/pytest_dst?keep_images=true") await asyncio.sleep(1) project_configuration = { diff --git a/learning_loop_node/testing/__init__.py b/learning_loop_node/testing/__init__.py new file mode 100644 index 00000000..66b755a9 --- /dev/null +++ b/learning_loop_node/testing/__init__.py @@ -0,0 +1,38 @@ +"""Helpers for testing a node, shipped with the library. + +Unlike `learning_loop_node.tests` — the library's own suite, which is excluded from the wheel — +everything here is importable by a node repository. + +The pytest fixtures live in `learning_loop_node.testing.fixtures` and are not re-exported here, +so that this package stays importable without pytest. Opt into them from a `conftest.py`: + + pytest_plugins = ['learning_loop_node.testing.fixtures'] +""" + +from .detections import get_dummy_detections, get_dummy_metadata +from .detector import TestingDetectorFactory, TestingDetectorLogic +from .helpers import ( + assert_not_production_loop, + condition, + get_files_in_folder, + get_latest_model_id, + unzip, + update_attributes, +) +from .trainer import TestingTrainerLogic, assert_training_state, create_active_training_file + +__all__ = [ + 'TestingDetectorFactory', + 'TestingDetectorLogic', + 'TestingTrainerLogic', + 'assert_not_production_loop', + 'assert_training_state', + 'condition', + 'create_active_training_file', + 'get_dummy_detections', + 'get_dummy_metadata', + 'get_files_in_folder', + 'get_latest_model_id', + 'unzip', + 'update_attributes', +] diff --git a/learning_loop_node/testing/detections.py b/learning_loop_node/testing/detections.py new file mode 100644 index 00000000..1667d8bc --- /dev/null +++ b/learning_loop_node/testing/detections.py @@ -0,0 +1,36 @@ +from ..data_classes import ( + BoxDetection, + ClassificationDetection, + Detections, + Point, + PointDetection, + SegmentationDetection, + Shape, +) +from ..data_classes.image_metadata import ImageMetadata + + +def get_dummy_detections() -> Detections: + return Detections( + box_detections=[ + BoxDetection(category_name='some_category_name', x=1, y=2, height=3, width=4, + model_name='some_model', confidence=.42, category_id='some_id')], + point_detections=[ + PointDetection(category_name='some_category_name_2', x=10, y=12, + model_name='some_model', confidence=.42, category_id='some_id_2')], + segmentation_detections=[ + SegmentationDetection(category_name='some_category_name_3', + shape=Shape(points=[Point(x=1, y=1)]), + model_name='some_model', confidence=.42, + category_id='some_id_3')], + classification_detections=[ + ClassificationDetection(category_name='some_category_name_4', model_name='some_model', + confidence=.42, category_id='some_id_4')]) + + +def get_dummy_metadata() -> ImageMetadata: + detections = get_dummy_detections() + return ImageMetadata(box_detections=detections.box_detections, + point_detections=detections.point_detections, + segmentation_detections=detections.segmentation_detections, + classification_detections=detections.classification_detections) diff --git a/learning_loop_node/tests/detector/testing_detector.py b/learning_loop_node/testing/detector.py similarity index 76% rename from learning_loop_node/tests/detector/testing_detector.py rename to learning_loop_node/testing/detector.py index a7841c00..d53f899a 100644 --- a/learning_loop_node/tests/detector/testing_detector.py +++ b/learning_loop_node/testing/detector.py @@ -3,11 +3,9 @@ import numpy as np -from learning_loop_node.data_classes import ImagesMetadata, ModelInformation - -from ...data_classes import ImageMetadata -from ...detector.detector_logic import DetectorLogic -from ..test_helper import get_dummy_metadata +from ..data_classes import ImageMetadata, ImagesMetadata, ModelInformation +from ..detector.detector_logic import DetectorLogic +from .detections import get_dummy_metadata class TestingDetectorLogic(DetectorLogic): diff --git a/learning_loop_node/testing/fixtures.py b/learning_loop_node/testing/fixtures.py new file mode 100644 index 00000000..d478e764 --- /dev/null +++ b/learning_loop_node/testing/fixtures.py @@ -0,0 +1,43 @@ +"""Pytest fixtures shared by the library's own suite and by node repositories. + +Opt in from a `conftest.py`: + + pytest_plugins = ['learning_loop_node.testing.fixtures'] + +The fixtures are deliberately *not* registered as a `pytest11` entry point: `data_folder` is +autouse and wipes a directory, which no project should get merely by installing the library. +""" + +import logging +import os +import shutil + +import pytest + +from ..globals import GLOBALS + +TEST_DATA_FOLDER = '/tmp/learning_loop_lib_data' + + +@pytest.fixture(autouse=True, scope='session') +def clear_loggers(): + """Remove handlers from all loggers""" + # see https://github.com/pytest-dev/pytest/issues/5502 + yield + + loggers = [logging.getLogger()] + list(logging.Logger.manager.loggerDict.values()) + for logger in loggers: + if not isinstance(logger, logging.Logger): + continue + handlers = getattr(logger, 'handlers', []) + for handler in handlers: + logger.removeHandler(handler) + + +@pytest.fixture(autouse=True, scope='function') +def data_folder(): + GLOBALS.data_folder = TEST_DATA_FOLDER + shutil.rmtree(GLOBALS.data_folder, ignore_errors=True) + os.makedirs(GLOBALS.data_folder, exist_ok=True) + yield + shutil.rmtree(GLOBALS.data_folder, ignore_errors=True) diff --git a/learning_loop_node/tests/test_helper.py b/learning_loop_node/testing/helpers.py similarity index 52% rename from learning_loop_node/tests/test_helper.py rename to learning_loop_node/testing/helpers.py index 808c65f7..ab39d5f5 100644 --- a/learning_loop_node/tests/test_helper.py +++ b/learning_loop_node/testing/helpers.py @@ -6,18 +6,31 @@ from glob import glob from typing import Callable -from ..data_classes import ( - BoxDetection, - ClassificationDetection, - Detections, - Point, - PointDetection, - SegmentationDetection, - Shape, -) -from ..data_classes.image_metadata import ImageMetadata +from ..helpers import environment_reader from ..loop_communication import LoopCommunicator +PRODUCTION_LOOP_HOST = 'learning-loop.ai' + + +def assert_not_production_loop() -> None: + """Refuse to run a destructive test fixture against the production Learning Loop. + + `LoopCommunicator` falls back to the production host when neither `LOOP_HOST` nor `HOST` is + set, so a forgotten `.env` aims project creation and deletion at real customer data. Every + fixture that creates or deletes a project should call this first. + """ + host = environment_reader.host() + if not host: + raise AssertionError( + 'LOOP_HOST is not set. Destructive test fixtures refuse to run, because ' + f'LoopCommunicator would fall back to the production loop at {PRODUCTION_LOOP_HOST}. ' + 'Point LOOP_HOST at a test instance, for example preview.learning-loop.ai.') + if host == PRODUCTION_LOOP_HOST: + raise AssertionError( + f'LOOP_HOST is {host}, the production Learning Loop. Destructive test fixtures create ' + 'and delete projects, so they refuse to run here. Point LOOP_HOST at a test instance, ' + 'for example preview.learning-loop.ai.') + def get_files_in_folder(folder: str): files = [entry for entry in glob(f'{folder}/**/*', recursive=True) if os.path.isfile(entry)] @@ -68,29 +81,3 @@ def _update_attribute_class_instance(obj, **kwargs) -> None: def _update_attribute_dict(obj: dict, **kwargs) -> None: for key, value in kwargs.items(): obj[key] = value - - -def get_dummy_detections() -> Detections: - return Detections( - box_detections=[ - BoxDetection(category_name='some_category_name', x=1, y=2, height=3, width=4, - model_name='some_model', confidence=.42, category_id='some_id')], - point_detections=[ - PointDetection(category_name='some_category_name_2', x=10, y=12, - model_name='some_model', confidence=.42, category_id='some_id_2')], - segmentation_detections=[ - SegmentationDetection(category_name='some_category_name_3', - shape=Shape(points=[Point(x=1, y=1)]), - model_name='some_model', confidence=.42, - category_id='some_id_3')], - classification_detections=[ - ClassificationDetection(category_name='some_category_name_4', model_name='some_model', - confidence=.42, category_id='some_id_4')]) - - -def get_dummy_metadata() -> ImageMetadata: - detections = get_dummy_detections() - return ImageMetadata(box_detections=detections.box_detections, - point_detections=detections.point_detections, - segmentation_detections=detections.segmentation_detections, - classification_detections=detections.classification_detections) diff --git a/learning_loop_node/tests/trainer/testing_trainer_logic.py b/learning_loop_node/testing/trainer.py similarity index 80% rename from learning_loop_node/tests/trainer/testing_trainer_logic.py rename to learning_loop_node/testing/trainer.py index b3f56a9d..7e8cad10 100644 --- a/learning_loop_node/tests/trainer/testing_trainer_logic.py +++ b/learning_loop_node/testing/trainer.py @@ -2,8 +2,29 @@ import time from typing import Dict, List, Optional -from ...data_classes import Context, Detections, ModelInformation, PretrainedModel, TrainingStateData -from ...trainer.trainer_logic import TrainerLogic +from ..data_classes import ( + Context, + Detections, + ModelInformation, + PretrainedModel, + Training, + TrainingStateData, +) +from ..trainer.trainer_logic import TrainerLogic +from .helpers import condition, update_attributes + + +def create_active_training_file(trainer: TrainerLogic, **kwargs) -> None: + update_attributes(trainer._training, **kwargs) # pylint: disable=protected-access + trainer.node.last_training_io.save(training=trainer.training) + + +async def assert_training_state(training: Training, state: str, timeout: float, interval: float) -> None: + try: + await condition(lambda: training.training_state == state, timeout=timeout, interval=interval) + except TimeoutError as exc: + msg = f"Trainer state should be '{state}' after {timeout} seconds, but is {training.training_state}" + raise AssertionError(msg) from exc class TestingTrainerLogic(TrainerLogic): diff --git a/learning_loop_node/tests/annotator/conftest.py b/learning_loop_node/tests/annotator/conftest.py index 97f62ea5..3536bc86 100644 --- a/learning_loop_node/tests/annotator/conftest.py +++ b/learning_loop_node/tests/annotator/conftest.py @@ -1,19 +1,17 @@ import asyncio import logging -import os -import shutil - -# ====================================== REDUNDANT FIXTURES IN ALL CONFTESTS ! ====================================== import sys import pytest -from ...globals import GLOBALS from ...loop_communication import LoopCommunicator +from ...testing import assert_not_production_loop +from ...testing.fixtures import clear_loggers, data_folder # noqa: F401 pylint: disable=unused-import @pytest.fixture() async def setup_test_project(): # pylint: disable=redefined-outer-name + assert_not_production_loop() loop_communicator = LoopCommunicator() try: await loop_communicator.delete("/zauberzeug/projects/pytest_nodelib_annotator?keep_images=true", timeout=10) @@ -29,27 +27,3 @@ async def setup_test_project(): # pylint: disable=redefined-outer-name yield await loop_communicator.delete("/zauberzeug/projects/pytest_nodelib_annotator?keep_images=true", timeout=10) await loop_communicator.shutdown() - - -@pytest.fixture(autouse=True, scope='session') -def clear_loggers(): - """Remove handlers from all loggers""" - # see https://github.com/pytest-dev/pytest/issues/5502 - yield - - loggers = [logging.getLogger()] + list(logging.Logger.manager.loggerDict.values()) - for logger in loggers: - if not isinstance(logger, logging.Logger): - continue - handlers = getattr(logger, 'handlers', []) - for handler in handlers: - logger.removeHandler(handler) - - -@pytest.fixture(autouse=True, scope='function') -def data_folder(): - GLOBALS.data_folder = '/tmp/learning_loop_lib_data' - shutil.rmtree(GLOBALS.data_folder, ignore_errors=True) - os.makedirs(GLOBALS.data_folder, exist_ok=True) - yield - shutil.rmtree(GLOBALS.data_folder, ignore_errors=True) diff --git a/learning_loop_node/tests/detector/conftest.py b/learning_loop_node/tests/detector/conftest.py index 4be0b030..5b3c1ecc 100644 --- a/learning_loop_node/tests/detector/conftest.py +++ b/learning_loop_node/tests/detector/conftest.py @@ -4,7 +4,6 @@ import logging import multiprocessing import os -import shutil import socket from dataclasses import asdict from glob import glob @@ -22,7 +21,8 @@ from ...detector.detector_node import DetectorNode from ...detector.outbox import Outbox from ...globals import GLOBALS -from .testing_detector import TestingDetectorFactory +from ...testing import TestingDetectorFactory +from ...testing.fixtures import clear_loggers, data_folder # noqa: F401 pylint: disable=unused-import logging.basicConfig(level=logging.INFO) @@ -177,29 +177,3 @@ async def detector_node(): node = DetectorNode(name="test_node", detector_factory=MockDetectorFactory()) await node._build_and_swap_detector(model_dir) return node - -# ====================================== REDUNDANT FIXTURES IN ALL CONFTESTS ! ====================================== - - -@pytest.fixture(autouse=True, scope='session') -def clear_loggers(): - """Remove handlers from all loggers""" - # see https://github.com/pytest-dev/pytest/issues/5502 - yield - - loggers = [logging.getLogger()] + list(logging.Logger.manager.loggerDict.values()) - for logger in loggers: - if not isinstance(logger, logging.Logger): - continue - handlers = getattr(logger, 'handlers', []) - for handler in handlers: - logger.removeHandler(handler) - - -@pytest.fixture(autouse=True, scope='function') -def data_folder(): - GLOBALS.data_folder = '/tmp/learning_loop_lib_data' - shutil.rmtree(GLOBALS.data_folder, ignore_errors=True) - os.makedirs(GLOBALS.data_folder, exist_ok=True) - yield - shutil.rmtree(GLOBALS.data_folder, ignore_errors=True) diff --git a/learning_loop_node/tests/detector/test_client_communication.py b/learning_loop_node/tests/detector/test_client_communication.py index 9b80c6de..d18688fe 100644 --- a/learning_loop_node/tests/detector/test_client_communication.py +++ b/learning_loop_node/tests/detector/test_client_communication.py @@ -10,8 +10,8 @@ from ...data_classes import ModelInformation from ...detector.detector_node import DetectorNode, _ActiveDetector from ...globals import GLOBALS +from ...testing import TestingDetectorLogic from .conftest import get_outbox_files -from .testing_detector import TestingDetectorLogic file_path = os.path.abspath(__file__) test_image_path = os.path.join(os.path.dirname(file_path), 'test.jpg') diff --git a/learning_loop_node/tests/detector/test_relevance_filter.py b/learning_loop_node/tests/detector/test_relevance_filter.py index b9055b86..fd0493d0 100644 --- a/learning_loop_node/tests/detector/test_relevance_filter.py +++ b/learning_loop_node/tests/detector/test_relevance_filter.py @@ -7,8 +7,8 @@ from ...data_classes import BoxDetection, ImageMetadata, PointDetection from ...detector.detector_node import DetectorNode, _ActiveDetector +from ...testing import TestingDetectorLogic from .conftest import get_outbox_files -from .testing_detector import TestingDetectorLogic file_path = os.path.abspath(__file__) test_image_path = os.path.join(os.path.dirname(file_path), 'test.jpg') diff --git a/learning_loop_node/tests/general/conftest.py b/learning_loop_node/tests/general/conftest.py index 8cc34d1b..7389b837 100644 --- a/learning_loop_node/tests/general/conftest.py +++ b/learning_loop_node/tests/general/conftest.py @@ -1,20 +1,19 @@ import asyncio import logging -import os -import shutil import sys import pytest from ...data_classes import Context from ...data_exchanger import DataExchanger -from ...globals import GLOBALS from ...loop_communication import LoopCommunicator +from ...testing import assert_not_production_loop +from ...testing.fixtures import clear_loggers, data_folder # noqa: F401 pylint: disable=unused-import @pytest.fixture(autouse=True, scope='function') async def create_project_for_module(): - + assert_not_production_loop() loop_communicator = LoopCommunicator() try: await loop_communicator.delete("/zauberzeug/projects/pytest_nodelib_general", timeout=10) @@ -40,29 +39,3 @@ async def data_exchanger(): dx = DataExchanger(context, loop_communicator) yield dx await loop_communicator.shutdown() - -# ====================================== REDUNDANT FIXTURES IN ALL CONFTESTS ! ====================================== - - -@pytest.fixture(autouse=True, scope='session') -def clear_loggers(): - """Remove handlers from all loggers""" - # see https://github.com/pytest-dev/pytest/issues/5502 - yield - - loggers = [logging.getLogger()] + list(logging.Logger.manager.loggerDict.values()) - for logger in loggers: - if not isinstance(logger, logging.Logger): - continue - handlers = getattr(logger, 'handlers', []) - for handler in handlers: - logger.removeHandler(handler) - - -@pytest.fixture(autouse=True, scope='function') -def data_folder(): - GLOBALS.data_folder = '/tmp/learning_loop_lib_data' - shutil.rmtree(GLOBALS.data_folder, ignore_errors=True) - os.makedirs(GLOBALS.data_folder, exist_ok=True) - yield - shutil.rmtree(GLOBALS.data_folder, ignore_errors=True) diff --git a/learning_loop_node/tests/general/test_downloader.py b/learning_loop_node/tests/general/test_downloader.py index a88ac831..3b18c094 100644 --- a/learning_loop_node/tests/general/test_downloader.py +++ b/learning_loop_node/tests/general/test_downloader.py @@ -5,7 +5,7 @@ from ...data_exchanger import DataExchanger from ...globals import GLOBALS from ...helpers.misc import create_image_folder, create_project_folder, create_training_folder, delete_corrupt_images -from .. import test_helper +from ...testing import get_files_in_folder, get_latest_model_id # Used by all Nodes @@ -13,11 +13,11 @@ async def test_download_model(data_exchanger: DataExchanger): _, _, trainings_folder = create_needed_folders() - model_id = await test_helper.get_latest_model_id(project='pytest_nodelib_general') + model_id = await get_latest_model_id(project='pytest_nodelib_general') await data_exchanger.download_model(trainings_folder, Context(organization='zauberzeug', project='pytest_nodelib_general'), model_id, 'mocked') - files = test_helper.get_files_in_folder(GLOBALS.data_folder) + files = get_files_in_folder(GLOBALS.data_folder) assert len(files) == 3, str(files) file_1 = f'{GLOBALS.data_folder}/zauberzeug/pytest_nodelib_general/trainings/some_uuid/file_1.txt' @@ -43,7 +43,7 @@ async def test_download_images(data_exchanger: DataExchanger): _, image_folder, _ = create_needed_folders() image_ids = await data_exchanger.fetch_image_uuids() await data_exchanger.download_images(image_ids, image_folder) - files = test_helper.get_files_in_folder(GLOBALS.data_folder) + files = get_files_in_folder(GLOBALS.data_folder) assert len(files) == 3 diff --git a/learning_loop_node/tests/trainer/conftest.py b/learning_loop_node/tests/trainer/conftest.py index fbb5c8a2..88efdcc8 100644 --- a/learning_loop_node/tests/trainer/conftest.py +++ b/learning_loop_node/tests/trainer/conftest.py @@ -1,6 +1,5 @@ import logging import os -import shutil import socket from multiprocessing import log_to_stderr @@ -8,9 +7,9 @@ import pytest from ...data_classes import Context -from ...globals import GLOBALS +from ...testing import TestingTrainerLogic +from ...testing.fixtures import clear_loggers, data_folder # noqa: F401 pylint: disable=unused-import from ...trainer.trainer_node import TrainerNode -from .testing_trainer_logic import TestingTrainerLogic # pylint: disable=protected-access @@ -72,30 +71,3 @@ async def test_initialized_trainer(): def is_port_in_use(port): with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: return s.connect_ex(('localhost', port)) == 0 - - -# ====================================== REDUNDANT FIXTURES IN ALL CONFTESTS ! ====================================== - - -@pytest.fixture(autouse=True, scope='session') -def clear_loggers(): - """Remove handlers from all loggers""" - # see https://github.com/pytest-dev/pytest/issues/5502 - yield - - loggers = [logging.getLogger()] + list(logging.Logger.manager.loggerDict.values()) - for logger in loggers: - if not isinstance(logger, logging.Logger): - continue - handlers = getattr(logger, 'handlers', []) - for handler in handlers: - logger.removeHandler(handler) - - -@pytest.fixture(autouse=True, scope='function') -def data_folder(): - GLOBALS.data_folder = '/tmp/learning_loop_lib_data' - shutil.rmtree(GLOBALS.data_folder, ignore_errors=True) - os.makedirs(GLOBALS.data_folder, exist_ok=True) - yield - shutil.rmtree(GLOBALS.data_folder, ignore_errors=True) diff --git a/learning_loop_node/tests/trainer/state_helper.py b/learning_loop_node/tests/trainer/state_helper.py deleted file mode 100644 index 567ec629..00000000 --- a/learning_loop_node/tests/trainer/state_helper.py +++ /dev/null @@ -1,22 +0,0 @@ -import logging - -from learning_loop_node.tests.test_helper import condition, update_attributes -from learning_loop_node.trainer.trainer_logic import TrainerLogic - -from ...data_classes import Training - - -def create_active_training_file(trainer: TrainerLogic, **kwargs) -> None: - update_attributes(trainer._training, **kwargs) # pylint: disable=protected-access - trainer.node.last_training_io.save(training=trainer.training) - - -async def assert_training_state(training: Training, state: str, timeout: float, interval: float) -> None: - try: - await condition(lambda: training.training_state == state, timeout=timeout, interval=interval) - except TimeoutError as exc: - msg = f"Trainer state should be '{state}' after {timeout} seconds, but is {training.training_state}" - raise AssertionError(msg) from exc - except Exception: - logging.exception('##### was ist das hier?') - raise diff --git a/learning_loop_node/tests/trainer/states/test_state_cleanup.py b/learning_loop_node/tests/trainer/states/test_state_cleanup.py index 90b0d283..425812b8 100644 --- a/learning_loop_node/tests/trainer/states/test_state_cleanup.py +++ b/learning_loop_node/tests/trainer/states/test_state_cleanup.py @@ -1,5 +1,4 @@ -from ..state_helper import create_active_training_file -from ..testing_trainer_logic import TestingTrainerLogic +from ....testing import TestingTrainerLogic, create_active_training_file # pylint: disable=protected-access diff --git a/learning_loop_node/tests/trainer/states/test_state_detecting.py b/learning_loop_node/tests/trainer/states/test_state_detecting.py index 6d9c42b3..87f87b72 100644 --- a/learning_loop_node/tests/trainer/states/test_state_detecting.py +++ b/learning_loop_node/tests/trainer/states/test_state_detecting.py @@ -1,10 +1,8 @@ import asyncio from ....enums import TrainerState +from ....testing import TestingTrainerLogic, assert_training_state, create_active_training_file, get_dummy_detections from ....trainer.trainer_logic import TrainerLogic -from ...test_helper import get_dummy_detections -from ..state_helper import assert_training_state, create_active_training_file -from ..testing_trainer_logic import TestingTrainerLogic # pylint: disable=protected-access diff --git a/learning_loop_node/tests/trainer/states/test_state_download_train_model.py b/learning_loop_node/tests/trainer/states/test_state_download_train_model.py index 96bfa6aa..4c59c54b 100644 --- a/learning_loop_node/tests/trainer/states/test_state_download_train_model.py +++ b/learning_loop_node/tests/trainer/states/test_state_download_train_model.py @@ -5,9 +5,7 @@ import pytest from ....enums import TrainerState -from ... import test_helper -from ..state_helper import assert_training_state, create_active_training_file -from ..testing_trainer_logic import TestingTrainerLogic +from ....testing import TestingTrainerLogic, assert_training_state, create_active_training_file, get_latest_model_id # pylint: disable=protected-access @@ -15,7 +13,7 @@ async def test_downloading_is_successful(test_initialized_trainer: TestingTrainerLogic): trainer = test_initialized_trainer - model_id = await test_helper.get_latest_model_id(project='demo') + model_id = await get_latest_model_id(project='demo') create_active_training_file(trainer, base_model_uuid=model_id, training_state=TrainerState.DataDownloaded) diff --git a/learning_loop_node/tests/trainer/states/test_state_prepare.py b/learning_loop_node/tests/trainer/states/test_state_prepare.py index 0c7d3568..42500de8 100644 --- a/learning_loop_node/tests/trainer/states/test_state_prepare.py +++ b/learning_loop_node/tests/trainer/states/test_state_prepare.py @@ -4,9 +4,8 @@ from ....data_classes import Context from ....enums import TrainerState +from ....testing import TestingTrainerLogic, assert_training_state, create_active_training_file from ....trainer.trainer_logic import TrainerLogic -from ..state_helper import assert_training_state, create_active_training_file -from ..testing_trainer_logic import TestingTrainerLogic # pylint: disable=protected-access diff --git a/learning_loop_node/tests/trainer/states/test_state_sync_confusion_matrix.py b/learning_loop_node/tests/trainer/states/test_state_sync_confusion_matrix.py index c002e51e..9337d186 100644 --- a/learning_loop_node/tests/trainer/states/test_state_sync_confusion_matrix.py +++ b/learning_loop_node/tests/trainer/states/test_state_sync_confusion_matrix.py @@ -6,10 +6,9 @@ ) from ....enums import TrainerState +from ....testing import TestingTrainerLogic, assert_training_state, create_active_training_file from ....trainer.trainer_logic import TrainerLogic from ....trainer.trainer_node import TrainerNode -from ..state_helper import assert_training_state, create_active_training_file -from ..testing_trainer_logic import TestingTrainerLogic # pylint: disable=protected-access diff --git a/learning_loop_node/tests/trainer/states/test_state_train.py b/learning_loop_node/tests/trainer/states/test_state_train.py index 89376ee8..ce1f8f95 100644 --- a/learning_loop_node/tests/trainer/states/test_state_train.py +++ b/learning_loop_node/tests/trainer/states/test_state_train.py @@ -1,9 +1,7 @@ from pytest_mock import MockerFixture from ....enums import TrainerState -from ...test_helper import condition -from ..state_helper import assert_training_state, create_active_training_file -from ..testing_trainer_logic import TestingTrainerLogic +from ....testing import TestingTrainerLogic, assert_training_state, condition, create_active_training_file # pylint: disable=protected-access diff --git a/learning_loop_node/tests/trainer/states/test_state_upload_detections.py b/learning_loop_node/tests/trainer/states/test_state_upload_detections.py index 0b0ec14f..8f91109c 100644 --- a/learning_loop_node/tests/trainer/states/test_state_upload_detections.py +++ b/learning_loop_node/tests/trainer/states/test_state_upload_detections.py @@ -6,10 +6,8 @@ from ....data_classes import BoxDetection, Context, Detections from ....enums import TrainerState from ....loop_communication import LoopCommunicator +from ....testing import TestingTrainerLogic, assert_training_state, create_active_training_file, get_dummy_detections from ....trainer.trainer_logic import TrainerLogic -from ...test_helper import get_dummy_detections -from ..state_helper import assert_training_state, create_active_training_file -from ..testing_trainer_logic import TestingTrainerLogic # pylint: disable=protected-access diff --git a/learning_loop_node/tests/trainer/states/test_state_upload_model.py b/learning_loop_node/tests/trainer/states/test_state_upload_model.py index caef4ac9..fd7e8a44 100644 --- a/learning_loop_node/tests/trainer/states/test_state_upload_model.py +++ b/learning_loop_node/tests/trainer/states/test_state_upload_model.py @@ -4,9 +4,8 @@ from ....data_classes import Context from ....enums import TrainerState +from ....testing import TestingTrainerLogic, assert_training_state, create_active_training_file from ....trainer.trainer_logic import TrainerLogic -from ..state_helper import assert_training_state, create_active_training_file -from ..testing_trainer_logic import TestingTrainerLogic # pylint: disable=protected-access diff --git a/learning_loop_node/tests/trainer/test_errors.py b/learning_loop_node/tests/trainer/test_errors.py index 348af1b1..557a4d91 100644 --- a/learning_loop_node/tests/trainer/test_errors.py +++ b/learning_loop_node/tests/trainer/test_errors.py @@ -4,8 +4,7 @@ import pytest from ...enums import TrainerState -from .state_helper import assert_training_state, create_active_training_file -from .testing_trainer_logic import TestingTrainerLogic +from ...testing import TestingTrainerLogic, assert_training_state, create_active_training_file # pylint: disable=protected-access diff --git a/learning_loop_node/tests/trainer/test_trainer_states.py b/learning_loop_node/tests/trainer/test_trainer_states.py index ccd4b498..d0042877 100644 --- a/learning_loop_node/tests/trainer/test_trainer_states.py +++ b/learning_loop_node/tests/trainer/test_trainer_states.py @@ -3,9 +3,9 @@ from ...data_classes import Context, Training from ...enums import TrainerState +from ...testing import TestingTrainerLogic from ...trainer.io_helpers import LastTrainingIO from ...trainer.trainer_node import TrainerNode -from .testing_trainer_logic import TestingTrainerLogic def create_training() -> Training: diff --git a/learning_loop_node/tests/unit/test_testing_package.py b/learning_loop_node/tests/unit/test_testing_package.py new file mode 100644 index 00000000..07d22f1f --- /dev/null +++ b/learning_loop_node/tests/unit/test_testing_package.py @@ -0,0 +1,156 @@ +"""Tests for `learning_loop_node.testing`, the helpers a node repository may import.""" + +import os +import re +import zipfile +from pathlib import Path + +import pytest + +from ... import testing +from ...testing.helpers import PRODUCTION_LOOP_HOST + +REPOSITORY_ROOT = Path(__file__).resolve().parents[3] + + +# ---------------------------------------------------------------- shipped with the wheel + +def _excluded_packages() -> list[str]: + """The `exclude` list of [tool.setuptools.packages.find], without a TOML parser. + + tomllib needs Python 3.11 and the library still targets 3.10. + """ + text = (REPOSITORY_ROOT / 'pyproject.toml').read_text() + section = text.split('[tool.setuptools.packages.find]', 1)[1].split('\n[', 1)[0] + exclude = re.search(r'exclude\s*=\s*\[(.*?)\]', section, re.DOTALL) + assert exclude is not None, 'no exclude list in [tool.setuptools.packages.find]' + return re.findall(r'"([^"]+)"', exclude.group(1)) + + +def test_the_testing_package_is_shipped(): + """The point of the package: a node can import it from the released wheel.""" + assert not any(pattern.startswith('learning_loop_node.testing') for pattern in _excluded_packages()) + + +def test_the_library_own_suite_stays_excluded(): + assert 'learning_loop_node.tests*' in _excluded_packages() + + +def test_everything_exported_exists(): + for name in testing.__all__: + assert hasattr(testing, name), name + + +# ---------------------------------------------------------------- the production loop guard + +def test_an_unset_host_is_refused(monkeypatch: pytest.MonkeyPatch): + """A missing .env must not silently aim a destructive fixture at production.""" + monkeypatch.delenv('LOOP_HOST', raising=False) + monkeypatch.delenv('HOST', raising=False) + with pytest.raises(AssertionError, match='LOOP_HOST is not set'): + testing.assert_not_production_loop() + + +def test_the_production_host_is_refused(monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv('HOST', raising=False) + monkeypatch.setenv('LOOP_HOST', PRODUCTION_LOOP_HOST) + with pytest.raises(AssertionError, match='production Learning Loop'): + testing.assert_not_production_loop() + + +def test_a_test_instance_is_allowed(monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv('HOST', raising=False) + monkeypatch.setenv('LOOP_HOST', 'preview.learning-loop.ai') + testing.assert_not_production_loop() + + +def test_the_bare_host_variable_is_checked_too(monkeypatch: pytest.MonkeyPatch): + """LoopCommunicator reads LOOP_HOST *or* HOST, so the guard has to read both.""" + monkeypatch.delenv('LOOP_HOST', raising=False) + monkeypatch.setenv('HOST', PRODUCTION_LOOP_HOST) + with pytest.raises(AssertionError, match='production Learning Loop'): + testing.assert_not_production_loop() + + +# ---------------------------------------------------------------- dummy detections + +def test_the_dummy_detections_cover_every_type(): + detections = testing.get_dummy_detections() + assert len(detections.box_detections) == 1 + assert len(detections.point_detections) == 1 + assert len(detections.segmentation_detections) == 1 + assert len(detections.classification_detections) == 1 + + +def test_the_dummy_metadata_carries_the_same_detections(): + detections, metadata = testing.get_dummy_detections(), testing.get_dummy_metadata() + assert metadata.box_detections == detections.box_detections + assert metadata.point_detections == detections.point_detections + assert metadata.segmentation_detections == detections.segmentation_detections + assert metadata.classification_detections == detections.classification_detections + + +# ---------------------------------------------------------------- condition + +async def test_condition_returns_as_soon_as_it_holds(): + calls = [] + + def eventually_true() -> bool: + calls.append(None) + return len(calls) >= 2 + + await testing.condition(eventually_true, timeout=1.0, interval=0.01) + assert len(calls) == 2 + + +async def test_condition_times_out(): + with pytest.raises(TimeoutError): + await testing.condition(lambda: False, timeout=0.05, interval=0.01) + + +# ---------------------------------------------------------------- update_attributes + +class _Thing: + def __init__(self) -> None: + self.name = 'before' + + +def test_update_attributes_sets_an_existing_attribute(): + thing = _Thing() + testing.update_attributes(thing, name='after') + assert thing.name == 'after' + + +def test_update_attributes_refuses_an_unknown_attribute(): + with pytest.raises(ValueError, match='does not have a property'): + testing.update_attributes(_Thing(), nope='after') + + +def test_update_attributes_adds_unknown_keys_to_a_dict(): + """A dict has no attributes to check, so it takes whatever it is given.""" + target: dict = {'name': 'before'} + testing.update_attributes(target, name='after', extra=1) + assert target == {'name': 'after', 'extra': 1} + + +# ---------------------------------------------------------------- files + +def test_get_files_in_folder_lists_files_recursively_and_sorted(tmp_path: Path): + (tmp_path / 'b').mkdir() + (tmp_path / 'b' / 'inner.txt').write_text('x') + (tmp_path / 'a.txt').write_text('x') + assert testing.get_files_in_folder(str(tmp_path)) == [ + str(tmp_path / 'a.txt'), str(tmp_path / 'b' / 'inner.txt')] + + +def test_unzip_replaces_the_target_folder(tmp_path: Path): + archive = tmp_path / 'archive.zip' + with zipfile.ZipFile(archive, 'w') as zip_: + zip_.writestr('kept.txt', 'x') + target = tmp_path / 'target' + target.mkdir() + (target / 'stale.txt').write_text('x') + + testing.unzip(str(archive), str(target)) + + assert sorted(os.listdir(target)) == ['kept.txt'] diff --git a/mock_detector/app_code/tests/conftest.py b/mock_detector/app_code/tests/conftest.py index 97178f36..c62568b8 100644 --- a/mock_detector/app_code/tests/conftest.py +++ b/mock_detector/app_code/tests/conftest.py @@ -3,7 +3,6 @@ import logging import multiprocessing import os -import shutil import socket from dataclasses import asdict from multiprocessing import Process @@ -16,6 +15,7 @@ from learning_loop_node.data_classes import Category, ModelInformation from learning_loop_node.detector.detector_node import DetectorNode from learning_loop_node.globals import GLOBALS +from learning_loop_node.testing.fixtures import clear_loggers, data_folder # noqa: F401 from ..mock_detector import MockDetectorFactory @@ -102,11 +102,3 @@ async def port_is(free: bool): return await asyncio.sleep(0.5) raise Exception(f'port {detector_port} is {"not" if free else ""} free') - - -@pytest.fixture(autouse=True, scope='function') -def data_folder(): - GLOBALS.data_folder = '/tmp/learning_loop_lib_data' - shutil.rmtree(GLOBALS.data_folder, ignore_errors=True) - yield - shutil.rmtree(GLOBALS.data_folder, ignore_errors=True) diff --git a/mock_trainer/app_code/tests/conftest.py b/mock_trainer/app_code/tests/conftest.py index e743c749..41f26f0c 100644 --- a/mock_trainer/app_code/tests/conftest.py +++ b/mock_trainer/app_code/tests/conftest.py @@ -1,10 +1,10 @@ import asyncio -import shutil import pytest -from learning_loop_node.globals import GLOBALS from learning_loop_node.loop_communication import LoopCommunicator +from learning_loop_node.testing import assert_not_production_loop +from learning_loop_node.testing.fixtures import clear_loggers, data_folder # noqa: F401 # pylint: disable=redefined-outer-name @@ -18,6 +18,7 @@ async def glc(): @pytest.fixture() async def setup_test_project1(glc: LoopCommunicator): + assert_not_production_loop() await glc.delete("/zauberzeug/projects/pytest_mock_trainer_test1?keep_images=true") await asyncio.sleep(1) project_configuration = { @@ -30,16 +31,9 @@ async def setup_test_project1(glc: LoopCommunicator): await asyncio.sleep(1) -@pytest.fixture(autouse=True, scope='function') -def data_folder(): - GLOBALS.data_folder = '/tmp/learning_loop_lib_data' - shutil.rmtree(GLOBALS.data_folder, ignore_errors=True) - yield - shutil.rmtree(GLOBALS.data_folder, ignore_errors=True) - - @pytest.fixture() async def setup_test_project2(glc: LoopCommunicator): + assert_not_production_loop() await glc.delete("/zauberzeug/projects/pytest_mock_trainer_test2?keep_images=true") await asyncio.sleep(1) project_configuration = { diff --git a/mock_trainer/app_code/tests/test_detections.py b/mock_trainer/app_code/tests/test_detections.py index 113f87f5..48f85934 100644 --- a/mock_trainer/app_code/tests/test_detections.py +++ b/mock_trainer/app_code/tests/test_detections.py @@ -4,9 +4,10 @@ # pylint: disable=protected-access,redefined-outer-name,unused-argument import pytest from fastapi.encoders import jsonable_encoder + from learning_loop_node.data_classes import Category, Context from learning_loop_node.globals import GLOBALS -from learning_loop_node.tests import test_helper +from learning_loop_node.testing import get_latest_model_id from learning_loop_node.trainer.trainer_node import TrainerNode from ..mock_trainer_logic import MockTrainerLogic @@ -17,7 +18,7 @@ async def test_all(): assert_image_count(0) assert GLOBALS.data_folder == '/tmp/learning_loop_lib_data' - latest_model_id = await test_helper.get_latest_model_id(project='pytest_mock_trainer_test1') + latest_model_id = await get_latest_model_id(project='pytest_mock_trainer_test1') trainer = MockTrainerLogic(model_format='mocked') node = TrainerNode(name='test', trainer_logic=trainer) diff --git a/pyproject.toml b/pyproject.toml index d05073f5..4d89005e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -25,6 +25,11 @@ dependencies = [ ] [project.optional-dependencies] +# what a node repository needs to use learning_loop_node.testing +testing = [ + "pytest>=7.0.0,<10.0.0", + "pytest-asyncio>=0.21.1,<0.22.0", +] dev = [ "pytest-flakefinder>=1.1.0,<2.0.0", "pytest-mock>=3.6.1,<4.0.0", diff --git a/run_tests.sh b/run_tests.sh index e8ae7b66..f1bc06dc 100755 --- a/run_tests.sh +++ b/run_tests.sh @@ -2,7 +2,22 @@ # shell script to run all tests # source local .env -set -o allexport; source .env; set +o allexport +if [ -f .env ]; then + set -o allexport; source .env; set +o allexport +fi + +# Every suite except `unit` creates and deletes projects on a real Learning Loop. Without +# LOOP_HOST, LoopCommunicator falls back to learning-loop.ai — production — so stop here +# rather than let the fixtures run there. The fixtures assert the same thing; this just fails +# before pytest starts. +if [ -z "$LOOP_HOST" ] || [ "$LOOP_HOST" = "learning-loop.ai" ]; then + echo "LOOP_HOST is ${LOOP_HOST:-unset}. The live-loop suites create and delete projects," >&2 + echo "so they must not run against production. Set LOOP_HOST in .env, e.g." >&2 + echo " LOOP_HOST=preview.learning-loop.ai" >&2 + echo "To run only the offline suite: python -m pytest learning_loop_node/tests/unit -v" >&2 + exit 1 +fi + # Check if argument is provided if [ $# -eq 1 ]; then # Run tests with filter @@ -26,4 +41,4 @@ python -m pytest learning_loop_node/tests/trainer -v python -m pytest learning_loop_node/tests/general -v python -m pytest mock_detector -v -python -m pytest mock_trainer -v \ No newline at end of file +python -m pytest mock_trainer -v diff --git a/uv.lock b/uv.lock index 1b373e19..d787f9bc 100644 --- a/uv.lock +++ b/uv.lock @@ -817,6 +817,10 @@ dev = [ { name = "pytest-mock" }, { name = "retry" }, ] +testing = [ + { name = "pytest" }, + { name = "pytest-asyncio" }, +] [package.metadata] requires-dist = [ @@ -833,7 +837,9 @@ requires-dist = [ { name = "pillow", specifier = ">=12.3.0,<13.0.0" }, { name = "psutil", specifier = ">=5.9.0,<8.0.0" }, { name = "pynvml", specifier = ">=11.4.1,<13.0.0" }, + { name = "pytest", marker = "extra == 'testing'", specifier = ">=7.0.0,<10.0.0" }, { name = "pytest-asyncio", marker = "extra == 'dev'", specifier = ">=0.21.1,<0.22.0" }, + { name = "pytest-asyncio", marker = "extra == 'testing'", specifier = ">=0.21.1,<0.22.0" }, { name = "pytest-flakefinder", marker = "extra == 'dev'", specifier = ">=1.1.0,<2.0.0" }, { name = "pytest-mock", marker = "extra == 'dev'", specifier = ">=3.6.1,<4.0.0" }, { name = "python-multipart", specifier = ">=0.0.31" }, @@ -843,7 +849,7 @@ requires-dist = [ { name = "tqdm", specifier = ">=4.66.3,<5.0.0" }, { name = "uvicorn", extras = ["standard"], specifier = ">=0.22.0,<1.0.0" }, ] -provides-extras = ["dev"] +provides-extras = ["testing", "dev"] [[package]] name = "multidict"