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 GraphBuilder, build_graph from app.milvus_store import MilvusVectorStore from app.schemas import QueryRequest, QueryResponse, RouteDecision from app.sql_store import IncidentRepository from app.tools import ( SQLQueryTool, TavilyWebSearchProvider, VectorSearchTool, WebSearchTool, ) class AgenticRAGService: """组装 VPN 故障诊断 Agent,并提供单次查询入口。""" def __init__(self, settings: Settings) -> None: self.settings = settings 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, ) vector_tool = VectorSearchTool(vector_store) incident_repository = IncidentRepository(settings.sqlite_path) sql_tool = SQLQueryTool(incident_repository) web_provider = TavilyWebSearchProvider( settings.tavily_api_key.get_secret_value() ) web_tool = WebSearchTool(web_provider) graph_builder = GraphBuilder( settings=settings, engine=create_decision_engine(settings), vector_tool=vector_tool, sql_tool=sql_tool, web_tool=web_tool, ) self.graph = build_graph(graph_builder) def invoke(self, request: QueryRequest) -> QueryResponse: """为每次请求创建独立 State 并运行 LangGraph。""" 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"), )