test_agent_knowledge_citations.py 7.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229
  1. import json
  2. from datetime import UTC, datetime
  3. from typing import Any
  4. from langchain_core.language_models.chat_models import BaseChatModel
  5. from langchain_core.messages import AIMessage
  6. from langchain_core.outputs import ChatGeneration, ChatResult
  7. from pydantic import Field
  8. from zbt.core.config import Settings
  9. from zbt.domains.agent.tools import build_agent_tool_registry
  10. from zbt.domains.catalog.repository import InMemoryCatalogRepository
  11. from zbt.domains.catalog.service import ProductCatalogService
  12. from zbt.domains.enrollment.repository import InMemoryEnrollmentRepository
  13. from zbt.domains.enrollment.service import EnrollmentService
  14. from zbt.domains.identity.models import H5User
  15. from zbt.domains.knowledge.index import InMemoryKnowledgeIndex
  16. from zbt.domains.knowledge.repository import InMemoryKnowledgeRepository
  17. from zbt.domains.knowledge.service import KnowledgeService
  18. from zbt.harness.kernel import AgentInvocation, AgentKernel
  19. NOW = datetime(2026, 7, 28, 10, tzinfo=UTC)
  20. class CitationAwareTestModel(BaseChatModel):
  21. cited_chunk_ids: list[str] = Field(default_factory=list)
  22. invocation_count: int = 0
  23. structured_tool_name: str | None = None
  24. @property
  25. def _llm_type(self) -> str:
  26. return "citation-aware-test-model"
  27. def bind_tools(
  28. self,
  29. tools: Any,
  30. *,
  31. tool_choice: str | None = None,
  32. **kwargs: Any,
  33. ) -> Any:
  34. del tool_choice, kwargs
  35. names = [_tool_name(tool) for tool in tools]
  36. self.structured_tool_name = next(
  37. (name for name in names if name == "AgentStructuredOutput"),
  38. None,
  39. )
  40. return self
  41. def _generate(
  42. self,
  43. messages: Any,
  44. stop: list[str] | None = None,
  45. run_manager: Any = None,
  46. **kwargs: Any,
  47. ) -> ChatResult:
  48. del messages, stop, run_manager, kwargs
  49. self.invocation_count += 1
  50. if self.invocation_count == 1:
  51. message = AIMessage(
  52. content="",
  53. tool_calls=[
  54. {
  55. "name": "search_insurance_knowledge",
  56. "args": {
  57. "query": "安心医疗险是否报销海外整形手术,比例是多少?",
  58. "limit": 5,
  59. },
  60. "id": "knowledge-search-1",
  61. }
  62. ],
  63. )
  64. elif self.structured_tool_name is not None:
  65. message = AIMessage(
  66. content="",
  67. tool_calls=[
  68. {
  69. "name": self.structured_tool_name,
  70. "args": {
  71. "message": (
  72. "当前知识库未明确海外整形手术是否属于保障责任,"
  73. "也未明确具体报销比例,无法确定。"
  74. ),
  75. "selected_product_codes": [],
  76. "suggested_action": None,
  77. "cited_knowledge_chunk_ids": self.cited_chunk_ids,
  78. },
  79. "id": "structured-answer-1",
  80. }
  81. ],
  82. )
  83. else:
  84. marker = json.dumps(self.cited_chunk_ids)
  85. message = AIMessage(
  86. content=(
  87. "当前知识库未明确海外整形手术是否属于保障责任,"
  88. "也未明确具体报销比例,无法确定。\n\n"
  89. f"<!-- ZBT_CITATIONS: {marker} -->"
  90. )
  91. )
  92. return ChatResult(generations=[ChatGeneration(message=message)])
  93. class CitationAwareTestGateway:
  94. def __init__(self, model: CitationAwareTestModel) -> None:
  95. self._model = model
  96. def chat_model(self, persona: str) -> BaseChatModel:
  97. assert persona == "customer"
  98. return self._model
  99. def test_unknown_answer_does_not_expose_retrieval_candidates_as_citations() -> None:
  100. kernel, user, _, _ = _kernel(cited_chunk_ids=[])
  101. reply = kernel.reply(
  102. AgentInvocation(
  103. persona="customer",
  104. message="安心医疗险是否可以报销海外整形手术?具体报销比例是多少?",
  105. h5_user=user,
  106. )
  107. )
  108. assert reply.invoked_tools == ("search_insurance_knowledge",)
  109. assert "无法确定" in reply.text
  110. assert "ZBT_CITATIONS" not in reply.text
  111. assert not any(card["type"] == "knowledge_sources" for card in reply.cards)
  112. def test_answer_only_exposes_selected_candidates_as_citations() -> None:
  113. kernel, user, chunk_ids, model = _kernel(cited_chunk_ids=[])
  114. model.cited_chunk_ids = [
  115. chunk_ids[0],
  116. chunk_ids[1],
  117. "not-returned-by-the-knowledge-tool",
  118. chunk_ids[0],
  119. ]
  120. reply = kernel.reply(
  121. AgentInvocation(
  122. persona="customer",
  123. message="非医学必需的美容项目是否属于保障责任?",
  124. h5_user=user,
  125. )
  126. )
  127. source_cards = [
  128. card for card in reply.cards if card["type"] == "knowledge_sources"
  129. ]
  130. assert len(source_cards) == 1
  131. assert len(source_cards[0]["items"]) == 1
  132. source = source_cards[0]["items"][0]
  133. assert source["excerpts"] == [
  134. "非医学必需的美容项目不属于本产品保障责任。",
  135. "具体报销比例以电子保单对应的保障计划为准。",
  136. ]
  137. def _kernel(
  138. *,
  139. cited_chunk_ids: list[str],
  140. ) -> tuple[AgentKernel, H5User, list[str], CitationAwareTestModel]:
  141. knowledge = KnowledgeService(
  142. InMemoryKnowledgeRepository(),
  143. lambda: NOW,
  144. InMemoryKnowledgeIndex(),
  145. )
  146. document = knowledge.create_document(
  147. title="安心医疗险保险条款",
  148. document_type="INSURANCE_TERMS",
  149. source_name="安心医疗险保险条款.md",
  150. product_code="MED-BASIC",
  151. content=(
  152. "非医学必需的美容项目不属于本产品保障责任。\n\n"
  153. "具体报销比例以电子保单对应的保障计划为准。"
  154. ),
  155. created_by="admin-01",
  156. )
  157. knowledge.index_document(document["document_id"])
  158. knowledge.test_search(document["document_id"], "美容项目是否保障")
  159. knowledge.publish_document(document["document_id"])
  160. chunk_ids = [
  161. f"{document['document_id']}:1:1",
  162. f"{document['document_id']}:1:2",
  163. ]
  164. catalog = ProductCatalogService(InMemoryCatalogRepository(), lambda: NOW)
  165. enrollment = EnrollmentService(
  166. InMemoryEnrollmentRepository(),
  167. catalog,
  168. lambda: NOW,
  169. "x" * 32,
  170. )
  171. tools = build_agent_tool_registry(catalog, enrollment, knowledge=knowledge)
  172. user = H5User(
  173. id="01H5USER000000000000000001",
  174. mobile="18800000001",
  175. mobile_masked="188****0001",
  176. display_name="测试用户",
  177. status="ACTIVE",
  178. created_at=NOW,
  179. )
  180. settings = Settings(
  181. app_env="test",
  182. jwt_access_secret="a" * 32,
  183. jwt_refresh_secret="b" * 32,
  184. field_encryption_key="c" * 32,
  185. langsmith_tracing=False,
  186. )
  187. model = CitationAwareTestModel(cited_chunk_ids=cited_chunk_ids)
  188. return (
  189. AgentKernel(
  190. settings=settings,
  191. model_gateway=CitationAwareTestGateway(model),
  192. tools=tools,
  193. ),
  194. user,
  195. chunk_ids,
  196. model,
  197. )
  198. def _tool_name(tool: Any) -> str:
  199. if hasattr(tool, "name"):
  200. return str(tool.name)
  201. if isinstance(tool, dict):
  202. function = tool.get("function", {})
  203. if isinstance(function, dict):
  204. return str(function.get("name", ""))
  205. return ""