233 lines
9.5 KiB
Python
233 lines
9.5 KiB
Python
"""
|
||
Контракт 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"]
|