| """Unit tests for confidence/conflicts.py.""" |
|
|
| from __future__ import annotations |
|
|
| from dataclasses import dataclass |
| from typing import Optional |
|
|
| from confidence.conflicts import ConflictDetector |
| from models.providers import ProviderCapability |
| from providers.base import ProviderResult |
|
|
|
|
| @dataclass |
| class FakeBox: |
| detector: str |
| confidence: float |
|
|
|
|
| @dataclass |
| class FakeMatch: |
| query_face_index: int |
| best_match: Optional[str] |
| distance: float |
| recognizer: str |
|
|
|
|
| def _make_result(provider: str, capability: ProviderCapability, |
| normalized: dict, success: bool = True) -> ProviderResult: |
| return ProviderResult( |
| provider=provider, |
| capability=capability, |
| success=success, |
| elapsed_ms=10.0, |
| normalized=normalized, |
| ) |
|
|
|
|
| class TestConflictDetector: |
| def test_no_conflicts_when_single_provider(self): |
| detector = ConflictDetector() |
| results = { |
| "haar": _make_result("haar", ProviderCapability.DETECTION, |
| {"num_faces": 1, "boxes": [{"x": 0, "y": 0, "w": 10, "h": 10}], |
| "confidences": [1.0]}), |
| } |
| boxes = [FakeBox(detector="haar", confidence=1.0)] |
| conflicts = detector.detect(results, boxes, []) |
| assert conflicts == [] |
|
|
| def test_face_count_mismatch_detected(self): |
| detector = ConflictDetector() |
| results = { |
| "haar": _make_result("haar", ProviderCapability.DETECTION, {"num_faces": 1}), |
| "dnn": _make_result("dnn", ProviderCapability.DETECTION, {"num_faces": 2}), |
| } |
| conflicts = detector.detect(results, [], []) |
| assert len(conflicts) == 1 |
| assert conflicts[0].kind == "face_count_mismatch" |
| assert "haar" in conflicts[0].providers |
| assert "dnn" in conflicts[0].providers |
|
|
| def test_face_count_agreement_no_conflict(self): |
| detector = ConflictDetector() |
| results = { |
| "haar": _make_result("haar", ProviderCapability.DETECTION, {"num_faces": 2}), |
| "dnn": _make_result("dnn", ProviderCapability.DETECTION, {"num_faces": 2}), |
| } |
| conflicts = detector.detect(results, [], []) |
| assert conflicts == [] |
|
|
| def test_match_disagreement_detected(self): |
| detector = ConflictDetector() |
| matches = [ |
| FakeMatch(query_face_index=0, best_match="alice", distance=0.3, recognizer="face_recognition"), |
| FakeMatch(query_face_index=0, best_match="bob", distance=0.4, recognizer="deepface"), |
| ] |
| conflicts = detector.detect({}, [], matches) |
| assert len(conflicts) == 1 |
| assert conflicts[0].kind == "match_disagreement" |
|
|
| def test_match_agreement_no_conflict(self): |
| detector = ConflictDetector() |
| matches = [ |
| FakeMatch(query_face_index=0, best_match="alice", distance=0.3, recognizer="face_recognition"), |
| FakeMatch(query_face_index=0, best_match="alice", distance=0.4, recognizer="deepface"), |
| ] |
| conflicts = detector.detect({}, [], matches) |
| assert conflicts == [] |
|
|
| def test_quality_disagreement_detected(self): |
| detector = ConflictDetector() |
| results = { |
| "image_quality": _make_result("image_quality", ProviderCapability.IMAGE_ANALYSIS, |
| {"quality_score": 0.9}), |
| "image_properties": _make_result("image_properties", ProviderCapability.IMAGE_ANALYSIS, |
| {"quality_score": 0.4}), |
| } |
| conflicts = detector.detect(results, [], []) |
| assert len(conflicts) == 1 |
| assert conflicts[0].kind == "quality_disagreement" |
|
|
| def test_format_mismatch_detected(self): |
| detector = ConflictDetector() |
| results = { |
| "exif": _make_result("exif", ProviderCapability.METADATA, {"format": "JPEG"}), |
| "xmp": _make_result("xmp", ProviderCapability.METADATA, {"format": "PNG"}), |
| } |
| conflicts = detector.detect(results, [], []) |
| assert len(conflicts) == 1 |
| assert conflicts[0].kind == "format_mismatch" |
|
|
| def test_failed_results_ignored(self): |
| detector = ConflictDetector() |
| results = { |
| "haar": _make_result("haar", ProviderCapability.DETECTION, {"num_faces": 1}, success=False), |
| "dnn": _make_result("dnn", ProviderCapability.DETECTION, {"num_faces": 2}), |
| } |
| |
| conflicts = detector.detect(results, [], []) |
| |
| |
| |
| assert conflicts == [] |
|
|