feat(ai): add tenant-safe hierarchical expense learning
This commit is contained in:
231
server/tests/test_tenant_safe_few_shot_foundation.py
Normal file
231
server/tests/test_tenant_safe_few_shot_foundation.py
Normal file
@@ -0,0 +1,231 @@
|
||||
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
|
||||
Reference in New Issue
Block a user