172 lines
5.8 KiB
Python
172 lines
5.8 KiB
Python
import pytest
|
|
|
|
from app.core.contracts import AgentRequest, AgentRequestMetadata, RequestContext
|
|
from app.core.errors import RecoverableAgentError
|
|
from app.core.knowledge_contracts import KnowledgeHit, KnowledgeSearchResult
|
|
from app.service.agent.base import BaseAgent
|
|
from app.service.agent.bootstrap import get_agent_factory
|
|
from app.service.agent.customer_service_agent import CustomerServiceAgent
|
|
from app.service.agent.customer_service_routing import CustomerServiceIntentRouter
|
|
|
|
|
|
def request(message: str, *, chitchat_streak: int = 0) -> AgentRequest:
|
|
return AgentRequest(
|
|
agent_type="customer_service",
|
|
message=message,
|
|
session_id="customer-service-session",
|
|
idempotency_key="customer-service-idempotency-key",
|
|
metadata=AgentRequestMetadata(chitchat_streak=chitchat_streak),
|
|
)
|
|
|
|
|
|
def context(role: str) -> RequestContext:
|
|
return RequestContext(
|
|
user_id="1",
|
|
trace_id="customer-service-trace",
|
|
roles=(role,),
|
|
permissions=("agent:run", "knowledge:query"),
|
|
data_scope="public" if role == "visitor" else "self",
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_visitor_account_question_only_returns_login_entry() -> None:
|
|
result = await CustomerServiceAgent().handle(
|
|
request("我的持仓收益是多少"), context("visitor")
|
|
)
|
|
|
|
assert result.text == "我无法查询账户数据,请先登录后前往“我的账户”查看相关状态。"
|
|
assert result.transfer_required is False
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_authenticated_account_question_only_returns_account_entry() -> None:
|
|
result = await CustomerServiceAgent().handle(request("查一下我的订单"), context("customer"))
|
|
|
|
assert result.text == "我无法查询账户数据,请前往“我的账户”查看相关状态。"
|
|
assert result.transfer_required is False
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_security_question_requires_human_transfer() -> None:
|
|
result = await CustomerServiceAgent().handle(
|
|
request("验证码已经发给别人了"), context("visitor")
|
|
)
|
|
|
|
assert result.transfer_required is True
|
|
assert result.transfer_reason == "security_notice"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_personalized_investment_advice_is_refused_and_transferred() -> None:
|
|
result = await CustomerServiceAgent().handle(
|
|
request("帮我推荐一只收益最高的基金"), context("customer")
|
|
)
|
|
|
|
assert result.transfer_required is True
|
|
assert result.transfer_reason == "compliance_refusal"
|
|
|
|
|
|
def test_customer_service_agent_is_registered_for_visitor_role() -> None:
|
|
agent = get_agent_factory().create("customer_service", context("visitor"))
|
|
|
|
assert isinstance(agent, CustomerServiceAgent)
|
|
|
|
|
|
def test_chitchat_streak_is_derived_from_continuous_prior_messages() -> None:
|
|
assert CustomerServiceIntentRouter.chitchat_streak(
|
|
("你好", "讲个笑话", "你开心吗"), "在吗"
|
|
) == 4
|
|
assert CustomerServiceIntentRouter.chitchat_streak(
|
|
("你好", "基金怎么开户"), "在吗"
|
|
) == 1
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_policy_question_uses_only_policy_knowledge(monkeypatch: pytest.MonkeyPatch) -> None:
|
|
captured: dict[str, object] = {}
|
|
|
|
async def fake_call_tool(
|
|
self: BaseAgent,
|
|
name: str,
|
|
arguments: dict[str, object],
|
|
*,
|
|
intent: str,
|
|
context: RequestContext,
|
|
) -> KnowledgeSearchResult:
|
|
captured.update(name=name, arguments=arguments, intent=intent)
|
|
return KnowledgeSearchResult(
|
|
hits=(
|
|
KnowledgeHit(
|
|
knowledge_id="101",
|
|
collection="fin_policy_collection",
|
|
snippet="赎回规则摘要",
|
|
answer="这是已审核的赎回公开规则。",
|
|
),
|
|
)
|
|
)
|
|
|
|
monkeypatch.setattr(BaseAgent, "call_tool", fake_call_tool)
|
|
|
|
result = await CustomerServiceAgent().handle(
|
|
request("基金赎回到账规则是什么"), context("customer")
|
|
)
|
|
|
|
assert result.text == "这是已审核的赎回公开规则。"
|
|
assert captured == {
|
|
"name": "query_knowledge",
|
|
"arguments": {
|
|
"query": "基金赎回到账规则是什么",
|
|
"intents": ("policy_explain",),
|
|
"top_k": 5,
|
|
},
|
|
"intent": "public_knowledge",
|
|
}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_knowledge_failure_requires_human_transfer(monkeypatch: pytest.MonkeyPatch) -> None:
|
|
async def unavailable_tool(
|
|
self: BaseAgent,
|
|
name: str,
|
|
arguments: dict[str, object],
|
|
*,
|
|
intent: str,
|
|
context: RequestContext,
|
|
) -> KnowledgeSearchResult:
|
|
raise RecoverableAgentError("知识检索不可用")
|
|
|
|
monkeypatch.setattr(BaseAgent, "call_tool", unavailable_tool)
|
|
|
|
result = await CustomerServiceAgent().handle(request("基金怎么开户"), context("visitor"))
|
|
|
|
assert result.transfer_required is True
|
|
assert result.transfer_reason == "knowledge_unavailable"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_fourth_chitchat_message_is_guided_once_without_knowledge_lookup(
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
) -> None:
|
|
async def unexpected_tool(
|
|
self: BaseAgent,
|
|
name: str,
|
|
arguments: dict[str, object],
|
|
*,
|
|
intent: str,
|
|
context: RequestContext,
|
|
) -> KnowledgeSearchResult:
|
|
raise AssertionError("闲聊不应调用公开知识工具")
|
|
|
|
monkeypatch.setattr(BaseAgent, "call_tool", unexpected_tool)
|
|
|
|
guided = await CustomerServiceAgent().handle(
|
|
request("你今天开心吗", chitchat_streak=4), context("visitor")
|
|
)
|
|
ordinary = await CustomerServiceAgent().handle(
|
|
request("你今天开心吗", chitchat_streak=5), context("visitor")
|
|
)
|
|
|
|
assert "基金业务" in guided.text
|
|
assert "基金业务" not in ordinary.text
|