diff --git a/livexface/__init__.py b/livexface/__init__.py index a183e25..3fa335e 100644 --- a/livexface/__init__.py +++ b/livexface/__init__.py @@ -7,6 +7,9 @@ VerifyResult, FaceMatch, IdentifyResult, + CrossCollectionSearchMatch, + CrossCollectionSearchResult, + SkippedCollection, LivenessResult, LivenessChallenge, ActiveLivenessResult, @@ -34,6 +37,9 @@ "VerifyResult", "FaceMatch", "IdentifyResult", + "CrossCollectionSearchMatch", + "CrossCollectionSearchResult", + "SkippedCollection", "LivenessResult", "LivenessChallenge", "ActiveLivenessResult", diff --git a/livexface/client.py b/livexface/client.py index 6c42d11..b062ba4 100644 --- a/livexface/client.py +++ b/livexface/client.py @@ -18,6 +18,7 @@ Face, VerifyResult, IdentifyResult, + CrossCollectionSearchResult, LivenessResult, ActiveLivenessResult, LivenessSession, @@ -337,6 +338,24 @@ def identify( ) return IdentifyResult.from_dict(resp) + def search( + self, + image: ImageInput, + collection_ids: Sequence[str] | None = None, + top_k: int = 5, + threshold: float | None = None, + ) -> CrossCollectionSearchResult: + """Search for matching faces across multiple or all collections.""" + fname, fbytes, ftype = _to_bytes_tuple(image) + files = {"image": (fname, fbytes, ftype)} + data: dict[str, str] = {"top_k": str(top_k)} + if collection_ids: + data["collection_ids"] = ",".join(collection_ids) + if threshold is not None: + data["threshold"] = str(threshold) + resp = self._c._request("POST", "/search", files=files, data=data) + return CrossCollectionSearchResult.from_dict(resp) + def liveness(self, collection_id: str, image: ImageInput) -> LivenessResult: """Passive liveness detection — check whether the face in the image is live.""" fname, fbytes, ftype = _to_bytes_tuple(image) diff --git a/livexface/types.py b/livexface/types.py index c1d7c5b..5bd381a 100644 --- a/livexface/types.py +++ b/livexface/types.py @@ -84,6 +84,55 @@ def from_dict(cls, d: dict[str, Any]) -> "IdentifyResult": ) +@dataclass +class CrossCollectionSearchMatch(FaceMatch): + collection_id: str = "" + + @classmethod + def from_dict(cls, d: dict[str, Any]) -> "CrossCollectionSearchMatch": + return cls( + face_id=d.get("faceId", ""), + external_id=d.get("externalId", ""), + confidence=float(d.get("confidence", 0.0)), + metadata=d.get("metadata"), + collection_id=d.get("collectionId", ""), + ) + + +@dataclass +class SkippedCollection: + id: str + name: str + reason: str + + @classmethod + def from_dict(cls, d: dict[str, Any]) -> "SkippedCollection": + return cls( + id=d.get("id", ""), + name=d.get("name", ""), + reason=d.get("reason", ""), + ) + + +@dataclass +class CrossCollectionSearchResult: + matches: list[CrossCollectionSearchMatch] + query_time_ms: int + collections_searched: int + skipped_collections: list[SkippedCollection] + + @classmethod + def from_dict(cls, d: dict[str, Any]) -> "CrossCollectionSearchResult": + return cls( + matches=[CrossCollectionSearchMatch.from_dict(m) for m in d.get("matches", [])], + query_time_ms=d.get("queryTimeMs", 0), + collections_searched=d.get("collectionsSearched", 0), + skipped_collections=[ + SkippedCollection.from_dict(c) for c in d.get("skippedCollections", []) + ], + ) + + @dataclass class LivenessResult: """Result of a passive liveness check. diff --git a/tests/test_contract.py b/tests/test_contract.py index 7269764..575bd49 100644 --- a/tests/test_contract.py +++ b/tests/test_contract.py @@ -50,6 +50,7 @@ "faces.delete": (("c1", "f1"), {}, NO_CONTENT), "faces.verify": (("c1", IMG, "f1"), {"threshold": 0.5}, DATA), "faces.identify": (("c1", IMG), {"top_k": 3, "threshold": 0.5}, DATA), + "faces.search": ((IMG,), {"collection_ids": ["c1", "c2"], "top_k": 3, "threshold": 0.5}, DATA), "faces.liveness": (("c1", IMG), {}, DATA), "faces.active_liveness": (("c1", [IMG] * 5), {}, DATA), "faces.create_liveness_session": ( diff --git a/tests/test_errors.py b/tests/test_errors.py index caf30a5..a43dc79 100644 --- a/tests/test_errors.py +++ b/tests/test_errors.py @@ -29,3 +29,21 @@ def test_api_error_exposes_details(mocker: Any) -> None: assert exc.value.code == "MULTIPLE_FACES" assert exc.value.details == {"faceCount": 2, "faces": []} assert exc.value.request_id == "r-1" + + +def test_search_parses_skips_and_typed_profile_mismatch(mocker: Any) -> None: + client = LiveXFace(api_key="lxf_test", base_url="http://api.test/api/v1") + ok = MagicMock(status_code=200, ok=True) + ok.json.return_value = {"success": True, "data": {"matches": [], "queryTimeMs": 7, "collectionsSearched": 1, "skippedCollections": [{"id": "c2", "name": "Legacy", "reason": "embedding_profile_mismatch"}]}} + request = mocker.patch.object(client._session, "request", return_value=ok) + result = client.faces.search(b"img", collection_ids=["c1", "c2"]) + assert result.skipped_collections[0].reason == "embedding_profile_mismatch" + assert request.call_args.kwargs["data"]["collection_ids"] == "c1,c2" + + conflict = MagicMock(status_code=409, ok=False) + conflict.json.return_value = {"success": False, "error": {"code": "EMBEDDING_PROFILE_MISMATCH", "message": "no compatible collections"}} + request.return_value = conflict + with pytest.raises(LiveXFaceApiError) as exc: + client.faces.search(b"img") + assert exc.value.status_code == 409 + assert exc.value.code == "EMBEDDING_PROFILE_MISMATCH"