def safe_json_dumps(obj): try: return json.dumps(obj, ensure_ascii=False, default=str) except Exception as e: return f"" import asyncio import html import json import uuid from typing import Any from fastapi import APIRouter, HTTPException, Depends from fastapi.responses import StreamingResponse from schemas.agent_output import AgentOutput from schemas.tool_input import ToolInput from schemas.tool_output import ToolOutput from schemas.chat_message_request import ChatMessageRequestDTO from schemas.chat_message_response import ChatMessageResponseDTO from schemas.message_feedback_request import MessageFeedbackRequestDTO 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.common.app_errors import AppError, ErrorCode from services.common.datetime_utils import DateTimeGenerator from services.integrations.ragflow_sync import RagflowSync from services.storage.structured_logger import get_structured_logger from services.storage.message_storage import get_message_storage from tools.sr_api_tool import SrApiQueryTool router = APIRouter() def _try_json_loads(value: Any) -> Any: if isinstance(value, (dict, list)): return value if isinstance(value, str): stripped = value.strip() if stripped.startswith("{") or stripped.startswith("["): try: return json.loads(stripped) except Exception: return value return value def _extract_total_rows(raw_sql_result: Any) -> int: parsed = _try_json_loads(raw_sql_result) if isinstance(parsed, dict): inner = _try_json_loads(parsed.get("text")) if "text" in parsed else parsed if isinstance(inner, dict): if isinstance(inner.get("total"), int): return int(inner.get("total") or 0) data = inner.get("data") if isinstance(data, list): return len(data) if isinstance(data, dict): rows = data.get("rows") or data.get("list") or data.get("records") if isinstance(rows, list): return len(rows) if isinstance(parsed, list): return len(parsed) return 0 def _build_final_answer(raw_sql_result: Any) -> str: total = _extract_total_rows(raw_sql_result) if total <= 0: return "未查询到符合条件的数据,请尝试调整筛选条件后再查询。" return f"查询完成,共返回 {total} 条记录。" def _extract_sql_rows(raw_sql_result: Any) -> list[dict[str, Any]]: parsed = _try_json_loads(raw_sql_result) payload = parsed if isinstance(parsed, dict): payload = _try_json_loads(parsed.get("text")) if "text" in parsed else parsed if isinstance(payload, dict): rows = payload.get("data") if isinstance(rows, list): normalized: list[dict[str, Any]] = [] for item in rows: if isinstance(item, dict): normalized.append(item) else: normalized.append({"value": item}) return normalized if isinstance(payload, list): normalized = [] for item in payload: if isinstance(item, dict): normalized.append(item) else: normalized.append({"value": item}) return normalized return [] def _rows_to_html_table(rows: list[dict[str, Any]]) -> str: if not rows: return "No data~" headers: list[str] = [] for row in rows: for key in row.keys(): if key not in headers: headers.append(str(key)) if not headers: return "No data~" thead = "".join(f"{html.escape(header)}" for header in headers) body_rows = [] for row in rows: cells = [] for header in headers: value = row.get(header) cell_text = "" if value is None else str(value) cells.append(f"{html.escape(cell_text)}") body_rows.append(f"{''.join(cells)}") return f"{thead}{''.join(body_rows)}
" def _extract_etl_version(rows: list[dict[str, Any]]) -> str: versions: list[tuple[int, str]] = [] for row in rows: if not isinstance(row, dict): continue etl_value = row.get("etl_time") if etl_value in (None, ""): continue etl_text = str(etl_value).strip() if not etl_text: continue try: bundle = DateTimeGenerator.bundle(etl_value, default_to_now=False) versions.append((bundle.epoch_millis, bundle.datetime_str)) except Exception: versions.append((0, etl_text)) if not versions: return "Unknown" versions.sort(key=lambda item: item[0], reverse=True) return versions[0][1] def _build_rich_answer_html(query: str, rows: list[dict[str, Any]], etl_version: str = None) -> str: safe_query = html.escape((query or "").strip()) row_count = len(rows or []) if etl_version is None: etl_version = _extract_etl_version(rows) if row_count <= 0: # 空数据返回纯文本,避免前端直接显示 HTML 标签 return ( f"Question: {safe_query}\n" "No data~\n" f"Rows: {row_count}\n" f"Data Version: {etl_version}" ) table_html = _rows_to_html_table(rows) return ( f"
Question: {safe_query}
" f"
{table_html}
" f"
Rows: {row_count}
" f"
Data Version: {etl_version}
" ) 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)}) def _require_query_text(payload: ChatMessageRequestDTO) -> str: query = payload.query if not isinstance(query, str) or not query.strip(): raise _to_http_error( AppError( code=ErrorCode.INVALID_REQUEST, message="query 不能为空", status_code=400, detail={"field": "query", "reason": "missing_or_blank"}, ) ) return query def _build_conversation_name(query: str, max_chars: int = 20) -> str: return (query or "")[:max_chars] def _storage_enabled(msg_storage) -> bool: return bool(getattr(msg_storage, "enabled", False)) def _handle_conversation(payload: ChatMessageRequestDTO, current_timestamp: int, msg_storage) -> str: provided_conversation_id = (payload.conversation_id or "").strip() if isinstance(payload.conversation_id, str) else "" if not _storage_enabled(msg_storage): return provided_conversation_id or uuid.uuid4().hex if not provided_conversation_id: conversation_id = uuid.uuid4().hex created = msg_storage.create_conversation( conversation_id=conversation_id, user=payload.user, name=_build_conversation_name(payload.query or ""), status="normal", introduction=None, created_at=current_timestamp, updated_at=current_timestamp, ) if not created: raise AppError( code=ErrorCode.CONVERSATION_CREATE_FAILED, message="会话创建失败", status_code=500, detail={"conversation_id": conversation_id}, ) return conversation_id conversation = msg_storage.get_conversation_by_id(provided_conversation_id) if conversation is None: raise AppError( code=ErrorCode.CONVERSATION_NOT_FOUND, message="会话不存在", status_code=400, detail={"conversation_id": provided_conversation_id}, ) updated = msg_storage.update_conversation_updated_at(provided_conversation_id, current_timestamp) if not updated: raise AppError( code=ErrorCode.CONVERSATION_UPDATE_FAILED, message="会话更新时间失败", status_code=500, detail={"conversation_id": provided_conversation_id}, ) return provided_conversation_id def _extract_answer_text(result_payload: Any) -> str: result_obj = (result_payload or {}).get("result") if isinstance(result_payload, dict) else None if isinstance(result_obj, dict): messages = result_obj.get("messages") if isinstance(messages, list): for msg in reversed(messages): content = getattr(msg, "content", None) if content: return str(content) context = result_obj.get("context") if isinstance(context, dict) and context.get("final_sql"): return str(context.get("final_sql")) return json.dumps(result_obj or {}, ensure_ascii=False, default=str) def _safe_save_message(msg_storage, **kwargs) -> bool: if hasattr(msg_storage, "save_message"): try: return bool(msg_storage.save_message(**kwargs)) except Exception: return False return False def _persist_message_or_raise(msg_storage, slog, trace_id: str, **kwargs) -> None: if not _storage_enabled(msg_storage): return saved = _safe_save_message(msg_storage, **kwargs) if not saved: slog.log( "ERROR", "message_save_failed", trace_id, payload={"conversation_id": kwargs.get("conversation_id"), "message_id": kwargs.get("message_id")}, ) raise _to_http_error( AppError( code=ErrorCode.MESSAGE_SAVE_FAILED, message="消息保存失败", status_code=500, detail={ "conversation_id": kwargs.get("conversation_id"), "message_id": kwargs.get("message_id"), }, ) ) @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: ChatMessageRequestDTO, workflow_manager=Depends(get_workflow_manager)): query = _require_query_text(payload) msg_storage = get_message_storage() current_timestamp = DateTimeGenerator.now().epoch_millis try: conversation_id = _handle_conversation(payload, current_timestamp, msg_storage) except Exception as e: raise _to_http_error(e) trace_id = uuid.uuid4().hex message_id = uuid.uuid4().hex slog = get_structured_logger() workflow_type = WorkflowType.CONVERSATION save_logs: list[str] = [f"run_workflow.start response_mode={payload.response_mode}"] slog.log("INFO", "run_workflow.start", trace_id, {"workflow_type": workflow_type.value, "response_mode": payload.response_mode}) try: result = workflow_manager.execute_workflow( workflow_type=workflow_type, user_input=query, session_id=conversation_id, user=payload.user, inputs=payload.inputs, files=[item.model_dump() for item in payload.files], ) save_logs.append(f"run_workflow.success session_id={result.get('session_id')}") slog.log("INFO", "run_workflow.success", trace_id, {"session_id": result.get("session_id")}) except Exception as e: save_logs.append(f"run_workflow.failed error={e}") slog.log("ERROR", "run_workflow.failed", trace_id, error_code=ErrorCode.INTERNAL_ERROR.value, payload={"error": str(e)}) raise _to_http_error(e) answer_text = _extract_answer_text(result) _persist_message_or_raise( msg_storage, slog, trace_id, conversation_id=conversation_id, message_id=message_id, query=query, answer=answer_text, workflow_type=WorkflowType.CONVERSATION.value, user=payload.user, metadata={ "trace_id": trace_id, "response_mode": payload.response_mode, "inputs": payload.inputs, "files": [item.model_dump() for item in payload.files], }, created_at=current_timestamp, updated_at=current_timestamp, logs=save_logs, ) if _storage_enabled(msg_storage): slog.log("INFO", "run_workflow.message_saved", trace_id, {"conversation_id": conversation_id, "message_id": message_id}) return AgentOutput( session_id=conversation_id, workflow_type=result["workflow_type"], result=result["result"], ) @router.post("/api/sql/generate") def generate_sql(payload: ChatMessageRequestDTO, workflow_manager=Depends(get_workflow_manager)): """仅生成 SQL,不调用 SR API""" query = _require_query_text(payload) msg_storage = get_message_storage() current_timestamp = DateTimeGenerator.now().epoch_millis try: conversation_id = _handle_conversation(payload, current_timestamp, msg_storage) except Exception as e: raise _to_http_error(e) trace_id = uuid.uuid4().hex slog = get_structured_logger() workflow_type = WorkflowType.CONVERSATION result = workflow_manager.execute_workflow( workflow_type=workflow_type, user_input=query, session_id=conversation_id, skip_sr_api=True, user=payload.user, inputs=payload.inputs, files=[item.model_dump() for item in payload.files], ) 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": conversation_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)): query = _require_query_text(payload) trace_id = uuid.uuid4().hex slog = get_structured_logger() msg_storage = get_message_storage() current_timestamp = DateTimeGenerator.now().epoch_millis try: conversation_id = _handle_conversation(payload, current_timestamp, msg_storage) except Exception as e: raise _to_http_error(e) if payload.response_mode != "streaming": raise _to_http_error( AppError( code=ErrorCode.INVALID_RESPONSE_MODE, message="/api/workflows/stream 仅支持 response_mode=streaming", status_code=400, ) ) task_id = uuid.uuid4().hex message_id = task_id chunk_size = 1024 def _build_stream_chunk(conversation_id: str, answer: str, event: str = "message") -> str: dto = ChatMessageResponseDTO( id=uuid.uuid4().hex, event=event, task_id=task_id, message_id=message_id, conversation_id=conversation_id, answer=answer, created_at=DateTimeGenerator.now().epoch_seconds, ) return f"data: {json.dumps(dto.model_dump(), ensure_ascii=False)}\n\n" def _record_stream_message(stage: str, active_conversation_id: str, final_answer: str, *, sql_text: str | None = None, execution_result: Any = None, extra_metadata: dict[str, Any] | None = None) -> None: if not _storage_enabled(msg_storage): return saved = _safe_save_message( msg_storage, conversation_id=active_conversation_id, message_id=message_id, query=query, answer=final_answer, workflow_type=WorkflowType.CONVERSATION.value, user=payload.user, sql_query=sql_text, execution_result=execution_result, metadata={ "trace_id": trace_id, "stage": stage, "inputs": payload.inputs, "files": [item.model_dump() for item in payload.files], **(extra_metadata or {}), }, created_at=current_timestamp, updated_at=current_timestamp, logs=stream_logs, ) if not saved: slog.log("ERROR", "stream.message_save_failed", trace_id, payload={"conversation_id": active_conversation_id, "message_id": message_id, "stage": stage}) else: slog.log("INFO", "stream.message_saved", trace_id, {"conversation_id": active_conversation_id, "message_id": message_id, "stage": stage}) async def event_stream(): active_conversation_id = conversation_id answer_parts: list[str] = [] try: stream_logs.append(json.dumps({"event": "start", "conversation_id": active_conversation_id}, ensure_ascii=False)) slog.log("INFO", "stream.start", trace_id, {"workflow_type": WorkflowType.CONVERSATION.value}) # 执行 workflow(skip_sr_api=False,SQL 执行在 workflow 内部完成) result = await asyncio.to_thread( workflow_manager.execute_workflow, WorkflowType.CONVERSATION, query, active_conversation_id, skip_sr_api=False, # 在 workflow 内部执行 SQL user=payload.user, inputs=payload.inputs, files=[item.model_dump() for item in payload.files], ) active_conversation_id = str(result.get("session_id") or active_conversation_id or task_id) stream_logs.append(json.dumps({"event": "workflow_success", "session_id": active_conversation_id}, ensure_ascii=False)) context = (result.get("result") or {}).get("context") or {} # 记录关键节点信息 workflow_steps = { "event": "workflow_steps", "normalized_input": context.get("normalized_input"), "query_mode": context.get("query_mode"), "table_name": context.get("table_name"), "current_step": context.get("current_step"), "is_empty_result": context.get("is_empty_result"), } stream_logs.append(json.dumps(workflow_steps, ensure_ascii=False)) sql_text = str(context.get("final_sql") or "") stream_logs.append(json.dumps({"event": "sql_generated", "sql": sql_text}, ensure_ascii=False)) # 从 context 获取完整表名 table_name = None sql_plan = context.get("sql_plan") or {} if sql_plan.get("data_source"): table_name = sql_plan["data_source"] elif context.get("table_name"): table_name = context["table_name"] elif isinstance(context.get("table_match"), dict): table_name = context["table_match"].get("table_name") stream_logs.append(json.dumps({"event": "table_name_extracted", "table_name": table_name}, ensure_ascii=False)) # 查询表的 etl_time etl_version = None if table_name: try: etl_version = get_table_etl_time(table_name) stream_logs.append(json.dumps({"event": "etl_version", "table": table_name, "etl_version": etl_version}, ensure_ascii=False)) except Exception as e: stream_logs.append(json.dumps({"event": "etl_version_failed", "table": table_name, "error": str(e)}, ensure_ascii=False)) etl_version = None else: stream_logs.append(json.dumps({"event": "etl_version_skipped", "reason": "no_table"}, ensure_ascii=False)) # 从 context 获取结果(SQL 已在 workflow 内执行) sr_api_result = context.get("sr_api_result") result_rows = _extract_sql_rows(sr_api_result) if sr_api_result else [] sample_rows = result_rows[:3] if result_rows else [] row_count = len(result_rows) stream_logs.append(json.dumps({"event": "result_rows", "count": row_count, "sample": sample_rows}, ensure_ascii=False)) # 检查是否为空结果 is_empty_result = context.get("is_empty_result", False) if is_empty_result and context.get("formatted_answer"): result_text = context["formatted_answer"] stream_logs.append(json.dumps({"event": "empty_result_formatted"}, ensure_ascii=False)) elif sr_api_result: result_text = _build_rich_answer_html(query, result_rows, etl_version=etl_version) stream_logs.append(json.dumps({"event": "rich_answer_html", "etl_version": etl_version}, ensure_ascii=False)) else: result_text = "查询执行完成,但未获取到结果数据。" stream_logs.append(json.dumps({"event": "no_result"}, ensure_ascii=False)) for index in range(0, len(result_text), chunk_size): chunk = result_text[index:index + chunk_size] answer_parts.append(chunk) yield _build_stream_chunk(active_conversation_id, chunk) yield _build_stream_chunk(active_conversation_id, "", event="message_end") _record_stream_message( "success", active_conversation_id, "".join(answer_parts), sql_text=sql_text, execution_result={"status": "success", "row_count": row_count, "sample": sample_rows, "is_empty": is_empty_result}, ) except Exception as e: stream_logs.append(f"stream.failed error={e}") slog.log("ERROR", "stream.failed", trace_id, error_code=ErrorCode.INTERNAL_ERROR.value, payload={"error": str(e)}) error_conversation_id = str(active_conversation_id or payload.conversation_id or task_id) error_text = str(e) answer_parts.append(error_text) yield _build_stream_chunk(error_conversation_id, error_text) yield _build_stream_chunk(error_conversation_id, "", event="message_end") _record_stream_message( "exception", error_conversation_id, "".join(answer_parts), extra_metadata={"error": str(e)}, ) stream_logs: list[str] = [] return StreamingResponse(event_stream(), media_type="text/event-stream") @router.post("/api/messages/feedback") def write_message_feedback(payload: MessageFeedbackRequestDTO): msg_storage = get_message_storage() updated = msg_storage.update_feedback_by_message_id( message_id=payload.message_id, feedback=payload.feedback, feedback_content=payload.feedback_content, ) if not updated: raise _to_http_error( AppError( code=ErrorCode.INVALID_REQUEST, message="反馈写回失败,message_id 不存在或存储未启用", status_code=400, detail={"field": "message_id", "reason": "not_found_or_storage_disabled"}, ) ) return {"ok": True, "message_id": payload.message_id} @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)) def get_table_etl_time(table_name: str) -> str: """ 查询指定表的最大 etl_time,若失败或无数据则返回 Unknown。 """ if not table_name: return "Unknown" try: tool = SrApiQueryTool() sql = f"SELECT MAX(etl_time) AS etl_time FROM {table_name}" result = tool.run(json.dumps({"sql": sql}, ensure_ascii=False)) rows = _extract_sql_rows(result) if rows and rows[0].get("etl_time"): etl_value = rows[0]["etl_time"] # 使用 DateTimeGenerator 转换为格式化字符串 try: bundle = DateTimeGenerator.bundle(etl_value, default_to_now=False) return bundle.datetime_str except Exception: return str(etl_value) except Exception: pass return "Unknown"