service.py 2.9 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182838485
  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 GraphBuilder, build_graph
  6. from app.milvus_store import MilvusVectorStore
  7. from app.schemas import QueryRequest, QueryResponse, RouteDecision
  8. from app.sql_store import IncidentRepository
  9. from app.tools import (
  10. SQLQueryTool,
  11. TavilyWebSearchProvider,
  12. VectorSearchTool,
  13. WebSearchTool,
  14. )
  15. class AgenticRAGService:
  16. """组装 VPN 故障诊断 Agent,并提供单次查询入口。"""
  17. def __init__(self, settings: Settings) -> None:
  18. self.settings = settings
  19. embeddings = create_embedding_provider(settings)
  20. vector_store = MilvusVectorStore(
  21. uri=settings.milvus_uri,
  22. token=settings.milvus_token.get_secret_value(),
  23. collection_name=settings.milvus_collection,
  24. embeddings=embeddings,
  25. )
  26. vector_tool = VectorSearchTool(vector_store)
  27. incident_repository = IncidentRepository(settings.sqlite_path)
  28. sql_tool = SQLQueryTool(incident_repository)
  29. web_provider = TavilyWebSearchProvider(
  30. settings.tavily_api_key.get_secret_value()
  31. )
  32. web_tool = WebSearchTool(web_provider)
  33. graph_builder = GraphBuilder(
  34. settings=settings,
  35. engine=create_decision_engine(settings),
  36. vector_tool=vector_tool,
  37. sql_tool=sql_tool,
  38. web_tool=web_tool,
  39. )
  40. self.graph = build_graph(graph_builder)
  41. def invoke(self, request: QueryRequest) -> QueryResponse:
  42. """为每次请求创建独立 State 并运行 LangGraph。"""
  43. state = self.graph.invoke(
  44. {
  45. "original_query": request.query,
  46. "current_query": request.query,
  47. "session_id": request.session_id,
  48. "debug": request.debug,
  49. "evidence": [],
  50. "tool_results": [],
  51. "executed_queries": [],
  52. "retrieval_round": 0,
  53. "errors": [],
  54. "trace": [],
  55. },
  56. config={"recursion_limit": 20},
  57. )
  58. route = state.get("route_decision")
  59. if route is None:
  60. route = RouteDecision(
  61. needs_retrieval=False,
  62. intent="internal_error",
  63. reason_code="MISSING_ROUTE_DECISION",
  64. )
  65. return QueryResponse(
  66. answer=state.get("final_answer", "系统未生成诊断结论。"),
  67. citations=state.get("evidence", []),
  68. route=route,
  69. executed_queries=(
  70. state.get("executed_queries", []) if request.debug else []
  71. ),
  72. trace=state.get("trace", []) if request.debug else [],
  73. termination_reason=state.get("termination_reason", "unknown"),
  74. )