bone_2026/tests/test_labeling_api.py

233 lines
9.5 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
Контракт API ручной разметки (`/api/v1/labeling/*`).
Проверяется не только форма ответов, но и два свойства, которые легко потерять:
вердикт нельзя сохранить для снимка вне датасета (интерфейс ходит по путям,
приходящим от клиента), а выгрузка должна читаться тем же `load_labels_csv`, что
и построенная разметка, иначе её нельзя передать в обучение как `--labels-csv`.
"""
import csv
import os
import sys
from pathlib import Path
import pytest
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
DATASET = Path("dataset_hack")
CHECKPOINT = Path("models/dxa_model.pth")
pytest.importorskip("fastapi")
pytest.importorskip("httpx")
needs_dataset = pytest.mark.skipif(
not DATASET.is_dir(),
reason="датасет dataset_hack недоступен",
)
VALID_VIOLATION = {
"spine": "artifact",
"hip_right": "rotation",
"hip_left": "rotation",
}
@pytest.fixture(scope="module")
def client():
os.environ.setdefault("DXA_MODEL_PATH", str(CHECKPOINT))
from fastapi.testclient import TestClient
from src.main import app
return TestClient(app)
@pytest.fixture
def manual_store(tmp_path, monkeypatch):
"""Вердикты пишутся в tmp, а не в боевой `labels/manual_labels.csv`."""
path = tmp_path / "manual_labels.csv"
monkeypatch.setattr("src.main.MANUAL_LABELS_PATH", str(path))
return path
def _first_item(client) -> dict:
response = client.get("/api/v1/labeling/items", params={"limit": 1})
assert response.status_code == 200
items = response.json()["items"]
assert items, "датасет должен содержать хотя бы один снимок"
return items[0]
@needs_dataset
class TestItems:
def test_reports_progress_and_dictionary(self, client, manual_store):
payload = client.get("/api/v1/labeling/items", params={"limit": 1}).json()
assert payload["total"] == payload["progress"]["total"]
assert payload["progress"]["reviewed"] == 0
assert set(payload["regions"]) == {"spine", "hip_right", "hip_left"}
codes = {v["code"] for v in payload["violations"]}
assert {"positioning", "axis_deviation", "artifact", "rotation", "roi_incorrect"} <= codes
def test_item_carries_provenance(self, client, manual_store):
item = _first_item(client)
for field in ("path", "file", "study", "region", "label", "base_label",
"label_source", "filename_conflict", "reviewed", "image_url"):
assert field in item, f"нет поля {field}"
assert item["label_source"] in ("manual", "table", "filename")
assert item["label"] in (0, 1)
def test_conflicts_come_first(self, client, manual_store):
items = client.get("/api/v1/labeling/items", params={"limit": 200}).json()["items"]
flags = [i["filename_conflict"] for i in items]
assert any(flags), "в датасете есть расхождения разметки с именем файла"
first_false = flags.index(False)
assert not any(flags[first_false:]), "расхождения должны идти первыми"
def test_region_filter(self, client, manual_store):
payload = client.get(
"/api/v1/labeling/items", params={"limit": 200, "region": "spine"}
).json()
assert payload["items"]
assert {i["region"] for i in payload["items"]} == {"spine"}
assert payload["total"] == payload["progress"]["by_region"]["spine"]["total"]
@needs_dataset
class TestImage:
def test_returns_png_for_dataset_image(self, client, manual_store):
item = _first_item(client)
response = client.get("/api/v1/labeling/image", params={"path": item["path"]})
assert response.status_code == 200
assert response.headers["content-type"] == "image/png"
assert response.content.startswith(b"\x89PNG")
@pytest.mark.parametrize(
"bad_path",
["../../etc/passwd", "/etc/passwd", "README.md", "dataset_hack/нет_такого.dcm"],
)
def test_refuses_paths_outside_dataset(self, client, manual_store, bad_path):
response = client.get("/api/v1/labeling/image", params={"path": bad_path})
assert response.status_code == 403, f"{bad_path} не должен отдаваться"
@needs_dataset
class TestVerdict:
def test_saves_and_shows_in_items(self, client, manual_store):
item = _first_item(client)
region = item["region"] or "spine"
response = client.post(
"/api/v1/labeling/verdict",
json={
"path": item["path"],
"region": region,
"quality": 1,
"violation": VALID_VIOLATION[region],
"comment": "размытие контура",
"reviewer": "специалист",
},
)
assert response.status_code == 200
assert response.json()["progress"]["reviewed"] == 1
stored = list(csv.DictReader(manual_store.open(newline="", encoding="utf-8")))
assert stored[0]["quality_class"] == "1"
assert stored[0]["reviewer"] == "специалист"
after = client.get("/api/v1/labeling/items", params={"limit": 200}).json()
reviewed = [i for i in after["items"] if i["path"] == item["path"]]
assert reviewed and reviewed[0]["reviewed"] is True
assert reviewed[0]["label_source"] == "manual"
assert reviewed[0]["label"] == 1
def test_good_image_rejects_violation_type(self, client, manual_store):
item = _first_item(client)
response = client.post(
"/api/v1/labeling/verdict",
json={"path": item["path"], "region": item["region"] or "spine",
"quality": 0, "violation": "artifact"},
)
assert response.status_code == 400
assert "не может быть типа нарушения" in response.json()["error"]
def test_unknown_violation_is_rejected(self, client, manual_store):
item = _first_item(client)
response = client.post(
"/api/v1/labeling/verdict",
json={"path": item["path"], "region": item["region"] or "spine",
"quality": 1, "violation": "ukladka_plohaya"},
)
assert response.status_code == 400
def test_outside_dataset_path_is_rejected(self, client, manual_store):
response = client.post(
"/api/v1/labeling/verdict",
json={"path": "/tmp/чужой.dcm", "region": "spine", "quality": 1, "violation": "artifact"},
)
assert response.status_code == 400
def test_delete_returns_image_to_built_labels(self, client, manual_store):
item = _first_item(client)
region = item["region"] or "spine"
client.post(
"/api/v1/labeling/verdict",
json={"path": item["path"], "region": region, "quality": 1,
"violation": VALID_VIOLATION[region]},
)
response = client.delete("/api/v1/labeling/verdict", params={"path": item["path"]})
assert response.status_code == 200
assert response.json()["progress"]["reviewed"] == 0
second = client.delete("/api/v1/labeling/verdict", params={"path": item["path"]})
assert second.status_code == 404
@needs_dataset
class TestExport:
def test_reviewed_export_is_loadable_by_training(self, client, manual_store):
from src.dxa.excel_labels import load_labels_csv
item = _first_item(client)
region = item["region"] or "spine"
client.post(
"/api/v1/labeling/verdict",
json={"path": item["path"], "region": region, "quality": 1,
"violation": VALID_VIOLATION[region], "reviewer": "специалист"},
)
response = client.get("/api/v1/labeling/export", params={"scope": "reviewed"})
assert response.status_code == 200
assert "text/csv" in response.headers["content-type"]
path = manual_store.parent / "export.csv"
path.write_bytes(response.content)
table = load_labels_csv(path)
assert table[item["path"]]["quality_class"] == "1"
assert table[item["path"]]["anatomical_region"] == region
def test_all_export_covers_whole_dataset(self, client, manual_store):
payload = client.get("/api/v1/labeling/items", params={"limit": 1}).json()
response = client.get("/api/v1/labeling/export", params={"scope": "all"})
rows = list(csv.DictReader(response.content.decode("utf-8-sig").splitlines()))
assert len(rows) == payload["progress"]["total"]
assert all(row["label_rule"] in ("table", "filename", "manual") for row in rows)
def test_bad_scope_rejected(self, client, manual_store):
response = client.get("/api/v1/labeling/export", params={"scope": "everything"})
assert response.status_code == 400
def test_xlsx_export(self, client, manual_store):
response = client.get(
"/api/v1/labeling/export", params={"scope": "reviewed", "fmt": "xlsx"}
)
assert response.status_code == 200
assert "spreadsheetml" in response.headers["content-type"]
assert response.content.startswith(b"PK")
def test_label_page_is_served(client):
response = client.get("/label")
assert response.status_code == 200
assert "text/html" in response.headers["content-type"]