diff --git a/gateway/main.py b/gateway/main.py index ceea5f2..3bedcc3 100644 --- a/gateway/main.py +++ b/gateway/main.py @@ -824,6 +824,11 @@ async def vectors_search_vector(payload: Dict[str, Any]): return await _post_json(f"{QDRANT_SVC_URL}/search/vector", payload, preserve_client_errors=True, upstream_name="qdrant") +@app.post("/vectors/search/point") +async def vectors_search_point(payload: Dict[str, Any]): + return await _post_json(f"{QDRANT_SVC_URL}/search/point", payload, preserve_client_errors=True, upstream_name="qdrant") + + @app.post("/vectors/delete") async def vectors_delete(payload: Dict[str, Any]): return await _post_json(f"{QDRANT_SVC_URL}/delete", payload, preserve_client_errors=True, upstream_name="qdrant") diff --git a/qdrant/main.py b/qdrant/main.py index a3f92af..81fe525 100644 --- a/qdrant/main.py +++ b/qdrant/main.py @@ -169,6 +169,16 @@ class SearchVectorRequest(BaseModel): exact: bool = False indexed_only: bool = False +class SearchPointRequest(BaseModel): + id: int = Field(ge=1) + limit: int = Field(default=5, ge=1, le=100) + score_threshold: Optional[float] = Field(default=None, ge=0.0, le=1.0) + collection: Optional[str] = None + filter_metadata: Dict[str, Any] = Field(default_factory=dict) + hnsw_ef: Optional[int] = Field(default=None, ge=1, le=512) + exact: bool = False + indexed_only: bool = False + class DeleteRequest(BaseModel): ids: List[str] @@ -271,6 +281,22 @@ def _id_filter(original_id: str) -> Filter: return Filter(must=[FieldCondition(key="_original_id", match=MatchValue(value=original_id))]) +def _lookup_point_id(raw: str): + """Parse an existing Qdrant point id without generating a new id.""" + try: + value = int(raw) + if value < 0: + raise ValueError + return value + except (TypeError, ValueError): + pass + + try: + return str(uuid.UUID(raw)) + except (TypeError, ValueError, AttributeError): + raise HTTPException(400, "Invalid point ID. Expected an unsigned integer or UUID.") + + def _point_id(raw: Optional[str]) -> str: """Return a Qdrant-compatible point id. @@ -592,6 +618,43 @@ async def search_file( vector = await _embed_bytes(data) return _do_search(vector, int(limit), score_threshold, collection, filter_metadata, hnsw_ef, exact, indexed_only) +@app.post("/search/point") +def search_point(req: SearchPointRequest): + """ + Search using an already-indexed Qdrant point. + + This avoids re-downloading and re-embedding the source artwork through CLIP. + """ + col = _col(req.collection) + + try: + points = client.retrieve( + collection_name=col, + ids=[req.id], + with_vectors=True, + with_payload=False, + ) + except Exception as e: + raise HTTPException(500, f"Qdrant point lookup failed: {e}") + + if not points: + raise HTTPException(404, f"Point '{req.id}' not found") + + vector = points[0].vector + + if not isinstance(vector, list) or not vector: + raise HTTPException(500, f"Point '{req.id}' has no usable vector") + + return _do_search( + vector, + req.limit, + req.score_threshold, + req.collection, + req.filter_metadata, + req.hnsw_ef, + req.exact, + req.indexed_only, + ) @app.post("/search/vector") def search_vector(req: SearchVectorRequest): @@ -655,7 +718,8 @@ def delete_points(req: DeleteRequest): def get_point(point_id: str, collection: Optional[str] = None): col = _col(collection) try: - points = client.retrieve(collection_name=col, ids=[point_id], with_vectors=True) + pid = _lookup_point_id(point_id) + points = client.retrieve(collection_name=col, ids=[pid], with_vectors=True) if not points: raise HTTPException(404, f"Point '{point_id}' not found") p = points[0] diff --git a/tests/test_gateway_point_search.py b/tests/test_gateway_point_search.py new file mode 100644 index 0000000..474d7e9 --- /dev/null +++ b/tests/test_gateway_point_search.py @@ -0,0 +1,29 @@ +from __future__ import annotations + +import asyncio +import unittest +from unittest.mock import AsyncMock, patch + +import gateway.main as module + + +class GatewayPointSearchTests(unittest.TestCase): + def test_point_search_proxies_to_qdrant_point_endpoint(self): + async def run(): + with patch.object(module, "_post_json", new_callable=AsyncMock) as post: + post.return_value = {"results": []} + result = await module.vectors_search_point({"id": 1284, "limit": 5}) + + post.assert_awaited_once_with( + f"{module.QDRANT_SVC_URL}/search/point", + {"id": 1284, "limit": 5}, + preserve_client_errors=True, + upstream_name="qdrant", + ) + self.assertEqual(result, {"results": []}) + + asyncio.run(run()) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_qdrant_clip_semantics.py b/tests/test_qdrant_clip_semantics.py index fd7e156..83a6992 100644 --- a/tests/test_qdrant_clip_semantics.py +++ b/tests/test_qdrant_clip_semantics.py @@ -3,6 +3,7 @@ from __future__ import annotations import asyncio import logging import unittest +from types import SimpleNamespace from unittest.mock import patch import httpx @@ -75,5 +76,41 @@ class QdrantClipSemanticsTests(unittest.IsolatedAsyncioTestCase): module.logger.removeHandler(handler) +class QdrantPointSearchTests(unittest.TestCase): + def test_point_search_retrieves_vector_then_queries(self): + point = SimpleNamespace(id=1284, vector=[0.1, 0.2], payload=None) + result = SimpleNamespace(points=[SimpleNamespace(id=59303, score=0.91, payload={"id": 59303})]) + fake = SimpleNamespace() + fake.retrieve = lambda **kwargs: [point] + fake.query_points = lambda **kwargs: result + + with patch.object(module, "client", fake): + response = module.search_point(module.SearchPointRequest(id=1284, limit=5)) + + self.assertEqual(response["count"], 1) + self.assertEqual(response["results"][0]["id"], 59303) + + def test_missing_point_is_404_without_search(self): + calls = [] + fake = SimpleNamespace( + retrieve=lambda **kwargs: calls.append(kwargs) or [], + query_points=lambda **kwargs: self.fail("query_points must not run"), + ) + + with patch.object(module, "client", fake): + with self.assertRaises(module.HTTPException) as raised: + module.search_point(module.SearchPointRequest(id=1284)) + + self.assertEqual(raised.exception.status_code, 404) + self.assertEqual(calls[0]["ids"], [1284]) + + def test_point_lookup_does_not_generate_an_id(self): + self.assertEqual(module._lookup_point_id("1284"), 1284) + self.assertEqual(module._lookup_point_id("0001284"), 1284) + with self.assertRaises(module.HTTPException) as raised: + module._lookup_point_id("not-a-point") + self.assertEqual(raised.exception.status_code, 400) + + if __name__ == "__main__": unittest.main()