143 lines
5.0 KiB
Python
143 lines
5.0 KiB
Python
"""L0 contract: the pinned Qdrant image performs Italian BM25 inference server-side."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import time
|
|
from pathlib import Path
|
|
|
|
import pytest
|
|
import requests
|
|
import yaml
|
|
from testcontainers.core.container import DockerContainer
|
|
|
|
pytestmark = [pytest.mark.l0]
|
|
|
|
|
|
def qdrant_image() -> str:
|
|
compose = Path(__file__).resolve().parents[3] / "compose.yaml"
|
|
image = yaml.safe_load(compose.read_text(encoding="utf-8"))["services"]["qdrant"]["image"]
|
|
assert isinstance(image, str) and "@sha256:" in image
|
|
return image
|
|
|
|
|
|
def request_ok(method: str, url: str, **kwargs: object) -> dict:
|
|
response = requests.request(method, url, timeout=10, **kwargs)
|
|
response.raise_for_status()
|
|
payload = response.json()
|
|
assert isinstance(payload, dict)
|
|
return payload
|
|
|
|
|
|
def wait_for_qdrant(base_url: str) -> None:
|
|
deadline = time.monotonic() + 30
|
|
while time.monotonic() < deadline:
|
|
try:
|
|
if requests.get(f"{base_url}/healthz", timeout=1).ok:
|
|
return
|
|
except requests.RequestException:
|
|
pass
|
|
time.sleep(0.25)
|
|
pytest.fail("the pinned Qdrant container did not become healthy")
|
|
|
|
|
|
def point_ids(base_url: str, collection: str) -> list[int]:
|
|
result = request_ok(
|
|
"POST",
|
|
f"{base_url}/collections/{collection}/points/scroll",
|
|
json={"limit": 100, "with_payload": True, "with_vector": False},
|
|
)
|
|
return sorted(point["id"] for point in result["result"]["points"])
|
|
|
|
|
|
def dense_result_id(base_url: str, collection: str, query: list[float]) -> int:
|
|
result = request_ok(
|
|
"POST",
|
|
f"{base_url}/collections/{collection}/points/query",
|
|
json={"query": query, "limit": 1, "with_payload": False},
|
|
)
|
|
return result["result"]["points"][0]["id"]
|
|
|
|
|
|
def test_pinned_qdrant_image_indexes_and_queries_italian_bm25_server_side():
|
|
with DockerContainer(qdrant_image()).with_exposed_ports(6333) as qdrant:
|
|
base_url = f"http://{qdrant.get_container_host_ip()}:{qdrant.get_exposed_port(6333)}"
|
|
wait_for_qdrant(base_url)
|
|
collection = "italian_bm25_contract"
|
|
request_ok(
|
|
"PUT",
|
|
f"{base_url}/collections/{collection}",
|
|
json={
|
|
"vectors": {"size": 4, "distance": "Cosine"},
|
|
},
|
|
)
|
|
legacy_points = [
|
|
{"id": 10, "vector": [1.0, 0.0, 0.0, 0.0], "payload": {"record_kind": "schema_table"}},
|
|
{"id": 11, "vector": [0.0, 1.0, 0.0, 0.0], "payload": {"record_kind": "schema_column"}},
|
|
{"id": 12, "vector": [0.0, 0.0, 1.0, 0.0], "payload": {"record_kind": "memory"}},
|
|
{"id": 13, "vector": [0.0, 0.0, 0.0, 1.0], "payload": {"record_kind": "solved_question"}},
|
|
]
|
|
request_ok(
|
|
"PUT",
|
|
f"{base_url}/collections/{collection}/points?wait=true",
|
|
json={"points": legacy_points},
|
|
)
|
|
ids_before = point_ids(base_url, collection)
|
|
dense_before = [
|
|
dense_result_id(base_url, collection, point["vector"])
|
|
for point in legacy_points
|
|
]
|
|
assert ids_before == [10, 11, 12, 13]
|
|
assert dense_before == ids_before
|
|
|
|
request_ok(
|
|
"PUT",
|
|
f"{base_url}/collections/{collection}/vectors/bm25",
|
|
json={"sparse": {"modifier": "idf"}},
|
|
)
|
|
configuration = request_ok("GET", f"{base_url}/collections/{collection}")["result"]["config"]["params"]
|
|
assert configuration["vectors"] == {"size": 4, "distance": "Cosine"}
|
|
assert configuration["sparse_vectors"] == {"bm25": {"modifier": "idf"}}
|
|
assert point_ids(base_url, collection) == ids_before
|
|
assert [
|
|
dense_result_id(base_url, collection, point["vector"])
|
|
for point in legacy_points
|
|
] == dense_before
|
|
|
|
document = {"model": "qdrant/bm25", "options": {"language": "italian"}}
|
|
request_ok(
|
|
"PUT",
|
|
f"{base_url}/collections/{collection}/points?wait=true",
|
|
json={
|
|
"points": [
|
|
{
|
|
"id": 1,
|
|
"vector": {
|
|
"": [0.1, 0.2, 0.3, 0.4],
|
|
"bm25": {**document, "text": "ricovero per cardiomiopatia dilatativa"},
|
|
},
|
|
},
|
|
{
|
|
"id": 2,
|
|
"vector": {
|
|
"": [0.4, 0.3, 0.2, 0.1],
|
|
"bm25": {**document, "text": "controllo dermatologico programmato"},
|
|
},
|
|
},
|
|
]
|
|
},
|
|
)
|
|
|
|
result = request_ok(
|
|
"POST",
|
|
f"{base_url}/collections/{collection}/points/query",
|
|
json={
|
|
"query": {**document, "text": "cardiomiopatia"},
|
|
"using": "bm25",
|
|
"limit": 2,
|
|
"with_payload": False,
|
|
},
|
|
)
|
|
|
|
points = result["result"]["points"]
|
|
assert [point["id"] for point in points] == [1]
|