from __future__ import annotations from app.config import Settings from app.decision_engine import create_decision_engine from app.embeddings import create_embedding_provider from app.graph import GraphDependencies, build_graph from app.milvus_store import MilvusVectorStore from app.schemas import QueryRequest, QueryResponse, RouteDecision from app.sql_store import OrderRepository from app.tools import ( SQLQueryTool, TavilyWebSearchProvider, VectorSearchTool, VectorStore, WebSearchProvider, WebSearchTool, ) class AgenticRAGService: def __init__( self, settings: Settings, vector_store: VectorStore | None = None, web_search_provider: WebSearchProvider | None = None, ) -> None: self.settings = settings # 测试可注入 Fake Store;生产路径才创建真实 Milvus 客户端。 if vector_store is None: embeddings = create_embedding_provider(settings) vector_store = MilvusVectorStore( uri=settings.milvus_uri, token=settings.milvus_token.get_secret_value(), collection_name=settings.milvus_collection, embeddings=embeddings, ) if web_search_provider is None: # Tavily Provider 延迟到真正调用 web_search 时才校验 Key 和依赖。 web_search_provider = TavilyWebSearchProvider( settings.tavily_api_key.get_secret_value() ) dependencies = GraphDependencies( settings=settings, engine=create_decision_engine(settings), vector_tool=VectorSearchTool(vector_store), sql_tool=SQLQueryTool(OrderRepository(settings.sqlite_path)), web_tool=WebSearchTool(web_search_provider), ) self.graph = build_graph(dependencies) def invoke(self, request: QueryRequest) -> QueryResponse: # 每次请求创建独立初始状态,避免跨会话共享 Evidence 或错误信息。 state = self.graph.invoke( { "original_query": request.query, "current_query": request.query, "session_id": request.session_id, "debug": request.debug, "evidence": [], "tool_results": [], "executed_queries": [], "retrieval_round": 0, "errors": [], "trace": [], }, config={"recursion_limit": 20}, ) route = state.get("route_decision") if route is None: route = RouteDecision( needs_retrieval=False, intent="internal_error", reason_code="MISSING_ROUTE_DECISION", ) return QueryResponse( answer=state.get("final_answer", "系统未生成答案。"), citations=state.get("evidence", []), route=route, executed_queries=state.get("executed_queries", []) if request.debug else [], trace=state.get("trace", []) if request.debug else [], termination_reason=state.get("termination_reason", "unknown"), )