fix: harden task3 qdrant compatibility boundaries
This commit is contained in:
@@ -1,167 +1,11 @@
|
||||
import json
|
||||
from uuid import NAMESPACE_URL, uuid5
|
||||
|
||||
import pytest
|
||||
import requests
|
||||
from qdrant_test_helpers import FakeQdrantHttp, _write_record
|
||||
|
||||
from tht.adapters.vector.qdrant import QdrantVectorStore, point_id
|
||||
from tht.ports.vector import VectorStoreError, VectorWriteRecord
|
||||
from tht.vectorstore.records import VectorRecord
|
||||
|
||||
|
||||
class FakeResponse:
|
||||
def __init__(self, status_code: int, payload=None, text: str | None = None):
|
||||
self.status_code = status_code
|
||||
self._payload = payload
|
||||
self.text = text if text is not None else (
|
||||
"" if payload is None else json.dumps(payload)
|
||||
)
|
||||
|
||||
@property
|
||||
def ok(self) -> bool:
|
||||
return 200 <= self.status_code < 300
|
||||
|
||||
def json(self):
|
||||
if isinstance(self._payload, Exception):
|
||||
raise self._payload
|
||||
return self._payload
|
||||
|
||||
|
||||
class FakeQdrantHttp:
|
||||
def __init__(self, *, dimension=1024, distance="Cosine"):
|
||||
self.dimension = dimension
|
||||
self.distance = distance
|
||||
self.collection = None
|
||||
self.payload_indexes: set[str] = set()
|
||||
self.payload_index_types: dict[str, str] = {}
|
||||
self.points: dict[str, dict] = {}
|
||||
self.calls: list[tuple[str, str, dict | None]] = []
|
||||
self.fail_request: Exception | None = None
|
||||
self.malformed_query = False
|
||||
self.malformed_scroll = False
|
||||
self.scroll_pages: list[dict] | None = None
|
||||
|
||||
def request(self, method, url, *, json=None, timeout=None):
|
||||
self.calls.append((method, url, json))
|
||||
if self.fail_request is not None:
|
||||
raise self.fail_request
|
||||
|
||||
path = url.split("://", 1)[-1].split("/", 1)[-1]
|
||||
path = "/" + path.split("?", 1)[0]
|
||||
|
||||
if method == "GET" and path == "/collections/workspace-semantic":
|
||||
if self.collection is None:
|
||||
return FakeResponse(404, {"status": "error"})
|
||||
return FakeResponse(200, {
|
||||
"result": {
|
||||
"config": {
|
||||
"params": {
|
||||
"vectors": {"size": self.dimension, "distance": self.distance}
|
||||
}
|
||||
},
|
||||
"payload_schema": {
|
||||
field: {"data_type": self.payload_index_types.get(field, "keyword")} for field in sorted(self.payload_indexes)
|
||||
},
|
||||
}
|
||||
})
|
||||
|
||||
if method == "PUT" and path == "/collections/workspace-semantic":
|
||||
self.collection = json
|
||||
self.dimension = json["vectors"]["size"]
|
||||
self.distance = json["vectors"]["distance"]
|
||||
return FakeResponse(200, {"status": "ok"})
|
||||
|
||||
if method == "PUT" and path == "/collections/workspace-semantic/index":
|
||||
self.payload_indexes.add(json["field_name"])
|
||||
return FakeResponse(200, {"status": "ok"})
|
||||
|
||||
if method == "PUT" and path == "/collections/workspace-semantic/points":
|
||||
if self.collection is None:
|
||||
return FakeResponse(404, {"status": "error"})
|
||||
for point in json["points"]:
|
||||
self.points[point["id"]] = point
|
||||
return FakeResponse(200, {"result": {"status": "acknowledged"}})
|
||||
|
||||
if method == "POST" and path == "/collections/workspace-semantic/points/query":
|
||||
if self.malformed_query:
|
||||
return FakeResponse(200, {"result": {"points": "nope"}})
|
||||
wanted = _match_points(self.points.values(), json["filter"])
|
||||
scored = sorted(
|
||||
(
|
||||
{
|
||||
"id": point["id"],
|
||||
"score": point.get("score", 0.9),
|
||||
"payload": point["payload"],
|
||||
}
|
||||
for point in wanted
|
||||
),
|
||||
key=lambda point: (-point["score"], point["payload"]["record_key"]),
|
||||
)
|
||||
return FakeResponse(200, {"result": {"points": scored[: json["limit"]]}})
|
||||
|
||||
if method == "POST" and path == "/collections/workspace-semantic/points/scroll":
|
||||
if self.malformed_scroll:
|
||||
return FakeResponse(200, {"result": {"points": "bad"}})
|
||||
if self.scroll_pages is not None:
|
||||
offset = json.get("offset")
|
||||
for page in self.scroll_pages:
|
||||
if page["offset"] == offset:
|
||||
filtered = _match_points(page["points"], json["filter"])
|
||||
return FakeResponse(200, {
|
||||
"result": {
|
||||
"points": filtered,
|
||||
"next_page_offset": page["next_page_offset"],
|
||||
}
|
||||
})
|
||||
raise AssertionError(("unexpected offset", offset, self.scroll_pages))
|
||||
wanted = sorted(
|
||||
_match_points(self.points.values(), json["filter"]),
|
||||
key=lambda point: point["payload"]["record_key"],
|
||||
)
|
||||
return FakeResponse(200, {"result": {"points": wanted, "next_page_offset": None}})
|
||||
|
||||
if method == "POST" and path == "/collections/workspace-semantic/points/delete":
|
||||
doomed = [point["id"] for point in _match_points(self.points.values(), json["filter"])]
|
||||
for point_id_value in doomed:
|
||||
self.points.pop(point_id_value, None)
|
||||
return FakeResponse(200, {"result": {"status": "acknowledged"}})
|
||||
|
||||
raise AssertionError((method, path, json))
|
||||
|
||||
|
||||
def _match_points(points, flt):
|
||||
matches = []
|
||||
must = flt["must"]
|
||||
for point in points:
|
||||
payload = point["payload"]
|
||||
if all(_match_clause(payload, clause) for clause in must):
|
||||
matches.append(point)
|
||||
return matches
|
||||
|
||||
|
||||
def _match_clause(payload, clause):
|
||||
key = clause["key"]
|
||||
match = clause["match"]
|
||||
if "value" in match:
|
||||
return payload.get(key) == match["value"]
|
||||
if "any" in match:
|
||||
return payload.get(key) in set(match["any"])
|
||||
raise AssertionError(clause)
|
||||
|
||||
|
||||
def _write_record(record_id: str, kind: str, *, metadata=None):
|
||||
return VectorWriteRecord(
|
||||
record=VectorRecord(
|
||||
id=record_id,
|
||||
kind=kind,
|
||||
ref=f"ref:{record_id}",
|
||||
title=f"title:{record_id}",
|
||||
content=f"content:{record_id}",
|
||||
metadata=metadata or {},
|
||||
),
|
||||
embedding=[0.1] * 1024,
|
||||
content_hash="sha256:" + "a" * 64,
|
||||
)
|
||||
from tht.ports.vector import VectorStoreError
|
||||
|
||||
|
||||
def _store(fake: FakeQdrantHttp, *, collection_lifecycle="create_if_missing") -> QdrantVectorStore:
|
||||
@@ -183,6 +27,20 @@ def test_point_id_is_deterministic_uuidv5():
|
||||
|
||||
|
||||
|
||||
def test_require_existing_requires_an_explicit_embedding_dimension():
|
||||
fake = FakeQdrantHttp()
|
||||
with pytest.raises(ValueError, match="expected dimension"):
|
||||
QdrantVectorStore(
|
||||
base_url="http://qdrant:6333",
|
||||
collection="workspace-semantic",
|
||||
workspace_id="demo",
|
||||
expected_dimension=None,
|
||||
collection_lifecycle="require_existing",
|
||||
request=fake.request,
|
||||
)
|
||||
assert fake.calls == []
|
||||
|
||||
|
||||
def test_require_existing_refuses_missing_collection_without_mutations():
|
||||
fake = FakeQdrantHttp()
|
||||
store = _store(fake, collection_lifecycle="require_existing")
|
||||
@@ -255,8 +113,11 @@ def test_require_existing_write_fails_after_collection_is_deleted_without_recrea
|
||||
collection_lifecycle="require_existing", request=request,
|
||||
)
|
||||
|
||||
with pytest.raises(VectorStoreError):
|
||||
from tht.ports.vector import SemanticIndexIncompatibleError
|
||||
|
||||
with pytest.raises(SemanticIndexIncompatibleError) as caught:
|
||||
store.upsert("memory", [_write_record("memory:1", "memory")])
|
||||
assert caught.value.code == "semantic_index_incompatible"
|
||||
assert not [call for call in fake.calls if call[0] == "PUT" and call[1].endswith("/collections/workspace-semantic")]
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user