init
This commit is contained in:
+516
-240
@@ -1,30 +1,180 @@
|
||||
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 time
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
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 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.app_errors import AppError, ErrorCode
|
||||
from services.ragflow_sync import RagflowSync
|
||||
from services.structured_logger import get_structured_logger
|
||||
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)
|
||||
@@ -42,6 +192,122 @@ def _to_http_error(e: Exception) -> HTTPException:
|
||||
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 {
|
||||
@@ -57,50 +323,90 @@ def nacos_status(nacos_manager=Depends(get_nacos_manager)):
|
||||
|
||||
|
||||
@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})
|
||||
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:
|
||||
workflow_type = _resolve_workflow_type(payload.workflow_type)
|
||||
conversation_id = _handle_conversation(payload, current_timestamp, msg_storage)
|
||||
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)
|
||||
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=payload.query,
|
||||
session_id=payload.conversation_id,
|
||||
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=result["session_id"],
|
||||
session_id=conversation_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)):
|
||||
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()
|
||||
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)
|
||||
workflow_type = WorkflowType.CONVERSATION
|
||||
|
||||
result = workflow_manager.execute_workflow(
|
||||
workflow_type=workflow_type,
|
||||
user_input=payload.query,
|
||||
session_id=payload.conversation_id,
|
||||
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 {}
|
||||
@@ -113,7 +419,7 @@ def generate_sql(payload: AgentInput, workflow_manager=Depends(get_workflow_mana
|
||||
slog.log("INFO", "generate_sql.success", trace_id, {"sql_len": len(sql_text)})
|
||||
|
||||
return {
|
||||
"session_id": result.get("session_id"),
|
||||
"session_id": conversation_id,
|
||||
"workflow_type": result.get("workflow_type"),
|
||||
"sql": sql_text,
|
||||
}
|
||||
@@ -121,122 +427,200 @@ def generate_sql(payload: AgentInput, workflow_manager=Depends(get_workflow_mana
|
||||
|
||||
@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_WORKFLOW_TYPE, message="/api/workflows/stream 仅支持 response_mode=streaming", status_code=400))
|
||||
raise _to_http_error(
|
||||
AppError(
|
||||
code=ErrorCode.INVALID_RESPONSE_MODE,
|
||||
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
|
||||
message_id = task_id
|
||||
chunk_size = 1024
|
||||
|
||||
def _build_message(conversation_id: str, answer: str) -> str:
|
||||
def _build_stream_chunk(conversation_id: str, answer: str, event: str = "message") -> str:
|
||||
dto = ChatMessageResponseDTO(
|
||||
id=uuid.uuid4().hex,
|
||||
event="message",
|
||||
event=event,
|
||||
task_id=task_id,
|
||||
message_id=uuid.uuid4().hex,
|
||||
message_id=message_id,
|
||||
conversation_id=conversation_id,
|
||||
answer=answer,
|
||||
created_at=int(time.time()),
|
||||
created_at=DateTimeGenerator.now().epoch_seconds,
|
||||
)
|
||||
return f"event: message\ndata: {json.dumps(dto.model_dump(), ensure_ascii=False)}\n\n"
|
||||
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})
|
||||
# 1) 先仅生成 SQL(不执行 SR API)
|
||||
|
||||
# 执行 workflow(skip_sr_api=False,SQL 执行在 workflow 内部完成)
|
||||
result = await asyncio.to_thread(
|
||||
workflow_manager.execute_workflow,
|
||||
WorkflowType.CONVERSATION,
|
||||
payload.query,
|
||||
payload.conversation_id,
|
||||
skip_sr_api=True,
|
||||
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],
|
||||
)
|
||||
conversation_id = str(result.get("session_id") or payload.conversation_id or task_id)
|
||||
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))
|
||||
|
||||
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
|
||||
# 查询表的 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))
|
||||
|
||||
# 2) 先流式返回 SQL
|
||||
yield _build_message(conversation_id, sql_text)
|
||||
# 从 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))
|
||||
|
||||
# 3) 异步执行 SQL,并及时流式返回执行结果
|
||||
tool = SrApiQueryTool()
|
||||
task = asyncio.create_task(
|
||||
asyncio.to_thread(tool.run, json.dumps({"sql": sql_text}, 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)},
|
||||
)
|
||||
|
||||
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"
|
||||
|
||||
stream_logs: list[str] = []
|
||||
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/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)
|
||||
@@ -302,133 +686,25 @@ def update_sql_gen():
|
||||
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
|
||||
|
||||
def get_table_etl_time(table_name: str) -> str:
|
||||
"""
|
||||
查询指定表的最大 etl_time,若失败或无数据则返回 Unknown。
|
||||
"""
|
||||
if not table_name:
|
||||
return "Unknown"
|
||||
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")
|
||||
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"
|
||||
|
||||
Reference in New Issue
Block a user