from fastapi import APIRouter, HTTPException, Depends from fastapi.responses import StreamingResponse 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 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.ragflow_sync import RagflowSync router = APIRouter() def _resolve_workflow_type(value: str) -> WorkflowType: try: return WorkflowType(value) except Exception as e: raise ValueError(f"不支持的工作流类型: {value}") from 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)): try: workflow_type = _resolve_workflow_type(payload.workflow_type) except ValueError as e: raise HTTPException(status_code=400, detail=str(e)) result = workflow_manager.execute_workflow( workflow_type=workflow_type, user_input=payload.input, session_id=payload.session_id, ) return AgentOutput( session_id=result["session_id"], workflow_type=result["workflow_type"], result=result["result"], ) @router.post("/api/workflows/stream") def run_workflow_stream(payload: AgentInput, workflow_manager=Depends(get_workflow_manager)): try: workflow_type = _resolve_workflow_type(payload.workflow_type) except ValueError as e: raise HTTPException(status_code=400, detail=str(e)) if workflow_type != WorkflowType.CONVERSATION: raise HTTPException(status_code=400, detail="仅支持对话工作流的流式输出") def _extract_output_text(result: dict) -> str: context = (result.get("context") or {}) if isinstance(result, dict) else {} if "sr_api_result" in context: return str(context.get("sr_api_result") or "") messages = result.get("messages") if isinstance(result, dict) else None if messages: last = messages[-1] if hasattr(last, "content"): return str(last.content or "") return "" def event_stream(): try: result = workflow_manager.execute_workflow( workflow_type=workflow_type, user_input=payload.input, session_id=payload.session_id, ) text = _extract_output_text(result.get("result") or {}) if not text: yield "event: end\ndata: [DONE]\n\n" return chunk_size = 512 for i in range(0, len(text), chunk_size): chunk = text[i : i + chunk_size] yield f"data: {chunk}\n\n" yield "event: end\ndata: [DONE]\n\n" except Exception as e: yield f"event: error\ndata: {str(e)}\n\n" return StreamingResponse(event_stream(), media_type="text/event-stream") @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(config: dict): """更新表名检索知识库配置""" syncer = RagflowSync() try: result = syncer.update_dataset(syncer._table_retrieval_dataset_id, config) 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(config: dict): """更新 SQL 生成知识库配置""" syncer = RagflowSync() try: result = syncer.update_dataset(syncer._sql_gen_dataset_id, config) return {"ok": True, "result": result} except Exception as e: raise HTTPException(status_code=500, detail=str(e))