359 lines
17 KiB
Python
359 lines
17 KiB
Python
"""
|
||
Тесты контракта API для веб-интерфейса.
|
||
|
||
Проверяют, что ответы эндпоинтов содержат поля, которые читает
|
||
`src/api/static/js/dxa-app.js`. Ранее эти ключи разошлись: панель деталей
|
||
показывала прочерки и одинаковые значения при любом клике, потому что API
|
||
отдавал `metrics` другой структуры. Тесты фиксируют контракт.
|
||
"""
|
||
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_assets = pytest.mark.skipif(
|
||
not (DATASET.is_dir() and CHECKPOINT.exists()),
|
||
reason="test DICOM files or trained checkpoint are not available",
|
||
)
|
||
|
||
# Поля, которые панель деталей читает напрямую из ответа.
|
||
DETAIL_TOP_LEVEL = [
|
||
"anatomical_region", "quality_class", "quality_label", "violation_type",
|
||
"violation_type_label", "violation_type_is_heuristic", "violation_type_note",
|
||
"reason", "confidence", "confidence_per_class", "threshold_probability",
|
||
"region_confidence", "overall_quality", "severity", "metrics",
|
||
"view_quality", "reasons", "metrics_note",
|
||
]
|
||
# Ключи внутри metrics, которые панель использует.
|
||
METRICS_KEYS = ["sharpness_laplacian", "sharpness_fft", "roi", "calibration"]
|
||
ROI_KEYS = ["size", "bounding_box", "valid", "reason", "issues"]
|
||
|
||
|
||
@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)
|
||
|
||
|
||
def _test_files():
|
||
return sorted(DATASET.glob("*.dcm"))
|
||
|
||
|
||
@needs_assets
|
||
class TestDetailedContract:
|
||
@pytest.fixture(scope="class")
|
||
def payloads(self, client):
|
||
out = []
|
||
for path in _test_files():
|
||
response = client.post(
|
||
"/api/v1/analyze/detailed",
|
||
files={"file": (path.name, path.read_bytes(), "application/dicom")},
|
||
params={"include_visualization": "true"},
|
||
)
|
||
assert response.status_code == 200, f"{path.name}: {response.text[:200]}"
|
||
out.append((path.name, response.json()))
|
||
return out
|
||
|
||
def test_all_documented_fields_present(self, payloads):
|
||
for name, data in payloads:
|
||
missing = [k for k in DETAIL_TOP_LEVEL if k not in data]
|
||
assert not missing, f"{name} missing {missing}"
|
||
|
||
def test_confidence_per_class_is_coherent(self, payloads):
|
||
for name, data in payloads:
|
||
cpc = data["confidence_per_class"]
|
||
assert set(cpc) == {"correct", "violation"}, name
|
||
assert cpc["correct"] + cpc["violation"] == pytest.approx(1.0, abs=1e-3), name
|
||
# Вероятность класса не должна расходиться с полем confidence.
|
||
assert data["confidence"] == pytest.approx(cpc["violation"], abs=1e-3), name
|
||
|
||
def test_metrics_have_expected_structure(self, payloads):
|
||
for name, data in payloads:
|
||
metrics = data["metrics"]
|
||
assert isinstance(metrics, dict) and metrics, f"{name}: metrics is empty"
|
||
missing = [k for k in METRICS_KEYS if k not in metrics]
|
||
assert not missing, f"{name} metrics missing {missing}"
|
||
for key in ROI_KEYS:
|
||
assert key in metrics["roi"], f"{name} roi missing {key}"
|
||
|
||
def test_metrics_differ_between_distinct_images(self, payloads):
|
||
"""
|
||
Метрики должны различаться для разных снимков.
|
||
|
||
Одинаковые значения были главным симптомом бага: пороги эвристик были
|
||
насыщены, и панель выглядела «не обновляющейся» при кликах.
|
||
"""
|
||
sharp = [d["metrics"]["sharpness_laplacian"] for _, d in payloads]
|
||
assert len(set(sharp)) > 1, "sharpness does not vary across images"
|
||
|
||
def test_duplicate_images_share_metrics(self, payloads):
|
||
"""Побайтные дубликаты должны давать идентичные метрики."""
|
||
by_study = {}
|
||
for name, data in payloads:
|
||
key = (data["study_uid"], data["image_uid"])
|
||
by_study.setdefault(key, []).append(data["metrics"]["sharpness_laplacian"])
|
||
for key, values in by_study.items():
|
||
if len(values) > 1:
|
||
assert max(values) == pytest.approx(min(values)), f"duplicates differ for {key}"
|
||
|
||
def test_region_specific_sections_match_region(self, payloads):
|
||
for name, data in payloads:
|
||
region = data["anatomical_region"]
|
||
if region == "spine":
|
||
assert "spine_completeness" in data, name
|
||
assert data["spine_completeness"], f"{name}: spine section empty"
|
||
if region.startswith("hip"):
|
||
assert data.get("hip_completeness") or data.get("hip_rotation"), name
|
||
|
||
def test_no_uncalibrated_verdicts_are_exposed(self, payloads):
|
||
"""
|
||
Эвристики не должны выглядеть как заключение.
|
||
|
||
Проверка фиксирует, что наружу не отдаются некалиброванные вердикты:
|
||
например число «позвонков» (эвристика выдавала 46) или тексты вида
|
||
«Позвонок 1 обрезан».
|
||
"""
|
||
for name, data in payloads:
|
||
spine = data.get("spine_completeness") or {}
|
||
assert "num_vertebrae" not in spine, f"{name}: raw vertebrae count exposed"
|
||
assert "issues" not in spine, f"{name}: raw issue texts exposed"
|
||
assert "reason" not in spine, f"{name}: raw verdict text exposed"
|
||
blob = str(data)
|
||
assert "Позвонок 1 обрезан" not in blob, f"{name}: verdict leaked into payload"
|
||
|
||
def test_quality_label_is_not_duplicated_in_badge(self, payloads):
|
||
"""Подпись класса не должна совпадать с текстом бейджа (было «OK OK»)."""
|
||
for name, data in payloads:
|
||
if data["quality_class"] == 0:
|
||
assert data["quality_label"] != "OK", name
|
||
|
||
def test_visualizations_are_valid_base64(self, payloads):
|
||
import base64
|
||
|
||
for name, data in payloads:
|
||
for field in ("image", "mask"):
|
||
raw = data.get(field)
|
||
assert raw, f"{name}: {field} is empty"
|
||
decoded = base64.b64decode(raw)
|
||
assert decoded[:8] == b"\x89PNG\r\n\x1a\n", f"{name}: {field} is not PNG"
|
||
|
||
def test_reasons_explain_the_decision(self, payloads):
|
||
for name, data in payloads:
|
||
reasons = data["reasons"]
|
||
assert isinstance(reasons, list) and reasons, name
|
||
text = " ".join(reasons)
|
||
assert data["anatomical_region"] in text or "область" in text, name
|
||
|
||
def test_metrics_note_warns_about_calibration(self, payloads):
|
||
for name, data in payloads:
|
||
assert "не калиброваны" in data["metrics_note"], name
|
||
|
||
def test_violation_type_comes_from_canonical_dictionary(self, payloads):
|
||
"""
|
||
Тип нарушения — код из единого словаря, а подпись — из него же.
|
||
|
||
Раньше интерфейс держал собственную таблицу подписей, и коды экспертной
|
||
таблицы (`artifact`, `roi_incorrect`, `positioning`) показывались как
|
||
есть. Тест фиксирует, что сервер отдаёт канонический код и подпись.
|
||
"""
|
||
from src.dxa.violations import VIOLATION_LABELS, VIOLATION_TYPES
|
||
|
||
for name, data in payloads:
|
||
code = data["violation_type"]
|
||
assert code == "" or code in VIOLATION_TYPES, f"{name}: unknown code {code}"
|
||
if code:
|
||
assert data["violation_type_label"] == VIOLATION_LABELS[code], name
|
||
assert data["violation_type_is_heuristic"] is True, name
|
||
assert "эвристик" in data["violation_type_note"], name
|
||
else:
|
||
assert data["violation_type_label"] == "", name
|
||
assert data["violation_type_is_heuristic"] is False, name
|
||
assert data["violation_type_note"] == "", name
|
||
|
||
def test_violation_type_matches_quality_class(self, payloads):
|
||
for name, data in payloads:
|
||
if data["quality_class"] == 0:
|
||
assert data["violation_type"] == "", name
|
||
else:
|
||
assert data["violation_type"] != "", name
|
||
|
||
|
||
@needs_assets
|
||
class TestHealthAndModelCard:
|
||
"""`/health` и `/api/v1/model` — источник сведений о модели для интерфейса."""
|
||
|
||
def test_health_reports_model_provenance(self, client):
|
||
data = client.get("/api/v1/health").json()
|
||
assert data["status"] == "ok" and data["model_loaded"] is True
|
||
model = data["model"]
|
||
assert model["loaded"] is True and model["exists"] is True
|
||
for key in ("path", "backbone", "head", "epoch", "threshold_logit",
|
||
"threshold_probability", "labels_csv", "val_metrics"):
|
||
assert key in model, f"health.model missing {key}"
|
||
# Разметка из экспертной таблицы — то, на чём обучен рабочий чекпоинт.
|
||
assert model["labels_csv"] == "labels/labels_images.csv"
|
||
assert 0.0 < model["threshold_probability"] < 1.0
|
||
assert 0.0 <= model["val_metrics"]["roc_auc"] <= 1.0
|
||
|
||
def test_model_card_covers_canonical_dictionary(self, client):
|
||
from src.dxa.violations import VIOLATION_TYPES
|
||
|
||
card = client.get("/api/v1/model").json()
|
||
codes = [v["code"] for v in card["violations"]]
|
||
assert codes == list(VIOLATION_TYPES)
|
||
for entry in card["violations"]:
|
||
assert entry["label"] and entry["note"]
|
||
assert entry["scope"] in ("spine", "hip", "any")
|
||
assert entry["source"] in ("expert_table", "condition_doctor", "system")
|
||
|
||
def test_model_card_metrics_are_intervals(self, client):
|
||
card = client.get("/api/v1/model").json()
|
||
cmp = card["comparison"]
|
||
for metric, variants in cmp["metrics"].items():
|
||
for name, triple in variants.items():
|
||
assert len(triple) == 3, f"{metric}/{name} is not [value, lo, hi]"
|
||
assert triple[1] <= triple[0] <= triple[2], f"{metric}/{name} CI inverted"
|
||
for metric, triple in cmp["paired_delta_table_minus_union"].items():
|
||
assert triple[1] <= triple[0] <= triple[2], f"delta {metric} CI inverted"
|
||
|
||
def test_model_card_reports_label_rule(self, client):
|
||
card = client.get("/api/v1/model").json()
|
||
assert card["model"]["label_rule"] == "table"
|
||
assert "экспертной таблицы" in card["labels"]["rule"]
|
||
# Правило выбиралось измерением: в карточке должно быть и решение, и основание.
|
||
assert card["comparison"]["decision"]
|
||
assert card["comparison"]["paired_wins"]
|
||
|
||
def test_model_card_explains_deployed_metrics(self, client):
|
||
card = client.get("/api/v1/model").json()
|
||
caveat = card["model"]["caveat"]
|
||
# Чекпоинт — обычный прогон, а не лучший из выборки; честная оценка рядом.
|
||
assert "seed по умолчанию" in caveat
|
||
assert card["model"]["val_roc_auc"] == pytest.approx(
|
||
card["comparison"]["metrics"]["roc_auc"]["table"][0], abs=0.02
|
||
)
|
||
|
||
def test_model_card_loads_model_on_cold_start(self, client):
|
||
"""
|
||
Первый же запрос к сервису должен отдавать полную карточку.
|
||
|
||
Блок `loaded` берётся из метаданных чекпоинта, поэтому эндпоинт обязан
|
||
сам инициировать загрузку модели. Здесь состояние сбрасывается, чтобы
|
||
воспроизвести холодный старт независимо от порядка тестов.
|
||
"""
|
||
import src.main as main
|
||
|
||
saved = (main.dxa_model, main.dxa_preprocess, main.dxa_threshold, main.dxa_metadata)
|
||
try:
|
||
main.dxa_model = None
|
||
main.dxa_metadata = {}
|
||
card = client.get("/api/v1/model").json()
|
||
assert card["loaded"]["loaded"] is True
|
||
assert card["loaded"]["val_metrics"], "метрики чекпоинта пусты после холодного старта"
|
||
assert card["loaded"]["val_metrics"]["roc_auc"] > 0.5
|
||
finally:
|
||
main.dxa_model, main.dxa_preprocess, main.dxa_threshold, main.dxa_metadata = saved
|
||
|
||
def test_frontend_reads_only_existing_card_fields(self, client):
|
||
"""
|
||
Фронтенд не должен обращаться к полям карточки, которых в ней нет.
|
||
|
||
Такой рассинхрон уже случался: ключи в карточке переименовали, а в
|
||
`dxa-app.js` остались старые — на панели «О модели» появлялись плитки с
|
||
прочерком вместо чисел.
|
||
"""
|
||
import re
|
||
from pathlib import Path
|
||
|
||
card = client.get("/api/v1/model").json()
|
||
source = Path("src/api/static/js/dxa-app.js").read_text(encoding="utf-8")
|
||
used = set(re.findall(r"\bds\.([a-z_]+)", source))
|
||
used |= set(re.findall(r"\bdeployed\.([a-z_]+)", source))
|
||
missing_dataset = sorted(k for k in used if k not in card["dataset"] and k not in card["model"])
|
||
assert not missing_dataset, f"фронтенд читает несуществующие поля: {missing_dataset}"
|
||
|
||
# Плитки блока «Данные» не должны получать прочерк из-за смены ключей.
|
||
tiles = re.findall(r"card\('([^']+)',\s*ds\.([a-z_]+)", source)
|
||
assert tiles, "не найдено ни одной плитки блока «Данные»"
|
||
for _title, key in tiles:
|
||
assert key in card["dataset"], f"плитка «{_title}» читает отсутствующий ключ ds.{key}"
|
||
|
||
def test_model_card_matches_loaded_checkpoint(self, client):
|
||
"""Карточка и фактически загруженный чекпоинт не должны расходиться."""
|
||
card = client.get("/api/v1/model").json()
|
||
loaded = card["loaded"]
|
||
assert loaded["backbone"] == card["model"]["backbone"]
|
||
assert loaded["head"] == card["model"]["head"]
|
||
assert loaded["labels_csv"] == card["model"]["labels_csv"]
|
||
assert loaded["threshold_logit"] == pytest.approx(
|
||
card["model"]["threshold_logit"], abs=1e-4
|
||
)
|
||
|
||
def test_model_card_lists_limitations_and_sources(self, client):
|
||
card = client.get("/api/v1/model").json()
|
||
assert len(card["limitations"]) >= 5
|
||
assert all(isinstance(text, str) and text for text in card["limitations"])
|
||
assert card["sources"]
|
||
assert card["snapshot_date"]
|
||
|
||
|
||
@needs_assets
|
||
class TestBasicContract:
|
||
def test_analyze_returns_basic_fields(self, client):
|
||
path = _test_files()[0]
|
||
response = client.post(
|
||
"/api/v1/analyze",
|
||
files={"file": (path.name, path.read_bytes(), "application/dicom")},
|
||
)
|
||
assert response.status_code == 200
|
||
data = response.json()
|
||
for key in ("anatomical_region", "quality_class", "confidence", "processing_status"):
|
||
assert key in data
|
||
# Базовый эндпоинт не должен тянуть тяжёлые метрики
|
||
assert data["metrics"] == {}
|
||
|
||
|
||
@needs_assets
|
||
class TestExportContract:
|
||
def test_export_has_required_columns(self, client):
|
||
import io
|
||
|
||
import pandas as pd
|
||
|
||
files = [(f.name, f.read_bytes(), "application/dicom") for f in _test_files()]
|
||
response = client.post("/api/v1/export", files=[("files", f) for f in files])
|
||
assert response.status_code == 200
|
||
df = pd.read_excel(io.BytesIO(response.content))
|
||
required = [
|
||
"path_to_study", "study_uid", "image_uid", "anatomical_region",
|
||
"quality_class", "violation_type", "processing_status", "time_of_processing",
|
||
]
|
||
# Столбцы задания идут первыми и в заданном порядке.
|
||
assert list(df.columns)[: len(required)] == required
|
||
assert len(df) == len(_test_files())
|
||
assert (df["processing_status"] == "Success").all()
|
||
|
||
|
||
@needs_assets
|
||
class TestErrorHandling:
|
||
def test_garbage_upload_returns_500_not_crash(self, client):
|
||
response = client.post(
|
||
"/api/v1/analyze",
|
||
files={"file": ("junk.dcm", b"not a dicom", "application/dicom")},
|
||
)
|
||
assert response.status_code == 500
|
||
assert "error" in response.json()
|