Files
duocthu/apps/ai-service/tests/test_live_datastores.py
T

173 lines
6.1 KiB
Python

from __future__ import annotations
import json
import os
import uuid
from contextlib import suppress
from functools import lru_cache
from pathlib import Path
import pytest
pytestmark = pytest.mark.skipif(
os.getenv("RUN_INTEGRATION") != "1",
reason="set RUN_INTEGRATION=1 with local PostgreSQL and Qdrant running",
)
ROOT = Path(__file__).resolve().parents[3]
CHUNKS = ROOT / "ingestion/data/processed/chunks.jsonl"
PDF = ROOT / "ingestion/data/raw/duoc-thu-quoc-gia-viet-nam-2018.pdf"
MIGRATION = Path(__file__).resolve().parents[1] / "migrations/001_rag_retrieval_trace.sql"
@lru_cache
def _first_real_chunk() -> dict:
with CHUNKS.open(encoding="utf-8") as handle:
record = json.loads(next(handle))
import fitz
from ingestion.extract.page_map import build_page_map
with fitz.open(PDF) as document:
page_map = build_page_map(document)
physical_start, physical_end = record["source_page_range"]
printed_start = page_map[physical_start]
printed_end = page_map[physical_end]
assert printed_start is not None and printed_end is not None
record["printed_page_range"] = [printed_start, printed_end]
return record
def test_real_qdrant_round_trip_uses_real_chunk_and_printed_folio():
from qdrant_client import QdrantClient
from qdrant_client.models import Distance, PointStruct, VectorParams
from adapters.embedding import LocalHashQueryEmbedder
from adapters.qdrant import QdrantRetriever
client = QdrantClient(url="http://localhost:6333")
collection = f"integration_{uuid.uuid4().hex}"
embedder = LocalHashQueryEmbedder(32)
record = _first_real_chunk()
try:
client.create_collection(
collection_name=collection,
vectors_config=VectorParams(size=32, distance=Distance.COSINE),
)
client.upsert(
collection_name=collection,
points=[PointStruct(
id=str(uuid.uuid4()),
vector=embedder.embed_query(record["text"]),
payload=record,
)],
wait=True,
)
hits = QdrantRetriever(client, collection, embedder).search(
record["text"], record["drug_id"], 3,
)
assert [hit.document.doc_id for hit in hits] == [record["chunk_id"]]
assert hits[0].document.text == record["text"]
assert hits[0].document.source_refs[0].printed_page_range == tuple(
record["printed_page_range"]
)
finally:
with suppress(Exception):
client.delete_collection(collection)
def test_real_postgres_migration_insert_and_read_back():
from adapters.postgres import PostgresTraceRepository
repository = PostgresTraceRepository(
"postgresql://duoc_thu:duoc_thu@localhost:5432/duoc_thu"
)
repository.migrate(MIGRATION)
trace_id = repository.save(
query="Liều abacavir?",
subject_scope="human",
intent="fact_lookup",
decision="answerable",
reason="grounded_evidence_available",
resolved_drug_id="abacavir",
citations=({
"chunk_id": "abacavir__ten_chung_quoc_te__0",
"printed_page_start": 101,
"printed_page_end": 103,
},),
)
stored = repository.get(trace_id)
assert stored is not None
assert stored.query == "Liều abacavir?"
assert stored.resolved_drug_id == "abacavir"
assert stored.citations[0]["printed_page_start"] == 101
def test_api_round_trip_uses_qdrant_and_persists_postgres_trace():
from fastapi.testclient import TestClient
from qdrant_client import QdrantClient
from qdrant_client.models import Distance, PointStruct, VectorParams
from adapters.embedding import LocalHashQueryEmbedder
from adapters.postgres import PostgresTraceRepository
from adapters.qdrant import QdrantParentStore, QdrantRetriever
from config import Settings
from main import create_app
from rag.answer import GroundedAnswerService
from rag.routing import CatalogDrugResolver, QueryRoutingService
from rag.service import EvidencePolicy, RetrievalService
qdrant = QdrantClient(url="http://localhost:6333")
collection = f"integration_{uuid.uuid4().hex}"
embedder = LocalHashQueryEmbedder(32)
record = dict(_first_real_chunk())
traces = PostgresTraceRepository(
"postgresql://duoc_thu:duoc_thu@localhost:5432/duoc_thu"
)
traces.migrate(MIGRATION)
try:
qdrant.create_collection(
collection_name=collection,
vectors_config=VectorParams(size=32, distance=Distance.COSINE),
)
qdrant.upsert(
collection_name=collection,
points=[PointStruct(
id=str(uuid.uuid4()),
vector=embedder.embed_query(record["text"]),
payload=record,
)],
wait=True,
)
retrieval = RetrievalService(
QdrantRetriever(qdrant, collection, embedder),
QdrantParentStore(qdrant, collection),
EvidencePolicy(minimum_score=0.01),
)
answers = GroundedAnswerService(QueryRoutingService(
retrieval,
CatalogDrugResolver({record["drug_id"]: {record["drug_name"]}}),
))
app = create_app(
settings=Settings(), answer_service=answers, trace_writer=traces,
)
response = TestClient(app).post("/v1/rag/query", json={
"query": record["text"],
"subject_scope": "human",
"intent": "fact_lookup",
})
assert response.status_code == 200
body = response.json()
assert body["decision"] == "answerable"
assert body["citations"][0]["chunk_id"] == record["chunk_id"]
assert body["citations"][0]["printed_page_start"] == (
record["printed_page_range"][0]
)
stored = traces.get(body["trace_id"])
assert stored is not None
assert stored.decision == "answerable"
assert stored.citations[0]["chunk_id"] == record["chunk_id"]
finally:
with suppress(Exception):
qdrant.delete_collection(collection)