import asyncio import json import time import uuid from fastapi import APIRouter, HTTPException, Depends from fastapi.responses import StreamingResponse from config import Config from schemas.agent_input import AgentInput from schemas.agent_output import AgentOutput from schemas.tool_input import ToolInput from schemas.tool_output import ToolOutput from schemas.chat_message_response import ChatMessageResponseDTO from schemas.chat_message_request import ChatMessageRequestDTO from schemas.super_agent import SuperAgentRequest, SuperAgentResponse, SuperAgentStreamEvent from workflows.workflow_manager import WorkflowType from api.dependencies import get_workflow_manager, get_nacos_manager, get_service_config, get_tool_router, get_prompt_manager from services.app_errors import AppError, ErrorCode from services.ragflow_sync import RagflowSync from services.structured_logger import get_structured_logger from tools.sr_api_tool import SrApiQueryTool router = APIRouter() def _resolve_workflow_type(value: str) -> WorkflowType: try: return WorkflowType(value) except Exception as e: raise AppError( code=ErrorCode.INVALID_WORKFLOW_TYPE, message=f"不支持的工作流类型: {value}", status_code=400, ) from e def _to_http_error(e: Exception) -> HTTPException: if isinstance(e, AppError): return HTTPException(status_code=e.status_code, detail=e.to_dict()) return HTTPException(status_code=500, detail={"code": ErrorCode.INTERNAL_ERROR.value, "message": str(e)}) @router.get("/health") def health_check(service_config=Depends(get_service_config)): return { "status": "ok", "service_name": service_config.service_name, "model_section": service_config.metadata.get("model_section", "") } @router.get("/nacos/status") def nacos_status(nacos_manager=Depends(get_nacos_manager)): return nacos_manager.status() @router.post("/api/workflows", response_model=AgentOutput) def run_workflow(payload: AgentInput, workflow_manager=Depends(get_workflow_manager)): trace_id = uuid.uuid4().hex slog = get_structured_logger() slog.log("INFO", "run_workflow.start", trace_id, {"workflow_type": payload.workflow_type}) try: workflow_type = _resolve_workflow_type(payload.workflow_type) except Exception as e: slog.log("ERROR", "run_workflow.invalid_type", trace_id, error_code=ErrorCode.INVALID_WORKFLOW_TYPE.value, payload={"workflow_type": payload.workflow_type}) raise _to_http_error(e) try: result = workflow_manager.execute_workflow( workflow_type=workflow_type, user_input=payload.query, session_id=payload.conversation_id, ) slog.log("INFO", "run_workflow.success", trace_id, {"session_id": result.get("session_id")}) except Exception as e: slog.log("ERROR", "run_workflow.failed", trace_id, error_code=ErrorCode.INTERNAL_ERROR.value, payload={"error": str(e)}) raise _to_http_error(e) return AgentOutput( session_id=result["session_id"], workflow_type=result["workflow_type"], result=result["result"], ) @router.post("/api/sql/generate") def generate_sql(payload: AgentInput, workflow_manager=Depends(get_workflow_manager)): """仅生成 SQL,不调用 SR API""" trace_id = uuid.uuid4().hex slog = get_structured_logger() try: workflow_type = _resolve_workflow_type(payload.workflow_type) except Exception as e: slog.log("ERROR", "generate_sql.invalid_type", trace_id, error_code=ErrorCode.INVALID_WORKFLOW_TYPE.value) raise _to_http_error(e) result = workflow_manager.execute_workflow( workflow_type=workflow_type, user_input=payload.query, session_id=payload.conversation_id, skip_sr_api=True, ) context = (result.get("result") or {}).get("context") or {} sql_text = context.get("final_sql") if not sql_text: e = AppError(code=ErrorCode.SQL_GENERATION_FAILED, message="SQL 生成失败") slog.log("ERROR", "generate_sql.failed", trace_id, error_code=e.code.value, payload={"context_keys": list(context.keys())}) raise _to_http_error(e) slog.log("INFO", "generate_sql.success", trace_id, {"sql_len": len(sql_text)}) return { "session_id": result.get("session_id"), "workflow_type": result.get("workflow_type"), "sql": sql_text, } @router.post("/api/workflows/stream") def run_workflow_stream(payload: ChatMessageRequestDTO, workflow_manager=Depends(get_workflow_manager)): trace_id = uuid.uuid4().hex slog = get_structured_logger() if payload.response_mode != "streaming": raise _to_http_error(AppError(code=ErrorCode.INVALID_WORKFLOW_TYPE, message="/api/workflows/stream 仅支持 response_mode=streaming", status_code=400)) stream_cfg = Config.get_section("stream") progress_interval = float(stream_cfg.get("progress_interval", 0.3)) task_id = uuid.uuid4().hex def _build_message(conversation_id: str, answer: str) -> str: dto = ChatMessageResponseDTO( id=uuid.uuid4().hex, event="message", task_id=task_id, message_id=uuid.uuid4().hex, conversation_id=conversation_id, answer=answer, created_at=int(time.time()), ) return f"event: message\ndata: {json.dumps(dto.model_dump(), ensure_ascii=False)}\n\n" async def event_stream(): try: slog.log("INFO", "stream.start", trace_id, {"workflow_type": WorkflowType.CONVERSATION.value}) # 1) 先仅生成 SQL(不执行 SR API) result = await asyncio.to_thread( workflow_manager.execute_workflow, WorkflowType.CONVERSATION, payload.query, payload.conversation_id, skip_sr_api=True, ) conversation_id = str(result.get("session_id") or payload.conversation_id or task_id) context = (result.get("result") or {}).get("context") or {} sql_text = str(context.get("final_sql") or "") if not sql_text: reason = "SQL 生成失败,可能是表未匹配或对应 SQL 提示词不存在" slog.log("ERROR", "stream.sql_generation_failed", trace_id, error_code=ErrorCode.SQL_GENERATION_FAILED.value, payload={"conversation_id": conversation_id}) yield _build_message(conversation_id, reason) yield "event: end\ndata: [DONE]\n\n" return # 2) 先流式返回 SQL yield _build_message(conversation_id, sql_text) # 3) 异步执行 SQL,并及时流式返回执行结果 tool = SrApiQueryTool() task = asyncio.create_task( asyncio.to_thread(tool.run, json.dumps({"sql": sql_text}, ensure_ascii=False)) ) while not task.done(): yield _build_message(conversation_id, "executing_sql") await asyncio.sleep(progress_interval) sql_result = await task slog.log("INFO", "stream.sql_executed", trace_id, {"result_len": len(str(sql_result))}) yield _build_message(conversation_id, str(sql_result)) yield "event: end\ndata: [DONE]\n\n" except Exception as e: slog.log("ERROR", "stream.failed", trace_id, error_code=ErrorCode.INTERNAL_ERROR.value, payload={"error": str(e)}) conversation_id = str(payload.conversation_id or task_id) yield _build_message(conversation_id, str(e)) yield "event: end\ndata: [DONE]\n\n" return StreamingResponse(event_stream(), media_type="text/event-stream") @router.get("/api/workflows/list") def list_workflows(workflow_manager=Depends(get_workflow_manager)): """列出所有可用工作流""" workflows = workflow_manager.get_available_workflows() result = [] for name in workflows: info = workflow_manager.get_workflow_info(name) if info: result.append(info) return {"workflows": result} @router.get("/api/workflows/{workflow_name}") def get_workflow_detail(workflow_name: str, workflow_manager=Depends(get_workflow_manager)): """获取工作流详情""" info = workflow_manager.get_workflow_info(workflow_name) if not info: raise HTTPException(status_code=404, detail=f"工作流不存在: {workflow_name}") return info @router.get("/api/tools/list") def list_tools(tool_router=Depends(get_tool_router)): """列出所有可用工具""" tools = tool_router.list_tools() result = [] for name in tools: info = tool_router.get_tool_info(name) if info: result.append(info) return {"tools": result} @router.get("/api/tools/{tool_name}") def get_tool_detail(tool_name: str, tool_router=Depends(get_tool_router)): """获取工具详情""" info = tool_router.get_tool_info(tool_name) if not info: raise HTTPException(status_code=404, detail=f"工具不存在: {tool_name}") return info @router.get("/api/tools/stats") def get_tools_stats(tool_router=Depends(get_tool_router)): """获取工具执行统计""" return tool_router.get_all_stats() @router.post("/api/tools/execute", response_model=ToolOutput) def run_tool(payload: ToolInput, tool_router=Depends(get_tool_router)): result = tool_router.call(payload.tool_name, payload.payload) return ToolOutput(**result) @router.post("/api/prompts/reload") def reload_prompts(prompt_manager=Depends(get_prompt_manager)): prompt_manager.reload() return {"ok": True} @router.post("/api/ragflow/table-retrieval/reload") def reload_table_retrieval(): syncer = RagflowSync() result = syncer.upload_table_retrieval() return {"ok": True, "result": result} @router.post("/api/ragflow/table-retrieval/upload") def upload_table_retrieval(): """上传表名检索模板文档""" syncer = RagflowSync() try: result = syncer.upload_table_retrieval() return {"ok": True, "result": result} except Exception as e: raise HTTPException(status_code=500, detail=str(e)) @router.put("/api/ragflow/table-retrieval/update") def update_table_retrieval(): """更新表名检索文档(仅文档内容)""" syncer = RagflowSync() try: result = syncer.update_table_retrieval_documents() return {"ok": True, "result": result} except Exception as e: raise HTTPException(status_code=500, detail=str(e)) @router.post("/api/ragflow/sql-gen/upload") def upload_sql_gen(): """上传 SQL 生成提示词文档""" syncer = RagflowSync() try: result = syncer.upload_sql_gen() return {"ok": True, "result": result} except Exception as e: raise HTTPException(status_code=500, detail=str(e)) @router.put("/api/ragflow/sql-gen/update") def update_sql_gen(): """更新 SQL 生成文档(仅文档内容)""" syncer = RagflowSync() try: result = syncer.update_sql_gen_documents() return {"ok": True, "result": result} except Exception as e: raise HTTPException(status_code=500, detail=str(e)) @router.post("/api/super-agent/query", response_model=SuperAgentResponse) def super_agent_query(payload: SuperAgentRequest, workflow_manager=Depends(get_workflow_manager)): """Super Agent 同步查询接口""" trace_id = uuid.uuid4().hex slog = get_structured_logger() slog.log("INFO", "super_agent.query.start", trace_id, { "query": payload.query[:100], "workflow_type": payload.workflow_type, "user_id": payload.user_id, }) conversation_id = payload.conversation_id or uuid.uuid4().hex try: workflow_type = _resolve_workflow_type(payload.workflow_type) except Exception as e: slog.log("ERROR", "super_agent.query.invalid_type", trace_id, error_code=ErrorCode.INVALID_WORKFLOW_TYPE.value) return SuperAgentResponse( conversation_id=conversation_id, workflow_type=payload.workflow_type, status="error", error=f"不支持的工作流类型: {payload.workflow_type}", ) try: result = workflow_manager.execute_workflow( workflow_type=workflow_type, user_input=payload.query, session_id=conversation_id, ) context = (result.get("result") or {}).get("context") or {} sql_text = context.get("final_sql") sr_api_result = context.get("sr_api_result") slog.log("INFO", "super_agent.query.success", trace_id, { "conversation_id": conversation_id, "has_sql": bool(sql_text), "has_result": bool(sr_api_result), }) return SuperAgentResponse( conversation_id=conversation_id, workflow_type=workflow_type.value, status="success", sql=sql_text, result=str(sr_api_result) if sr_api_result else None, metadata={"trace_id": trace_id}, ) except Exception as e: slog.log("ERROR", "super_agent.query.failed", trace_id, error_code=ErrorCode.INTERNAL_ERROR.value, payload={"error": str(e)}) return SuperAgentResponse( conversation_id=conversation_id, workflow_type=payload.workflow_type, status="error", error=str(e), metadata={"trace_id": trace_id}, ) @router.post("/api/super-agent/stream") def super_agent_stream(payload: SuperAgentRequest, workflow_manager=Depends(get_workflow_manager)): """Super Agent 流式查询接口""" trace_id = uuid.uuid4().hex slog = get_structured_logger() stream_cfg = Config.get_section("stream") progress_interval = float(stream_cfg.get("progress_interval", 0.3)) conversation_id = payload.conversation_id or uuid.uuid4().hex def _build_sse_event(event: str, data: str) -> str: dto = SuperAgentStreamEvent( conversation_id=conversation_id, event=event, data=data, timestamp=int(time.time() * 1000), ) return f"event: {event}\ndata: {json.dumps(dto.model_dump(), ensure_ascii=False)}\n\n" async def event_stream(): try: slog.log("INFO", "super_agent.stream.start", trace_id, { "query": payload.query[:100], "user_id": payload.user_id, }) workflow_type = _resolve_workflow_type(payload.workflow_type) result = await asyncio.to_thread( workflow_manager.execute_workflow, workflow_type, payload.query, conversation_id, skip_sr_api=True, ) context = (result.get("result") or {}).get("context") or {} sql_text = context.get("final_sql") if not sql_text: slog.log("ERROR", "super_agent.stream.sql_failed", trace_id, error_code=ErrorCode.SQL_GENERATION_FAILED.value) yield _build_sse_event("error", "SQL 生成失败") yield _build_sse_event("done", "") return yield _build_sse_event("sql_generated", sql_text) yield _build_sse_event("sql_executing", "") tool = SrApiQueryTool() task = asyncio.create_task( asyncio.to_thread(tool.run, json.dumps({"sql": sql_text}, ensure_ascii=False)) ) while not task.done(): yield _build_sse_event("sql_executing", "") await asyncio.sleep(progress_interval) sql_result = await task slog.log("INFO", "super_agent.stream.success", trace_id, {"result_len": len(str(sql_result))}) yield _build_sse_event("result", str(sql_result)) except Exception as e: slog.log("ERROR", "super_agent.stream.failed", trace_id, error_code=ErrorCode.INTERNAL_ERROR.value, payload={"error": str(e)}) yield _build_sse_event("error", str(e)) yield _build_sse_event("done", "") return StreamingResponse(event_stream(), media_type="text/event-stream")