Spaces:
Sleeping
Sleeping
File size: 6,184 Bytes
8ea8b52 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 | import pytest
from fastapi.testclient import TestClient
import numpy as np
import cv2
import tempfile
import os
from backend.oct_analyzer.api import app
import backend.oct_analyzer.api as api_module
from backend.oct_analyzer.classifier_integration import ClassifierWrapper, get_classifier
import backend.oct_analyzer.classifier_integration as ci
import backend.oct_analyzer.mvp_pipeline as mvp_pipeline
client = TestClient(app)
def test_predict_image_invalid_suffix():
response = client.post("/predict", files={"file": ("test.txt", b"dummy")})
assert response.status_code == 400
assert response.json() == {"detail": "Unsupported file type"}
def test_predict_image_success(monkeypatch):
class MockClassifier:
def predict(self, img_path, gradcam=True):
return {
"Level1": {"prediction": "NORMAL", "confidence": 0.99},
"Final_Diagnosis": "NORMAL",
"gradcams": {}
}
monkeypatch.setattr(api_module, "get_classifier", lambda: MockClassifier())
img = np.zeros((10, 10, 3), dtype=np.uint8)
with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as tmp:
cv2.imwrite(tmp.name, img)
tmp_path = tmp.name
try:
with open(tmp_path, "rb") as f:
response = client.post("/predict", files={"file": ("test.png", f, "image/png")})
assert response.status_code == 200
data = response.json()
assert data["Level1"]["prediction"] == "NORMAL"
assert data["Final_Diagnosis"] == "NORMAL"
finally:
os.remove(tmp_path)
def test_predict_image_error(monkeypatch):
class MockClassifier:
def predict(self, img_path, gradcam=True):
return {"error": "Mock error"}
monkeypatch.setattr(api_module, "get_classifier", lambda: MockClassifier())
img = np.zeros((10, 10, 3), dtype=np.uint8)
with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as tmp:
cv2.imwrite(tmp.name, img)
tmp_path = tmp.name
try:
with open(tmp_path, "rb") as f:
response = client.post("/predict", files={"file": ("test.png", f, "image/png")})
assert response.status_code == 400
assert response.json()["detail"] == "Mock error"
finally:
os.remove(tmp_path)
def test_classifier_wrapper_singleton(monkeypatch):
ci.ClassifierWrapper._instance = None
class MockPipeline:
def __init__(self, *args, **kwargs):
pass
def predict(self, image_path, gradcam=True):
return {"mock": True}
monkeypatch.setattr(ci, "OCTInferencePipeline", MockPipeline)
wrapper1 = get_classifier()
wrapper2 = get_classifier()
assert wrapper1 is wrapper2
assert wrapper1.predict("dummy_path") == {"mock": True}
def test_classifier_wrapper_missing_pipeline(monkeypatch):
ci.ClassifierWrapper._instance = None
monkeypatch.setattr(ci, "OCTInferencePipeline", None)
with pytest.raises(RuntimeError, match="OCTInferencePipeline is not available."):
ClassifierWrapper()
def test_process_scan_classifier_exception(monkeypatch):
def mock_get_classifier():
raise ValueError("Simulated classifier error")
monkeypatch.setattr(ci, "get_classifier", mock_get_classifier)
from backend.oct_analyzer.scan_types import NormalizedScan
scan = NormalizedScan(
volume=np.zeros((1, 10, 10), dtype=np.float32),
spacing_mm=(1.0, 1.0, 1.0),
source_format="vol",
metadata={},
warnings=[]
)
result = mvp_pipeline.process_scan(scan)
assert result["level1"] == {}
def test_process_scan_classifier_success(monkeypatch):
class MockClassifier:
def predict(self, img_path, gradcam=True):
return {
"Level1": {"prediction": "NORMAL", "confidence": 0.99},
"Final_Diagnosis": "NORMAL"
}
monkeypatch.setattr(ci, "get_classifier", lambda: MockClassifier())
from backend.oct_analyzer.scan_types import NormalizedScan
scan = NormalizedScan(
volume=np.zeros((1, 10, 10), dtype=np.float32),
spacing_mm=(1.0, 1.0, 1.0),
source_format="vol",
metadata={},
warnings=[]
)
result = mvp_pipeline.process_scan(scan)
assert result["diagnosis"] == "NORMAL"
assert result["confidence"] == 0.99
assert result["level1"] == {"prediction": "NORMAL", "confidence": 0.99}
def test_classifier_integration_import_error(monkeypatch):
import builtins
import importlib
real_import = builtins.__import__
def fake_import(name, *args, **kwargs):
if "scripts.inference_pipeline" in name:
raise ImportError("Simulated import error")
return real_import(name, *args, **kwargs)
monkeypatch.setattr(builtins, "__import__", fake_import)
# Reload the module to trigger the import again
importlib.reload(ci)
assert ci.OCTInferencePipeline is None
# Restore for other tests
importlib.reload(ci)
def test_main_cli(monkeypatch):
from backend.oct_analyzer import main
import builtins
# Mock load_oct_volume to simulate success
def mock_load(path):
return np.zeros((10, 10, 10)), (1.0, 1.0, 1.0)
monkeypatch.setattr(main, "load_oct_volume", mock_load)
# Mock pipeline
monkeypatch.setattr(main, "get_preprocessing_pipeline", lambda: lambda x: x)
# Mock flatten
monkeypatch.setattr(main, "flatten_volume_to_rpe", lambda x: x)
main.main()
def test_main_cli_execution(monkeypatch):
import runpy
import sys
# Mock load_oct_volume to prevent file errors
import backend.oct_analyzer.main as main
monkeypatch.setattr(main, "load_oct_volume", lambda path: (np.zeros((10,10,10)), (1,1,1)))
monkeypatch.setattr(main, "get_preprocessing_pipeline", lambda: lambda x: x)
monkeypatch.setattr(main, "flatten_volume_to_rpe", lambda x: x)
# Run module as main
runpy.run_module("backend.oct_analyzer.main", run_name="__main__")
|