feat(ai): add tenant-safe hierarchical expense learning
This commit is contained in:
@@ -1,12 +1,11 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from app.services.few_shot_retrieval import FewShotRetriever
|
||||
from app.services.few_shot_store import FewShotStore
|
||||
from app.services.risk_rule_generation import RiskRuleGenerationService
|
||||
from app.services.risk_rule_generation_prompt import build_risk_rule_compiler_messages
|
||||
|
||||
|
||||
@@ -36,7 +35,9 @@ def test_retrieve_returns_injection_blocks_with_token_budget() -> None:
|
||||
]
|
||||
retriever = FewShotRetriever(store)
|
||||
blocks = retriever.retrieve_for_risk_rule_generation(
|
||||
domain="expense", natural_language="同一发票重复报销"
|
||||
tenant_id="tenant-a",
|
||||
domain="expense",
|
||||
natural_language="同一发票重复报销",
|
||||
)
|
||||
assert len(blocks) == 2
|
||||
assert blocks[0]["score"] == 0.9
|
||||
@@ -51,7 +52,13 @@ def test_retrieve_returns_injection_blocks_with_token_budget() -> None:
|
||||
def test_retrieve_empty_case_text_returns_empty() -> None:
|
||||
store = MagicMock(spec=FewShotStore)
|
||||
retriever = FewShotRetriever(store)
|
||||
assert retriever.retrieve_for_risk_rule_generation(natural_language="") == []
|
||||
assert (
|
||||
retriever.retrieve_for_risk_rule_generation(
|
||||
tenant_id="tenant-a",
|
||||
natural_language="",
|
||||
)
|
||||
== []
|
||||
)
|
||||
store.search.assert_not_called()
|
||||
|
||||
|
||||
@@ -62,7 +69,10 @@ def test_retrieve_truncates_overlong_conclusion() -> None:
|
||||
_hit(0.9, "confirmed", long_text),
|
||||
]
|
||||
retriever = FewShotRetriever(store)
|
||||
blocks = retriever.retrieve_for_risk_rule_generation(natural_language="x")
|
||||
blocks = retriever.retrieve_for_risk_rule_generation(
|
||||
tenant_id="tenant-a",
|
||||
natural_language="x",
|
||||
)
|
||||
assert len(blocks) == 1
|
||||
# 超长结论应被截断到单条上限
|
||||
from app.services.few_shot_retrieval import SINGLE_SAMPLE_MAX_CHARS
|
||||
@@ -70,6 +80,45 @@ def test_retrieve_truncates_overlong_conclusion() -> None:
|
||||
assert len(blocks[0]["conclusion"]) <= SINGLE_SAMPLE_MAX_CHARS
|
||||
|
||||
|
||||
def test_rule_generation_forwards_authenticated_tenant_to_retriever(monkeypatch) -> None:
|
||||
fake_retriever = MagicMock()
|
||||
fake_retriever.retrieve_for_risk_rule_generation.return_value = []
|
||||
monkeypatch.setenv("FEW_SHOT_INJECTION_ENABLED", "true")
|
||||
monkeypatch.setattr(
|
||||
FewShotRetriever,
|
||||
"from_session",
|
||||
classmethod(lambda _cls, _session: fake_retriever),
|
||||
)
|
||||
service = RiskRuleGenerationService(MagicMock())
|
||||
|
||||
service._retrieve_few_shot_samples(
|
||||
tenant_id="tenant-smart-learning",
|
||||
domain="expense",
|
||||
natural_language="重复发票风险规则",
|
||||
)
|
||||
|
||||
fake_retriever.retrieve_for_risk_rule_generation.assert_called_once_with(
|
||||
tenant_id="tenant-smart-learning",
|
||||
domain="expense",
|
||||
natural_language="重复发票风险规则",
|
||||
)
|
||||
|
||||
|
||||
def test_rule_generation_without_tenant_disables_historical_injection(monkeypatch) -> None:
|
||||
from_session = MagicMock()
|
||||
monkeypatch.setenv("FEW_SHOT_INJECTION_ENABLED", "true")
|
||||
monkeypatch.setattr(FewShotRetriever, "from_session", from_session)
|
||||
|
||||
result = RiskRuleGenerationService(MagicMock())._retrieve_few_shot_samples(
|
||||
tenant_id=None,
|
||||
domain="expense",
|
||||
natural_language="重复发票风险规则",
|
||||
)
|
||||
|
||||
assert result == []
|
||||
from_session.assert_not_called()
|
||||
|
||||
|
||||
def test_build_prompt_merges_few_shot_into_examples() -> None:
|
||||
samples = [
|
||||
{
|
||||
@@ -89,7 +138,14 @@ def test_build_prompt_merges_few_shot_into_examples() -> None:
|
||||
expense_category=None,
|
||||
expense_category_label="",
|
||||
natural_language="重复发票规则",
|
||||
available_fields=[{"key": "attachment.invoice_no", "label": "发票号", "type": "string", "source": "attachment"}],
|
||||
available_fields=[
|
||||
{
|
||||
"key": "attachment.invoice_no",
|
||||
"label": "发票号",
|
||||
"type": "string",
|
||||
"source": "attachment",
|
||||
}
|
||||
],
|
||||
few_shot_samples=samples,
|
||||
)
|
||||
assert len(messages) == 2
|
||||
|
||||
Reference in New Issue
Block a user