435 lines
16 KiB
Python
435 lines
16 KiB
Python
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")
|