Files
ThothII/harness/tests/qdrant_test_helpers.py
T

168 lines
6.4 KiB
Python

"""Shared fake Qdrant HTTP boundary for adapter and CLI tests."""
import json
from tht.ports.vector import 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.collection is None:
return FakeResponse(404, {"status": "error"})
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.collection is None:
return FakeResponse(404, {"status": "error"})
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":
if self.collection is None:
return FakeResponse(404, {"status": "error"})
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,
)