Add stored-vector point search path
This commit is contained in:
@@ -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
@@ -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]
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -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()
|
||||||
|
|||||||
Reference in New Issue
Block a user