from __future__ import annotations from types import SimpleNamespace from unittest.mock import MagicMock, patch from sqlalchemy import create_engine, select from sqlalchemy.orm import Session, sessionmaker from sqlalchemy.pool import StaticPool from app.db.base import Base from app.models.few_shot_sample import FewShotSample from app.models.risk_observation import RiskObservation from app.services.few_shot_retrieval import FewShotRetriever from app.services.few_shot_store import FewShotStore, stable_vector_id def _session() -> Session: engine = create_engine( "sqlite+pysqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool, ) Base.metadata.create_all(engine) return sessionmaker(bind=engine)() def _observation(*, tenant_id: str, key: str) -> RiskObservation: return RiskObservation( tenant_id=tenant_id, observation_key=key, subject_type="expense_claim", subject_key="claim:1", risk_type="duplicate_invoice", risk_signal="duplicate_invoice", risk_level="high", ) def _sample( *, sample_id: str, tenant_id: str, sample_key: str, version: str, ) -> FewShotSample: return FewShotSample( id=sample_id, tenant_id=tenant_id, sample_key=sample_key, scene="expense_reimbursement", policy_ref="TRAVEL-001", rule_version=version, label="confirmed", case_text="同一发票重复报销", conclusion_text=f"历史结论 {version}", payload_json={"risk_signal": "duplicate_invoice"}, status="active", ) def test_keys_are_unique_inside_tenant_but_reusable_across_tenants() -> None: with _session() as db: db.add_all( [ _observation(tenant_id="tenant-a", key="same-key"), _observation(tenant_id="tenant-b", key="same-key"), _sample( sample_id="sample-a", tenant_id="tenant-a", sample_key="same-sample", version="v1", ), _sample( sample_id="sample-b", tenant_id="tenant-b", sample_key="same-sample", version="v1", ), ] ) db.commit() assert len(db.scalars(select(RiskObservation)).all()) == 2 assert len(db.scalars(select(FewShotSample)).all()) == 2 def test_vector_id_is_stable_and_tenant_scoped() -> None: first = stable_vector_id(tenant_id="tenant-a", sample_id="sample-1") repeated = stable_vector_id(tenant_id="tenant-a", sample_id="sample-1") other_tenant = stable_vector_id(tenant_id="tenant-b", sample_id="sample-1") assert first == repeated assert first != other_tenant def test_store_upsert_replaces_legacy_vector_and_writes_tenant_payload() -> None: provider = MagicMock() provider.embed.return_value = [[0.1, 0.2]] store = FewShotStore(provider) client = MagicMock() store._client = client sample = SimpleNamespace( id="sample-1", tenant_id="tenant-a", sample_key="key", scene="expense_reimbursement", policy_ref="TRAVEL-001", rule_version="v2", label="false_positive", domain="expense", risk_type="duplicate_invoice", risk_level="high", status="active", case_text="案例", conclusion_text="改判为误报", payload_json={}, vector_id="legacy-random-vector-id", ) with patch.object(store, "_ensure_collection", return_value=True): vector_id = store.upsert(sample) assert vector_id == stable_vector_id(tenant_id="tenant-a", sample_id="sample-1") client.delete.assert_called_once() points = client.upsert.call_args.kwargs["points"] point = points[0] assert point["id"] == vector_id assert point["payload"]["tenant_id"] == "tenant-a" assert point["payload"]["rule_version"] == "v2" assert point["payload"]["label"] == "false_positive" assert points[1]["id"] == "legacy-random-vector-id" assert points[1]["payload"]["label"] == "false_positive" def test_existing_collection_receives_required_payload_indexes() -> None: provider = MagicMock() store = FewShotStore(provider) client = MagicMock() client.get_collection.return_value = SimpleNamespace() store._client = client assert store._ensure_collection() is True fields = {call.kwargs["field_name"] for call in client.create_payload_index.call_args_list} assert {"tenant_id", "scene", "policy_ref", "rule_version", "status"} <= fields client.create_collection.assert_not_called() def test_search_always_filters_tenant_and_supports_rule_identity() -> None: provider = MagicMock() provider.embed.return_value = [[0.1, 0.2]] store = FewShotStore(provider) client = MagicMock() client.query_points.return_value = SimpleNamespace(points=[]) store._client = client with patch.object(store, "_ensure_collection", return_value=True): assert ( store.search( "重复发票", tenant_id="tenant-a", scene="expense_reimbursement", policy_ref="TRAVEL-001", rule_version="v2", ) == [] ) filter_payload = client.query_points.call_args.kwargs["query_filter"].model_dump() must = filter_payload["must"] assert any(item["key"] == "tenant_id" and item["match"]["value"] == "tenant-a" for item in must) assert any(item["key"] == "scene" for item in must) assert any(item["key"] == "policy_ref" for item in must) assert any(item["key"] == "rule_version" for item in must) def test_expense_case_retrieval_db_rechecks_tenant_status_and_marks_old_version() -> None: with _session() as db: current = _sample( sample_id="current", tenant_id="tenant-a", sample_key="current-key", version="v2", ) stale = _sample( sample_id="stale", tenant_id="tenant-a", sample_key="stale-key", version="v1", ) other_tenant = _sample( sample_id="other", tenant_id="tenant-b", sample_key="other-key", version="v2", ) db.add_all([current, stale, other_tenant]) db.commit() store = MagicMock(spec=FewShotStore) store.search.side_effect = [ [{"sample_id": "current", "score": 0.95}], [ {"sample_id": "current", "score": 0.95}, {"sample_id": "stale", "score": 0.8}, {"sample_id": "other", "score": 0.99}, ], ] retriever = FewShotRetriever(store, db) evidence = retriever.retrieve_for_expense_case( tenant_id="tenant-a", scene="expense_reimbursement", policy_ref="TRAVEL-001", rule_version="v2", query="重复发票", top_k=3, ) assert [item["sample_id"] for item in evidence] == ["current", "stale"] assert evidence[0]["version_status"] == "matched" assert evidence[0]["advisory_only"] is True assert evidence[1]["version_status"] == "stale" assert evidence[1]["stale"] is True def test_qdrant_unavailable_fails_closed() -> None: store = FewShotStore(MagicMock()) with patch.object(store, "_ensure_collection", return_value=False): assert store.search("案例", tenant_id="tenant-a") == [] assert store.upsert(SimpleNamespace(tenant_id="tenant-a", id="sample-1")) is None