711 lines
27 KiB
Python
711 lines
27 KiB
Python
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"
|