train_utility/tests/test_export_runner.py
2026-08-04 14:00:46 +04:00

244 lines
6.6 KiB
Python

from __future__ import annotations
import json
import os
import sys
from pathlib import Path
from types import ModuleType
from typing import Any
import pytest
from yolo_webui import export_runner
def _write_config(tmp_path: Path, data: Any) -> Path:
path = tmp_path / "export.json"
path.write_text(json.dumps(data), encoding="utf-8")
return path
def _write_model(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> Path:
monkeypatch.chdir(tmp_path)
models = tmp_path / "models"
models.mkdir()
model = models / "model.pt"
model.write_bytes(b"checkpoint")
return model
def _install_fake_ultralytics(
monkeypatch: pytest.MonkeyPatch,
result: Any,
) -> tuple[list[str], list[dict[str, Any]]]:
constructed: list[str] = []
export_calls: list[dict[str, Any]] = []
class FakeYOLO:
def __init__(self, model: str) -> None:
assert os.environ["ULTRALYTICS_SAFE_LOAD"] == "1"
constructed.append(model)
def export(self, **kwargs: Any) -> Any:
export_calls.append(kwargs)
return result
fake_ultralytics = ModuleType("ultralytics")
fake_ultralytics.YOLO = FakeYOLO # type: ignore[attr-defined]
monkeypatch.setitem(sys.modules, "ultralytics", fake_ultralytics)
return constructed, export_calls
def test_main_validates_config_and_reports_existing_export(
monkeypatch: pytest.MonkeyPatch,
tmp_path: Path,
capsys: pytest.CaptureFixture[str],
) -> None:
model = _write_model(tmp_path, monkeypatch)
result = tmp_path / "models" / "model.onnx"
result.write_bytes(b"onnx")
constructed, export_calls = _install_fake_ultralytics(monkeypatch, result)
monkeypatch.setenv("ULTRALYTICS_SAFE_LOAD", "0")
config = _write_config(
tmp_path,
{
"model": model.name,
"format": "ONNX",
"imgsz": 320,
"half": False,
"int8": False,
"dynamic": True,
"simplify": True,
"batch": 2,
"workspace": 2.5,
},
)
return_code = export_runner.main([str(config)])
output = capsys.readouterr()
assert return_code == 0
assert constructed == [str(model)]
assert export_calls == [
{
"format": "onnx",
"imgsz": 320,
"half": False,
"int8": False,
"dynamic": True,
"simplify": True,
"batch": 2,
"workspace": 2.5,
}
]
assert f"__YOLO_WEBUI_RESULT__:{result}" in output.out
@pytest.mark.parametrize(
"override",
[
{"format": "onnxx"},
{"format": ["onnx"]},
{"imgsz": "640"},
{"imgsz": True},
{"half": 1},
{"batch": 0},
{"workspace": float("inf")},
{"workspace": 10**400},
{"workspace": 64.5},
{"unexpected": "value"},
],
)
def test_invalid_export_values_are_rejected_before_loading_ultralytics(
override: dict[str, Any],
monkeypatch: pytest.MonkeyPatch,
tmp_path: Path,
capsys: pytest.CaptureFixture[str],
) -> None:
model = _write_model(tmp_path, monkeypatch)
constructed, _ = _install_fake_ultralytics(monkeypatch, model)
config = _write_config(
tmp_path,
{
"model": str(model),
"format": "onnx",
**override,
},
)
return_code = export_runner.main([str(config)])
output = capsys.readouterr()
assert return_code == 1
assert constructed == []
assert "Ошибка конфигурации экспорта:" in output.err
def test_non_object_json_is_reported_as_invalid_config(
monkeypatch: pytest.MonkeyPatch,
tmp_path: Path,
capsys: pytest.CaptureFixture[str],
) -> None:
monkeypatch.chdir(tmp_path)
return_code = export_runner.main([str(_write_config(tmp_path, ["model.pt"]))])
output = capsys.readouterr()
assert return_code == 1
assert "JSON-объектом" in output.err
def test_model_must_be_a_regular_file_inside_an_allowed_root(
monkeypatch: pytest.MonkeyPatch,
tmp_path: Path,
capsys: pytest.CaptureFixture[str],
) -> None:
monkeypatch.chdir(tmp_path)
(tmp_path / "models").mkdir()
outside = tmp_path / "outside.pt"
outside.write_bytes(b"checkpoint")
return_code = export_runner.main(
[str(_write_config(tmp_path, {"model": str(outside), "format": "onnx"}))]
)
output = capsys.readouterr()
assert return_code == 1
assert "разрешённом каталоге" in output.err
def test_symlink_cannot_escape_an_allowed_model_root(
monkeypatch: pytest.MonkeyPatch,
tmp_path: Path,
capsys: pytest.CaptureFixture[str],
) -> None:
monkeypatch.chdir(tmp_path)
models = tmp_path / "models"
models.mkdir()
outside = tmp_path / "outside.pt"
outside.write_bytes(b"checkpoint")
link = models / "linked.pt"
link.symlink_to(outside)
return_code = export_runner.main(
[str(_write_config(tmp_path, {"model": str(link), "format": "onnx"}))]
)
output = capsys.readouterr()
assert return_code == 1
assert "разрешённом каталоге" in output.err
def test_half_and_int8_cannot_be_enabled_together(
monkeypatch: pytest.MonkeyPatch,
tmp_path: Path,
capsys: pytest.CaptureFixture[str],
) -> None:
model = _write_model(tmp_path, monkeypatch)
return_code = export_runner.main(
[
str(
_write_config(
tmp_path,
{
"model": str(model),
"format": "onnx",
"half": True,
"int8": True,
},
)
)
]
)
output = capsys.readouterr()
assert return_code == 1
assert "нельзя включать одновременно" in output.err
@pytest.mark.parametrize("result", [None, [], "/missing/export.onnx"])
def test_missing_export_artifact_is_a_failure(
result: Any,
monkeypatch: pytest.MonkeyPatch,
tmp_path: Path,
capsys: pytest.CaptureFixture[str],
) -> None:
model = _write_model(tmp_path, monkeypatch)
_install_fake_ultralytics(monkeypatch, result)
return_code = export_runner.main(
[
str(
_write_config(
tmp_path,
{"model": str(model), "format": "onnx"},
)
)
]
)
output = capsys.readouterr()
assert return_code == 1
assert "__YOLO_WEBUI_RESULT__:" not in output.out
assert "Ошибка при экспорте модели:" in output.out