feat(ai): add tenant-safe hierarchical expense learning

This commit is contained in:
caoxiaozhu
2026-07-16 14:30:41 +08:00
parent 6bdf65bc24
commit ee88a36baf
65 changed files with 6909 additions and 232 deletions

View 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