100 lines
3.9 KiB
Python
100 lines
3.9 KiB
Python
import json
|
|
import unittest
|
|
from types import SimpleNamespace
|
|
|
|
import httpx
|
|
from fastapi import FastAPI
|
|
|
|
from api.chat.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()
|