Add stored-vector point search path

This commit is contained in:
2026-08-30 16:57:20 +02:00
parent 32ddf70e2f
commit 1e87431dd5
4 changed files with 136 additions and 1 deletions
+5
View File
@@ -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") 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") @app.post("/vectors/delete")
async def vectors_delete(payload: Dict[str, Any]): 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") return await _post_json(f"{QDRANT_SVC_URL}/delete", payload, preserve_client_errors=True, upstream_name="qdrant")
+65 -1
View File
@@ -169,6 +169,16 @@ class SearchVectorRequest(BaseModel):
exact: bool = False exact: bool = False
indexed_only: 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): class DeleteRequest(BaseModel):
ids: List[str] 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))]) 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: def _point_id(raw: Optional[str]) -> str:
"""Return a Qdrant-compatible point id. """Return a Qdrant-compatible point id.
@@ -592,6 +618,43 @@ async def search_file(
vector = await _embed_bytes(data) vector = await _embed_bytes(data)
return _do_search(vector, int(limit), score_threshold, collection, filter_metadata, hnsw_ef, exact, indexed_only) 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") @app.post("/search/vector")
def search_vector(req: SearchVectorRequest): def search_vector(req: SearchVectorRequest):
@@ -655,7 +718,8 @@ def delete_points(req: DeleteRequest):
def get_point(point_id: str, collection: Optional[str] = None): def get_point(point_id: str, collection: Optional[str] = None):
col = _col(collection) col = _col(collection)
try: 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: if not points:
raise HTTPException(404, f"Point '{point_id}' not found") raise HTTPException(404, f"Point '{point_id}' not found")
p = points[0] p = points[0]
+29
View File
@@ -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()
+37
View File
@@ -3,6 +3,7 @@ from __future__ import annotations
import asyncio import asyncio
import logging import logging
import unittest import unittest
from types import SimpleNamespace
from unittest.mock import patch from unittest.mock import patch
import httpx import httpx
@@ -75,5 +76,41 @@ class QdrantClipSemanticsTests(unittest.IsolatedAsyncioTestCase):
module.logger.removeHandler(handler) 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__": if __name__ == "__main__":
unittest.main() unittest.main()