feat: add qdrant vector adapter
This commit is contained in:
@@ -0,0 +1,312 @@
|
||||
import json
|
||||
from uuid import NAMESPACE_URL, uuid5
|
||||
|
||||
import pytest
|
||||
import requests
|
||||
|
||||
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.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
|
||||
|
||||
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": "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":
|
||||
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"}})
|
||||
wanted = sorted(
|
||||
_match_points(self.points.values(), json["filter"]),
|
||||
key=lambda point: point["payload"]["record_key"],
|
||||
)
|
||||
return FakeResponse(200, {"result": {"points": wanted}})
|
||||
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
def _store(fake: FakeQdrantHttp) -> QdrantVectorStore:
|
||||
return QdrantVectorStore(
|
||||
base_url="http://qdrant:6333",
|
||||
collection="workspace-semantic",
|
||||
workspace_id="demo",
|
||||
expected_dimension=1024,
|
||||
request=fake.request,
|
||||
)
|
||||
|
||||
|
||||
def test_point_id_is_deterministic_uuidv5():
|
||||
assert point_id("demo", "memory", "memory:1") == str(
|
||||
uuid5(NAMESPACE_URL, "thothii:demo:memory:memory:1")
|
||||
)
|
||||
|
||||
|
||||
def test_upsert_creates_collection_and_keyword_indexes_idempotently():
|
||||
fake = FakeQdrantHttp()
|
||||
store = _store(fake)
|
||||
|
||||
assert store.upsert("memory", [_write_record("memory:1", "memory")]) == 1
|
||||
assert store.upsert("memory", [_write_record("memory:1", "memory")]) == 1
|
||||
|
||||
creates = [call for call in fake.calls if call[0] == "PUT" and call[1].endswith("/collections/workspace-semantic")]
|
||||
assert len(creates) == 1
|
||||
assert creates[0][2] == {"vectors": {"size": 1024, "distance": "Cosine"}}
|
||||
assert fake.payload_indexes == {
|
||||
"content_hash",
|
||||
"document_id",
|
||||
"kind",
|
||||
"record_key",
|
||||
"record_kind",
|
||||
"vector_generation",
|
||||
"workspace_id",
|
||||
}
|
||||
|
||||
|
||||
def test_upsert_refuses_collection_dimension_or_distance_mismatch_without_recreating():
|
||||
fake = FakeQdrantHttp(dimension=384, distance="Dot")
|
||||
fake.collection = {"vectors": {"size": 384, "distance": "Dot"}}
|
||||
store = _store(fake)
|
||||
|
||||
with pytest.raises(VectorStoreError, match="Qdrant collection configuration mismatch"):
|
||||
store.upsert("memory", [_write_record("memory:1", "memory")])
|
||||
|
||||
creates = [call for call in fake.calls if call[0] == "PUT" and call[1].endswith("/collections/workspace-semantic")]
|
||||
assert creates == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("record", "semantic_kind"),
|
||||
[
|
||||
(_write_record("schema_column:patients.id", "schema_column"), "schema"),
|
||||
(
|
||||
_write_record(
|
||||
"demo:gen:11111111111111111111111111111111:chunk:1",
|
||||
"evidence",
|
||||
metadata={
|
||||
"workspace_id": "demo",
|
||||
"vector_generation": "gen:11111111111111111111111111111111",
|
||||
"document_id": "doc:abc",
|
||||
},
|
||||
),
|
||||
"evidence",
|
||||
),
|
||||
(_write_record("memory:1", "memory"), "memory"),
|
||||
],
|
||||
)
|
||||
def test_upsert_serializes_qdrant_point_payloads(record, semantic_kind):
|
||||
fake = FakeQdrantHttp()
|
||||
store = _store(fake)
|
||||
|
||||
store.upsert("memory" if semantic_kind == "memory" else "evidence" if semantic_kind == "evidence" else "schema_records", [record])
|
||||
|
||||
point = next(iter(fake.points.values()))
|
||||
assert point["id"] == point_id("demo", semantic_kind, record.record.id)
|
||||
assert point["vector"] == record.embedding
|
||||
assert point["payload"]["workspace_id"] == "demo"
|
||||
assert point["payload"]["kind"] == semantic_kind
|
||||
assert point["payload"]["record_kind"] == record.record.kind
|
||||
assert point["payload"]["record_key"] == record.record.id
|
||||
assert point["payload"]["content_hash"] == record.content_hash
|
||||
|
||||
|
||||
def test_search_filters_by_workspace_and_allowed_record_kinds():
|
||||
fake = FakeQdrantHttp()
|
||||
store = _store(fake)
|
||||
store.upsert("memory", [_write_record("memory:1", "memory")])
|
||||
other = next(iter(fake.points.values())).copy()
|
||||
other["id"] = point_id("other", "memory", "memory:2")
|
||||
other["payload"] = {**other["payload"], "workspace_id": "other", "record_key": "memory:2"}
|
||||
fake.points[other["id"]] = other
|
||||
solved = next(iter(fake.points.values())).copy()
|
||||
solved["id"] = point_id("demo", "memory", "solved:1")
|
||||
solved["payload"] = {**solved["payload"], "record_key": "solved:1", "record_kind": "solved_question"}
|
||||
fake.points[solved["id"]] = solved
|
||||
|
||||
hits = store.search(["memory"], [0.2] * 1024, limit=5, kinds=["memory"])
|
||||
|
||||
assert [hit.id for hit in hits] == ["memory:1"]
|
||||
query_call = next(call for call in fake.calls if call[0] == "POST" and call[1].endswith("/points/query?wait=true") is False and call[1].endswith("/points/query"))
|
||||
assert query_call[2]["filter"] == {
|
||||
"must": [
|
||||
{"key": "workspace_id", "match": {"value": "demo"}},
|
||||
{"key": "record_kind", "match": {"any": ["memory"]}},
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
def test_existing_hashes_health_and_exact_generation_inventory_and_delete():
|
||||
fake = FakeQdrantHttp()
|
||||
store = _store(fake)
|
||||
generation = "gen:" + "1" * 32
|
||||
keep = "gen:" + "2" * 32
|
||||
store.upsert("evidence", [
|
||||
_write_record(
|
||||
f"demo:{generation}:chunk:1",
|
||||
"evidence",
|
||||
metadata={"workspace_id": "demo", "vector_generation": generation, "document_id": "doc:1"},
|
||||
),
|
||||
_write_record(
|
||||
f"demo:{keep}:chunk:2",
|
||||
"evidence",
|
||||
metadata={"workspace_id": "demo", "vector_generation": keep, "document_id": "doc:2"},
|
||||
),
|
||||
])
|
||||
|
||||
assert store.existing_hashes("evidence", ["evidence"]) == {
|
||||
f"demo:{generation}:chunk:1": "sha256:" + "a" * 64,
|
||||
f"demo:{keep}:chunk:2": "sha256:" + "a" * 64,
|
||||
}
|
||||
assert store.list_evidence_generations("evidence", "demo") == [generation, keep]
|
||||
assert store.delete_generation("evidence", generation, "demo") == 1
|
||||
assert store.list_evidence_generations("evidence", "demo") == [keep]
|
||||
|
||||
health = store.health()
|
||||
assert health.ok is True
|
||||
assert health.read_reachable is True
|
||||
assert health.write_reachable is True
|
||||
assert health.observed_dimensions == (1024,)
|
||||
assert health.dimension_compatible is True
|
||||
|
||||
|
||||
def test_sanitizes_timeout_and_malformed_responses():
|
||||
fake = FakeQdrantHttp()
|
||||
store = _store(fake)
|
||||
fake.fail_request = requests.Timeout("dial tcp 10.0.0.9:6333: i/o timeout")
|
||||
|
||||
with pytest.raises(VectorStoreError, match="Qdrant request failed") as timeout:
|
||||
store.search(["memory"], [0.2] * 1024, limit=1)
|
||||
assert "10.0.0.9" not in str(timeout.value)
|
||||
|
||||
fake.fail_request = None
|
||||
store.upsert("memory", [_write_record("memory:1", "memory")])
|
||||
fake.malformed_query = True
|
||||
with pytest.raises(VectorStoreError, match="Qdrant returned malformed query response"):
|
||||
store.search(["memory"], [0.2] * 1024, limit=1)
|
||||
|
||||
fake.malformed_query = False
|
||||
fake.malformed_scroll = True
|
||||
with pytest.raises(VectorStoreError, match="Qdrant returned malformed scroll response"):
|
||||
store.existing_hashes("memory", ["memory"])
|
||||
@@ -3,14 +3,15 @@ from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from tht.adapters.vector.thoth_http import ThothHttpVectorStore
|
||||
from tht.adapters.vector.legacy_direct import LegacyDirectVectorStore
|
||||
from tht.adapters.vector.qdrant import QdrantVectorStore
|
||||
from tht.adapters.vector.thoth_http import ThothHttpVectorStore
|
||||
from tht.evidence.model import EvidenceDoc
|
||||
from tht.ports.vector import (
|
||||
VectorHit,
|
||||
VectorReadUnavailable,
|
||||
VectorRecord,
|
||||
VectorStore,
|
||||
VectorReadUnavailable,
|
||||
VectorWriteRecord,
|
||||
VectorWriteUnavailable,
|
||||
)
|
||||
@@ -158,11 +159,13 @@ def test_http_store_is_runtime_vector_store():
|
||||
|
||||
|
||||
def test_vector_contract_is_exported_from_public_packages():
|
||||
from tht.adapters.vector import QdrantVectorStore as PublicQdrantStore
|
||||
from tht.adapters.vector import ThothHttpVectorStore as PublicHttpStore
|
||||
from tht.ports import VectorReadUnavailable as PublicVectorReadUnavailable
|
||||
from tht.ports import VectorStore as PublicVectorStore
|
||||
from tht.ports import VectorWriteRecord as PublicVectorWriteRecord
|
||||
from tht.ports import VectorReadUnavailable as PublicVectorReadUnavailable
|
||||
|
||||
assert PublicQdrantStore is QdrantVectorStore
|
||||
assert PublicHttpStore is ThothHttpVectorStore
|
||||
assert PublicVectorStore is VectorStore
|
||||
assert PublicVectorWriteRecord is VectorWriteRecord
|
||||
@@ -248,3 +251,24 @@ def test_legacy_direct_search_requires_a_strict_positive_integer_limit(limit):
|
||||
|
||||
with pytest.raises(ValueError, match="positive integer"):
|
||||
store.search(["memory"], [0.1], limit=limit)
|
||||
|
||||
|
||||
def test_qdrant_store_is_runtime_vector_store():
|
||||
store = QdrantVectorStore(
|
||||
base_url="http://qdrant:6333",
|
||||
collection="workspace-semantic",
|
||||
workspace_id="demo",
|
||||
expected_dimension=1024,
|
||||
request=lambda *args, **kwargs: MagicMock(
|
||||
ok=True,
|
||||
status_code=200,
|
||||
text='{"result":{"config":{"params":{"vectors":{"size":1024,"distance":"Cosine"}}},"payload_schema":{}}}',
|
||||
json=lambda: {
|
||||
"result": {
|
||||
"config": {"params": {"vectors": {"size": 1024, "distance": "Cosine"}}},
|
||||
"payload_schema": {},
|
||||
}
|
||||
},
|
||||
),
|
||||
)
|
||||
assert isinstance(store, VectorStore)
|
||||
|
||||
Reference in New Issue
Block a user