tools.py 5.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172
  1. from __future__ import annotations
  2. from time import perf_counter
  3. from typing import Protocol
  4. from app.schemas import Evidence, ToolResult
  5. from app.sql_store import OrderRepository
  6. class VectorStore(Protocol):
  7. def search(
  8. self,
  9. query: str,
  10. top_k: int = 5,
  11. policy_type: str = "",
  12. ) -> list[Evidence]: ...
  13. class WebSearchProvider(Protocol):
  14. def search(self, query: str, max_results: int = 5) -> list[Evidence]: ...
  15. class TavilyWebSearchProvider:
  16. """把 Tavily 搜索结果转换为项目统一的 Evidence。"""
  17. def __init__(self, api_key: str) -> None:
  18. self.api_key = api_key
  19. def search(self, query: str, max_results: int = 5) -> list[Evidence]:
  20. if not self.api_key:
  21. raise RuntimeError("网页搜索未配置 TAVILY_API_KEY")
  22. try:
  23. from tavily import TavilyClient
  24. except ImportError as exc:
  25. raise RuntimeError(
  26. "缺少网页搜索依赖,请执行 uv sync --extra test --extra web"
  27. ) from exc
  28. client = TavilyClient(api_key=self.api_key)
  29. response = client.search(
  30. query=query,
  31. search_depth="basic",
  32. topic="general",
  33. max_results=max_results,
  34. include_answer=False,
  35. include_raw_content=False,
  36. )
  37. # 只保留可引用片段和来源,不把搜索服务生成的 answer 当作事实答案。
  38. evidence: list[Evidence] = []
  39. for item in response.get("results", []):
  40. url = str(item.get("url", "")).strip()
  41. content = str(item.get("content", "")).strip()
  42. if not url or not content:
  43. continue
  44. title = str(item.get("title", "网页结果")).strip()
  45. evidence.append(
  46. Evidence(
  47. source_type="web",
  48. source=url,
  49. content=f"{title}\n{content}",
  50. score=float(item.get("score", 0.0)),
  51. metadata={
  52. "title": title,
  53. "url": url,
  54. "published_date": item.get("published_date"),
  55. },
  56. )
  57. )
  58. return evidence
  59. class VectorSearchTool:
  60. name = "milvus_search"
  61. def __init__(self, store: VectorStore) -> None:
  62. self.store = store
  63. def invoke(self, query: str, policy_type: str = "", top_k: int = 5) -> ToolResult:
  64. started = perf_counter()
  65. try:
  66. # Adapter 负责把数据源异常统一转换为 ToolResult,Graph 不捕获 SDK 异常。
  67. evidence = self.store.search(query, top_k=top_k, policy_type=policy_type)
  68. return ToolResult(
  69. status="success" if evidence else "empty",
  70. tool=self.name,
  71. evidence=evidence,
  72. latency_ms=int((perf_counter() - started) * 1000),
  73. )
  74. except Exception as exc:
  75. return ToolResult(
  76. status="error",
  77. tool=self.name,
  78. latency_ms=int((perf_counter() - started) * 1000),
  79. retryable=True,
  80. error_code="MILVUS_SEARCH_FAILED",
  81. error_message=str(exc),
  82. )
  83. class SQLQueryTool:
  84. name = "sql_query"
  85. def __init__(self, repository: OrderRepository) -> None:
  86. self.repository = repository
  87. def invoke(
  88. self,
  89. user_id: str,
  90. days: int = 30,
  91. product_keyword: str = "",
  92. ) -> ToolResult:
  93. started = perf_counter()
  94. try:
  95. data, evidence = self.repository.summarize_recent(
  96. user_id=user_id,
  97. days=days,
  98. product_keyword=product_keyword,
  99. )
  100. return ToolResult(
  101. status="success",
  102. tool=self.name,
  103. data=[data],
  104. evidence=[evidence],
  105. latency_ms=int((perf_counter() - started) * 1000),
  106. )
  107. except (ValueError, RuntimeError) as exc:
  108. return ToolResult(
  109. status="error",
  110. tool=self.name,
  111. latency_ms=int((perf_counter() - started) * 1000),
  112. retryable=False,
  113. error_code="SQL_QUERY_REJECTED",
  114. error_message=str(exc),
  115. )
  116. class WebSearchTool:
  117. name = "web_search"
  118. def __init__(self, provider: WebSearchProvider) -> None:
  119. self.provider = provider
  120. def invoke(self, query: str, max_results: int = 5) -> ToolResult:
  121. started = perf_counter()
  122. try:
  123. evidence = self.provider.search(query, max_results=max_results)
  124. return ToolResult(
  125. status="success" if evidence else "empty",
  126. tool=self.name,
  127. data=[item.metadata for item in evidence],
  128. evidence=evidence,
  129. latency_ms=int((perf_counter() - started) * 1000),
  130. )
  131. except RuntimeError as exc:
  132. return ToolResult(
  133. status="error",
  134. tool=self.name,
  135. latency_ms=int((perf_counter() - started) * 1000),
  136. retryable=False,
  137. error_code="WEB_SEARCH_NOT_CONFIGURED",
  138. error_message=str(exc),
  139. )
  140. except Exception as exc:
  141. return ToolResult(
  142. status="error",
  143. tool=self.name,
  144. latency_ms=int((perf_counter() - started) * 1000),
  145. retryable=True,
  146. error_code="WEB_SEARCH_FAILED",
  147. error_message=str(exc),
  148. )