test_agent_api.py 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340
  1. import re
  2. from datetime import UTC, datetime
  3. from fastapi.testclient import TestClient
  4. from zbt.core.config import Settings
  5. from zbt.core.passwords import PasswordService
  6. from zbt.domains.agent.repository import InMemoryAgentRepository
  7. from zbt.domains.agent.runtime import AgentReply
  8. from zbt.domains.catalog.repository import InMemoryCatalogRepository
  9. from zbt.domains.enrollment.repository import InMemoryEnrollmentRepository
  10. from zbt.domains.identity.models import AdminUser
  11. from zbt.domains.identity.repository import InMemoryIdentityRepository
  12. from zbt.harness.kernel import AgentInvocation
  13. from zbt.infrastructure.redis.agent_state import InMemoryAgentStateStore
  14. from zbt.main import create_app
  15. class FixedAgentRuntime:
  16. def reply(self, invocation: AgentInvocation) -> AgentReply:
  17. assert "65岁" in invocation.message
  18. assert invocation.persona == "customer"
  19. return AgentReply(
  20. text="可以优先考虑银龄守护医疗险,我会继续核对地区和职业。",
  21. cards=[],
  22. trace_id="trace_test_001",
  23. )
  24. class FixedOperationRuntime:
  25. def reply(self, invocation: AgentInvocation) -> AgentReply:
  26. assert invocation.persona == "operation"
  27. assert invocation.admin_user is not None
  28. assert "订单" in invocation.message
  29. return AgentReply(
  30. text="当前共有24笔订单,其中18笔已出单。",
  31. cards=[
  32. {
  33. "type": "metrics",
  34. "version": "1.0",
  35. "title": "经营总览",
  36. "items": [
  37. {"label": "订单总量", "value": 24, "unit": "笔"},
  38. ],
  39. }
  40. ],
  41. trace_id="00000000-0000-0000-0000-000000000001",
  42. trace_url="https://smith.langchain.com/example-trace",
  43. invoked_tools=("get_operation_overview",),
  44. )
  45. class StreamingAgentRuntime:
  46. def reply(self, invocation: AgentInvocation) -> AgentReply:
  47. assert invocation.h5_user is not None
  48. return AgentReply(
  49. text="已为你找到银龄守护医疗险。",
  50. cards=[
  51. {
  52. "type": "product_recommendations",
  53. "version": "1.0",
  54. "title": "为你匹配的保障方案",
  55. "items": [],
  56. }
  57. ],
  58. trace_id="00000000-0000-0000-0000-000000000002",
  59. actions=[],
  60. invoked_tools=("list_available_products",),
  61. )
  62. def test_h5_user_can_create_customer_agent_thread() -> None:
  63. settings = Settings(
  64. app_env="test",
  65. jwt_access_secret="a" * 32,
  66. jwt_refresh_secret="b" * 32,
  67. field_encryption_key="c" * 32,
  68. )
  69. app = create_app(
  70. settings=settings,
  71. identity_repository=InMemoryIdentityRepository(),
  72. agent_repository=InMemoryAgentRepository(),
  73. )
  74. with TestClient(app) as client:
  75. token = client.post(
  76. "/api/v1/h5/auth/login",
  77. json={"mobile": "18800000001", "code": "147258"},
  78. ).json()["data"]["tokens"]["access_token"]
  79. response = client.post(
  80. "/api/v1/agent/threads",
  81. headers={"Authorization": f"Bearer {token}"},
  82. json={"title": "给父亲配置医疗险"},
  83. )
  84. assert response.status_code == 201
  85. assert response.json()["data"]["persona"] == "customer"
  86. assert response.json()["data"]["title"] == "给父亲配置医疗险"
  87. def test_h5_user_can_send_message_and_receive_agent_reply() -> None:
  88. settings = Settings(
  89. app_env="test",
  90. jwt_access_secret="a" * 32,
  91. jwt_refresh_secret="b" * 32,
  92. field_encryption_key="c" * 32,
  93. )
  94. app = create_app(
  95. settings=settings,
  96. identity_repository=InMemoryIdentityRepository(),
  97. agent_repository=InMemoryAgentRepository(),
  98. agent_runtime=FixedAgentRuntime(),
  99. )
  100. with TestClient(app) as client:
  101. token = client.post(
  102. "/api/v1/h5/auth/login",
  103. json={"mobile": "18800000001", "code": "147258"},
  104. ).json()["data"]["tokens"]["access_token"]
  105. headers = {"Authorization": f"Bearer {token}"}
  106. thread_id = client.post(
  107. "/api/v1/agent/threads",
  108. headers=headers,
  109. json={"title": "给父亲配置医疗险"},
  110. ).json()["data"]["id"]
  111. response = client.post(
  112. f"/api/v1/agent/threads/{thread_id}/messages",
  113. headers=headers,
  114. json={
  115. "content": {
  116. "type": "text",
  117. "text": "想给65岁的父亲买医疗险",
  118. },
  119. "client_message_id": "client-message-001",
  120. },
  121. )
  122. assert response.status_code == 202
  123. assert response.json()["data"]["status"] == "COMPLETED"
  124. assert "银龄守护医疗险" in response.json()["data"]["assistant_message"]["text"]
  125. assert re.fullmatch(
  126. r"ZBT-\d{8}-\d{4}",
  127. response.json()["data"]["service_no"],
  128. )
  129. assert (
  130. response.json()["data"]["assistant_message"]["service_no"]
  131. == response.json()["data"]["service_no"]
  132. )
  133. assert response.json()["data"]["trace_id"] == "trace_test_001"
  134. run_id = response.json()["data"]["run_id"]
  135. with TestClient(app) as client:
  136. token = client.post(
  137. "/api/v1/h5/auth/login",
  138. json={"mobile": "18800000001", "code": "147258"},
  139. ).json()["data"]["tokens"]["access_token"]
  140. headers = {"Authorization": f"Bearer {token}"}
  141. messages = client.get(
  142. f"/api/v1/agent/threads/{thread_id}/messages",
  143. headers=headers,
  144. )
  145. run = client.get(f"/api/v1/agent/runs/{run_id}", headers=headers)
  146. assert [item["role"] for item in messages.json()["data"]["items"]] == [
  147. "USER",
  148. "ASSISTANT",
  149. ]
  150. assert run.json()["data"]["status"] == "COMPLETED"
  151. assert run.json()["data"]["trace_id"] == "trace_test_001"
  152. def test_admin_can_use_operation_persona_with_the_same_agent_kernel_contract() -> None:
  153. settings = Settings(
  154. app_env="test",
  155. jwt_access_secret="a" * 32,
  156. jwt_refresh_secret="b" * 32,
  157. field_encryption_key="c" * 32,
  158. )
  159. identities = InMemoryIdentityRepository()
  160. identities.save_admin_user(
  161. AdminUser(
  162. id="01ADMIN0000000000000000001",
  163. username="admin",
  164. password_hash=PasswordService().hash("zaq1XSW@"),
  165. display_name="系统管理员",
  166. status="ACTIVE",
  167. roles=("SUPER_ADMIN",),
  168. permissions=("*",),
  169. data_scope="ALL",
  170. )
  171. )
  172. app = create_app(
  173. settings=settings,
  174. identity_repository=identities,
  175. catalog_repository=InMemoryCatalogRepository(),
  176. agent_repository=InMemoryAgentRepository(),
  177. agent_runtime=FixedOperationRuntime(),
  178. enrollment_repository=InMemoryEnrollmentRepository(),
  179. clock=lambda: datetime(2026, 7, 26, 8, tzinfo=UTC),
  180. )
  181. with TestClient(app) as client:
  182. token = client.post(
  183. "/api/v1/admin/auth/login",
  184. json={"username": "admin", "password": "zaq1XSW@"},
  185. ).json()["data"]["tokens"]["access_token"]
  186. headers = {"Authorization": f"Bearer {token}"}
  187. thread = client.post(
  188. "/api/v1/admin/agent/threads",
  189. headers=headers,
  190. json={"title": "订单经营分析"},
  191. )
  192. response = client.post(
  193. f"/api/v1/admin/agent/threads/{thread.json()['data']['id']}/messages",
  194. headers=headers,
  195. json={
  196. "content": {"type": "text", "text": "分析一下当前订单情况"},
  197. "client_message_id": "operation-message-001",
  198. },
  199. )
  200. runs = client.get("/api/v1/admin/agent/runs", headers=headers)
  201. assert thread.status_code == 201
  202. assert thread.json()["data"]["persona"] == "operation"
  203. assert response.status_code == 202
  204. data = response.json()["data"]
  205. assert data["invoked_tools"] == ["get_operation_overview"]
  206. assert data["trace_url"] == "https://smith.langchain.com/example-trace"
  207. assert data["assistant_message"]["cards"][0]["type"] == "metrics"
  208. assert runs.status_code == 200
  209. assert runs.json()["data"]["total"] == 1
  210. assert runs.json()["data"]["items"][0]["id"] == data["run_id"]
  211. def test_h5_agent_stream_emits_and_replays_runtime_events() -> None:
  212. settings = Settings(
  213. app_env="test",
  214. jwt_access_secret="a" * 32,
  215. jwt_refresh_secret="b" * 32,
  216. field_encryption_key="c" * 32,
  217. )
  218. state_store = InMemoryAgentStateStore()
  219. app = create_app(
  220. settings=settings,
  221. identity_repository=InMemoryIdentityRepository(),
  222. catalog_repository=InMemoryCatalogRepository(),
  223. agent_repository=InMemoryAgentRepository(),
  224. agent_runtime=StreamingAgentRuntime(),
  225. enrollment_repository=InMemoryEnrollmentRepository(),
  226. agent_state_store=state_store,
  227. )
  228. with TestClient(app) as client:
  229. token = client.post(
  230. "/api/v1/h5/auth/login",
  231. json={"mobile": "18800000001", "code": "147258"},
  232. ).json()["data"]["tokens"]["access_token"]
  233. headers = {"Authorization": f"Bearer {token}"}
  234. thread_id = client.post(
  235. "/api/v1/agent/threads",
  236. headers=headers,
  237. json={"title": "SSE测试"},
  238. ).json()["data"]["id"]
  239. response = client.post(
  240. f"/api/v1/agent/threads/{thread_id}/messages/stream",
  241. headers=headers,
  242. json={
  243. "content": {"type": "text", "text": "推荐父母医疗险"},
  244. "client_message_id": "stream-message-001",
  245. },
  246. )
  247. assert response.status_code == 200
  248. assert response.headers["content-type"].startswith("text/event-stream")
  249. assert response.text.index("event: run.started") < response.text.index(
  250. "event: tool.completed"
  251. )
  252. assert response.text.index("event: tool.completed") < response.text.index("event: ui.ready")
  253. assert response.text.index("event: ui.ready") < response.text.index("event: run.completed")
  254. run_id = next(
  255. line.removeprefix("id: ").split(":")[0]
  256. for line in response.text.splitlines()
  257. if line.startswith("id: ")
  258. )
  259. replay = client.get(
  260. f"/api/v1/agent/runs/{run_id}/events",
  261. headers=headers,
  262. )
  263. assert replay.status_code == 200
  264. assert replay.text.count("event: ") == 4
  265. assert "list_available_products" in replay.text
  266. def test_agent_rate_limit_is_enforced_before_model_invocation() -> None:
  267. settings = Settings(
  268. app_env="test",
  269. jwt_access_secret="a" * 32,
  270. jwt_refresh_secret="b" * 32,
  271. field_encryption_key="c" * 32,
  272. )
  273. app = create_app(
  274. settings=settings,
  275. identity_repository=InMemoryIdentityRepository(),
  276. catalog_repository=InMemoryCatalogRepository(),
  277. agent_repository=InMemoryAgentRepository(),
  278. agent_runtime=StreamingAgentRuntime(),
  279. enrollment_repository=InMemoryEnrollmentRepository(),
  280. agent_state_store=InMemoryAgentStateStore(rate_limit=1),
  281. )
  282. with TestClient(app) as client:
  283. token = client.post(
  284. "/api/v1/h5/auth/login",
  285. json={"mobile": "18800000001", "code": "147258"},
  286. ).json()["data"]["tokens"]["access_token"]
  287. headers = {"Authorization": f"Bearer {token}"}
  288. thread_id = client.post(
  289. "/api/v1/agent/threads",
  290. headers=headers,
  291. json={"title": "限流测试"},
  292. ).json()["data"]["id"]
  293. payload = {
  294. "content": {"type": "text", "text": "推荐医疗险"},
  295. "client_message_id": "rate-message",
  296. }
  297. first = client.post(
  298. f"/api/v1/agent/threads/{thread_id}/messages",
  299. headers=headers,
  300. json=payload,
  301. )
  302. second = client.post(
  303. f"/api/v1/agent/threads/{thread_id}/messages",
  304. headers=headers,
  305. json=payload,
  306. )
  307. assert first.status_code == 202
  308. assert second.status_code == 429
  309. assert second.json()["error"]["code"] == "AGENT_RATE_LIMITED"