service.py 3.1 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182
  1. from __future__ import annotations
  2. from app.config import Settings
  3. from app.decision_engine import create_decision_engine
  4. from app.embeddings import create_embedding_provider
  5. from app.graph import GraphDependencies, build_graph
  6. from app.milvus_store import MilvusVectorStore
  7. from app.schemas import QueryRequest, QueryResponse, RouteDecision
  8. from app.sql_store import OrderRepository
  9. from app.tools import (
  10. SQLQueryTool,
  11. TavilyWebSearchProvider,
  12. VectorSearchTool,
  13. VectorStore,
  14. WebSearchProvider,
  15. WebSearchTool,
  16. )
  17. class AgenticRAGService:
  18. def __init__(
  19. self,
  20. settings: Settings,
  21. vector_store: VectorStore | None = None,
  22. web_search_provider: WebSearchProvider | None = None,
  23. ) -> None:
  24. self.settings = settings
  25. # 测试可注入 Fake Store;生产路径才创建真实 Milvus 客户端。
  26. if vector_store is None:
  27. embeddings = create_embedding_provider(settings)
  28. vector_store = MilvusVectorStore(
  29. uri=settings.milvus_uri,
  30. token=settings.milvus_token.get_secret_value(),
  31. collection_name=settings.milvus_collection,
  32. embeddings=embeddings,
  33. )
  34. if web_search_provider is None:
  35. # Tavily Provider 延迟到真正调用 web_search 时才校验 Key 和依赖。
  36. web_search_provider = TavilyWebSearchProvider(
  37. settings.tavily_api_key.get_secret_value()
  38. )
  39. dependencies = GraphDependencies(
  40. settings=settings,
  41. engine=create_decision_engine(settings),
  42. vector_tool=VectorSearchTool(vector_store),
  43. sql_tool=SQLQueryTool(OrderRepository(settings.sqlite_path)),
  44. web_tool=WebSearchTool(web_search_provider),
  45. )
  46. self.graph = build_graph(dependencies)
  47. def invoke(self, request: QueryRequest) -> QueryResponse:
  48. # 每次请求创建独立初始状态,避免跨会话共享 Evidence 或错误信息。
  49. state = self.graph.invoke(
  50. {
  51. "original_query": request.query,
  52. "current_query": request.query,
  53. "session_id": request.session_id,
  54. "debug": request.debug,
  55. "evidence": [],
  56. "tool_results": [],
  57. "executed_queries": [],
  58. "retrieval_round": 0,
  59. "errors": [],
  60. "trace": [],
  61. },
  62. config={"recursion_limit": 20},
  63. )
  64. route = state.get("route_decision")
  65. if route is None:
  66. route = RouteDecision(
  67. needs_retrieval=False,
  68. intent="internal_error",
  69. reason_code="MISSING_ROUTE_DECISION",
  70. )
  71. return QueryResponse(
  72. answer=state.get("final_answer", "系统未生成答案。"),
  73. citations=state.get("evidence", []),
  74. route=route,
  75. executed_queries=state.get("executed_queries", []) if request.debug else [],
  76. trace=state.get("trace", []) if request.debug else [],
  77. termination_reason=state.get("termination_reason", "unknown"),
  78. )