Files
Mutual_Fund/tests/test_customer_agent_router.py
T

100 lines
3.9 KiB
Python

import json
import unittest
from types import SimpleNamespace
import httpx
from fastapi import FastAPI
from api.routers.customer_agent import router
from rag.intent import Intent
from service.customer_agent.chat import AnonymousCustomerAgent
from agent.customer_agent.session import AnonymousSessionService
from tests.test_agent_session import FakeRedis
class RouterTests(unittest.IsolatedAsyncioTestCase):
async def asyncSetUp(self):
redis = FakeRedis()
config = {
"agent.customer.session.ttl": "120",
"agent.customer.rate_limit.window_sec": "60",
"agent.customer.rate_limit.max_requests": "1",
"agent.customer.template.fallback_human": "请转人工客服",
}
session = AnonymousSessionService(redis, config_getter=config.get)
class Context:
async def append(self, *args):
pass
async def get(self, *args):
return []
agent = AnonymousCustomerAgent(
context=Context(),
rag_retrieve=lambda query, customer_id: [],
intent_recognize=lambda query: Intent.NO_MATCH,
generate_answer=lambda messages: "answer",
audit_writer=lambda **kwargs: None,
config_getter=config.get,
)
self.app = FastAPI()
self.app.state.customer_agent_runtime = SimpleNamespace(
redis=redis, session_service=session, agent=agent
)
self.app.include_router(router, prefix="/api/agent/customer")
async def test_create_chat_and_end_anonymous_session(self):
transport = httpx.ASGITransport(app=self.app)
async with httpx.AsyncClient(transport=transport, base_url="http://test") as client:
created = await client.post("/api/agent/customer/session/create")
self.assertEqual(created.status_code, 200)
session_id = created.json()["data"]["session_id"]
chat = await client.get(
"/api/agent/customer/chat",
params={"session_id": session_id, "query": "你好"},
headers={"X-Trace-Id": "trace-router"},
)
self.assertEqual(chat.status_code, 200)
event = json.loads(chat.text.removeprefix("data: ").strip())
self.assertEqual(event["trace_id"], "trace-router")
limited = await client.get(
"/api/agent/customer/chat",
params={"session_id": session_id, "query": "再次提问"},
)
self.assertEqual(limited.status_code, 429)
self.assertIn(int(limited.headers["Retry-After"]), range(1, 61))
ended = await client.post(
"/api/agent/customer/session/end", json={"session_id": session_id}
)
self.assertEqual(ended.status_code, 200)
async def test_query_over_2000_returns_400(self):
transport = httpx.ASGITransport(app=self.app)
async with httpx.AsyncClient(transport=transport, base_url="http://test") as client:
created = await client.post("/api/agent/customer/session/create")
session_id = created.json()["data"]["session_id"]
response = await client.get(
"/api/agent/customer/chat",
params={"session_id": session_id, "query": "x" * 2001},
)
self.assertEqual(response.status_code, 400)
self.assertEqual(response.json()["code"], 400)
async def test_missing_session_returns_404(self):
transport = httpx.ASGITransport(app=self.app)
async with httpx.AsyncClient(transport=transport, base_url="http://test") as client:
response = await client.get(
"/api/agent/customer/chat",
params={"session_id": "missing", "query": "你好"},
)
self.assertEqual(response.status_code, 404)
self.assertEqual(response.json()["code"], 404)
if __name__ == "__main__":
unittest.main()