Files
2026-03-24 18:07:22 +08:00

711 lines
27 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
def safe_json_dumps(obj):
try:
return json.dumps(obj, ensure_ascii=False, default=str)
except Exception as e:
return f"<unserializable: {e}>"
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"<th>{html.escape(header)}</th>" 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"<td>{html.escape(cell_text)}</td>")
body_rows.append(f"<tr>{''.join(cells)}</tr>")
return f"<table><thead><tr>{thead}</tr></thead><tbody>{''.join(body_rows)}</tbody></table>"
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"<div><strong>Question:</strong> {safe_query}</div>"
f"<div style='margin-top:8px;'>{table_html}</div>"
f"<div style='margin-top:8px;'><strong>Rows:</strong> {row_count}</div>"
f"<div><strong>Data Version:</strong> {etl_version}</div>"
)
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"