| 12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182 |
- 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"),
- )
|