from __future__ import annotations import sys import signal from pathlib import Path from types import ModuleType, SimpleNamespace from typing import Any from yolo_webui.config import MlflowConfig, TrainingConfig from yolo_webui.trainer import TrainingEvent, TrainingRunner class FakeProcess: def __init__(self) -> None: self.signals: list[int] = [] self.terminate_calls = 0 self.kill_calls = 0 def send_signal(self, signum: int) -> None: self.signals.append(signum) def terminate(self) -> None: self.terminate_calls += 1 def kill(self) -> None: self.kill_calls += 1 def poll(self) -> None: return None def test_runner_wires_yolo_callbacks_and_returns_output( monkeypatch: Any, tmp_path: Path ) -> None: settings_updates: list[dict[str, bool]] = [] constructed: list[tuple[str, str]] = [] train_arguments: list[dict[str, Any]] = [] class FakeSettings: def update(self, values: dict[str, bool]) -> None: settings_updates.append(values) class FakeYOLO: def __init__(self, model: str, task: str) -> None: constructed.append((model, task)) self.callbacks: dict[str, list[Any]] = {} self.trainer = SimpleNamespace( args=SimpleNamespace(epochs=2), epoch=0, metrics={}, fitness=0.0, best_fitness=0.0, stop=False, save_dir=tmp_path / "run", ) def add_callback(self, name: str, callback: Any) -> None: self.callbacks.setdefault(name, []).append(callback) def run_callbacks(self, name: str) -> None: for callback in self.callbacks.get(name, []): callback(self.trainer) def train(self, **kwargs: Any) -> None: train_arguments.append(kwargs) self.run_callbacks("on_train_start") for epoch in range(2): self.trainer.epoch = epoch self.trainer.metrics = {"metrics/mAP50": 0.5 + epoch / 10} self.trainer.fitness = 0.5 + epoch / 10 self.trainer.best_fitness = self.trainer.fitness self.run_callbacks("on_fit_epoch_end") self.run_callbacks("on_train_end") fake_ultralytics = ModuleType("ultralytics") fake_ultralytics.YOLO = FakeYOLO # type: ignore[attr-defined] fake_ultralytics.settings = FakeSettings() # type: ignore[attr-defined] fake_trainers = ModuleType("yolo_webui.ultralytics_trainers") class FakeTrainer: pass fake_trainers.trainer_for_task = lambda _task: FakeTrainer # type: ignore[attr-defined] monkeypatch.setitem(sys.modules, "ultralytics", fake_ultralytics) monkeypatch.setitem(sys.modules, "yolo_webui.ultralytics_trainers", fake_trainers) config = TrainingConfig( dataset="dataset.yaml", model="model.pt", task="pose", epochs=2, mlflow=MlflowConfig(enabled=False), ) events: list[TrainingEvent] = [] output = TrainingRunner().train(config, events.append) assert output == tmp_path / "run" assert constructed == [("models/model.pt", "pose")] assert settings_updates == [{"mlflow": False}] assert train_arguments[0]["data"] == "dataset.yaml" assert train_arguments[0]["verbose"] is True assert train_arguments[0]["trainer"] is FakeTrainer assert [event.kind for event in events] == ["info", "started", "epoch", "epoch", "success"] assert "mAP50=0.5" in events[2].message assert "mAP50=0.6" in events[3].message def test_final_checkpoint_validation_is_not_reported_as_an_extra_epoch() -> None: events: list[TrainingEvent] = [] trainer = SimpleNamespace( args=SimpleNamespace(epochs=1), epoch=1, metrics={"metrics/mAP50": 0.75}, validator=SimpleNamespace(training=False), ) TrainingRunner()._on_epoch_end(events.append)(trainer) assert events == [] def test_prepare_run_clears_previous_stop_request() -> None: runner = TrainingRunner() runner.request_stop() assert runner.stop_requested is True runner.prepare_run() assert runner.stop_requested is False def test_early_stop_is_delivered_after_subprocess_ready() -> None: runner = TrainingRunner() process = FakeProcess() runner.request_stop() runner.set_subprocess(process, ready=False) assert process.signals == [] runner.mark_subprocess_ready() assert process.signals == [signal.SIGTERM] assert process.terminate_calls == 0 runner.clear_subprocess() def test_ready_subprocess_receives_cooperative_signal_not_terminate() -> None: runner = TrainingRunner() process = FakeProcess() runner.set_subprocess(process, ready=True) runner.request_stop() assert process.signals == [signal.SIGTERM] assert process.terminate_calls == 0 runner.clear_subprocess()