249 lines
7.1 KiB
Python
249 lines
7.1 KiB
Python
"""
|
|
统一响应格式模块
|
|
|
|
提供标准化的 API 响应和流式事件格式
|
|
"""
|
|
|
|
from typing import Any, Dict, Generic, List, Optional, TypeVar, Literal
|
|
from pydantic import BaseModel, Field
|
|
import time
|
|
import uuid
|
|
|
|
|
|
T = TypeVar("T")
|
|
|
|
|
|
class ApiResponse(BaseModel, Generic[T]):
|
|
"""
|
|
统一 API 响应格式
|
|
|
|
所有 API 响应都使用这个格式,提供一致的响应结构
|
|
|
|
Usage:
|
|
@router.get("/users/{user_id}")
|
|
async def get_user(user_id: str) -> ApiResponse[User]:
|
|
user = await user_service.get(user_id)
|
|
return ApiResponse.success(data=user)
|
|
"""
|
|
|
|
code: str = Field(default="success", description="响应代码")
|
|
message: str = Field(default="", description="响应消息")
|
|
data: Optional[T] = Field(default=None, description="响应数据")
|
|
trace_id: Optional[str] = Field(default=None, description="追踪ID")
|
|
timestamp: int = Field(
|
|
default_factory=lambda: int(time.time() * 1000),
|
|
description="时间戳(毫秒)"
|
|
)
|
|
|
|
@classmethod
|
|
def success(cls, data: T = None, message: str = "", trace_id: Optional[str] = None) -> "ApiResponse[T]":
|
|
"""创建成功响应"""
|
|
return cls(
|
|
code="success",
|
|
message=message,
|
|
data=data,
|
|
trace_id=trace_id or uuid.uuid4().hex,
|
|
)
|
|
|
|
@classmethod
|
|
def error(
|
|
cls,
|
|
code: str = "error",
|
|
message: str = "",
|
|
data: T = None,
|
|
trace_id: Optional[str] = None,
|
|
) -> "ApiResponse[T]":
|
|
"""创建错误响应"""
|
|
return cls(
|
|
code=code,
|
|
message=message,
|
|
data=data,
|
|
trace_id=trace_id or uuid.uuid4().hex,
|
|
)
|
|
|
|
@classmethod
|
|
def from_exception(cls, exc: Exception, trace_id: Optional[str] = None) -> "ApiResponse[None]":
|
|
"""从异常创建错误响应"""
|
|
return cls.error(
|
|
code="internal_error",
|
|
message=str(exc),
|
|
trace_id=trace_id,
|
|
)
|
|
|
|
def is_success(self) -> bool:
|
|
"""判断是否成功"""
|
|
return self.code == "success"
|
|
|
|
|
|
class PagedResponse(BaseModel, Generic[T]):
|
|
"""
|
|
分页响应格式
|
|
|
|
用于返回分页数据
|
|
"""
|
|
|
|
items: List[T] = Field(default_factory=list, description="数据列表")
|
|
total: int = Field(default=0, description="总数")
|
|
page: int = Field(default=1, description="当前页")
|
|
page_size: int = Field(default=20, description="每页大小")
|
|
total_pages: int = Field(default=0, description="总页数")
|
|
|
|
@classmethod
|
|
def create(
|
|
cls,
|
|
items: List[T],
|
|
total: int,
|
|
page: int = 1,
|
|
page_size: int = 20,
|
|
) -> "PagedResponse[T]":
|
|
"""创建分页响应"""
|
|
total_pages = (total + page_size - 1) // page_size if page_size > 0 else 0
|
|
return cls(
|
|
items=items,
|
|
total=total,
|
|
page=page,
|
|
page_size=page_size,
|
|
total_pages=total_pages,
|
|
)
|
|
|
|
|
|
class StreamEvent(BaseModel):
|
|
"""
|
|
流式响应事件
|
|
|
|
用于 SSE (Server-Sent Events) 流式响应
|
|
|
|
Usage:
|
|
async def event_stream():
|
|
yield StreamEvent(event="start", data="Processing started")
|
|
# ... 处理逻辑
|
|
yield StreamEvent(event="result", data=json.dumps(result))
|
|
yield StreamEvent(event="done", data="")
|
|
"""
|
|
|
|
event: str = Field(..., description="事件类型")
|
|
data: str = Field(default="", description="事件数据")
|
|
event_id: Optional[str] = Field(default=None, description="事件ID")
|
|
retry: Optional[int] = Field(default=None, description="重试间隔(毫秒)")
|
|
|
|
def to_sse(self) -> str:
|
|
"""转换为 SSE 格式字符串"""
|
|
lines = [f"event: {self.event}"]
|
|
if self.event_id:
|
|
lines.append(f"id: {self.event_id}")
|
|
if self.retry:
|
|
lines.append(f"retry: {self.retry}")
|
|
lines.append(f"data: {self.data}")
|
|
lines.append("")
|
|
lines.append("")
|
|
return "\n".join(lines)
|
|
|
|
@classmethod
|
|
def message(cls, data: str, event_id: Optional[str] = None) -> "StreamEvent":
|
|
"""创建消息事件"""
|
|
return cls(event="message", data=data, event_id=event_id)
|
|
|
|
@classmethod
|
|
def done(cls) -> "StreamEvent":
|
|
"""创建完成事件"""
|
|
return cls(event="done", data="[DONE]")
|
|
|
|
@classmethod
|
|
def error(cls, message: str) -> "StreamEvent":
|
|
"""创建错误事件"""
|
|
return cls(event="error", data=message)
|
|
|
|
|
|
class WorkflowEvent(BaseModel):
|
|
"""
|
|
工作流事件
|
|
|
|
用于工作流执行过程中的状态通知
|
|
"""
|
|
|
|
workflow_id: str = Field(..., description="工作流ID")
|
|
event_type: Literal[
|
|
"started",
|
|
"node_started",
|
|
"node_completed",
|
|
"node_failed",
|
|
"completed",
|
|
"failed",
|
|
] = Field(..., description="事件类型")
|
|
node_name: Optional[str] = Field(None, description="节点名称")
|
|
data: Optional[Dict[str, Any]] = Field(None, description="事件数据")
|
|
error: Optional[str] = Field(None, description="错误信息")
|
|
timestamp: int = Field(
|
|
default_factory=lambda: int(time.time() * 1000),
|
|
description="时间戳"
|
|
)
|
|
|
|
@classmethod
|
|
def started(cls, workflow_id: str) -> "WorkflowEvent":
|
|
"""创建开始事件"""
|
|
return cls(workflow_id=workflow_id, event_type="started")
|
|
|
|
@classmethod
|
|
def node_started(cls, workflow_id: str, node_name: str) -> "WorkflowEvent":
|
|
"""创建节点开始事件"""
|
|
return cls(
|
|
workflow_id=workflow_id,
|
|
event_type="node_started",
|
|
node_name=node_name,
|
|
)
|
|
|
|
@classmethod
|
|
def node_completed(
|
|
cls,
|
|
workflow_id: str,
|
|
node_name: str,
|
|
data: Optional[Dict[str, Any]] = None,
|
|
) -> "WorkflowEvent":
|
|
"""创建节点完成事件"""
|
|
return cls(
|
|
workflow_id=workflow_id,
|
|
event_type="node_completed",
|
|
node_name=node_name,
|
|
data=data,
|
|
)
|
|
|
|
@classmethod
|
|
def completed(
|
|
cls,
|
|
workflow_id: str,
|
|
data: Optional[Dict[str, Any]] = None,
|
|
) -> "WorkflowEvent":
|
|
"""创建完成事件"""
|
|
return cls(workflow_id=workflow_id, event_type="completed", data=data)
|
|
|
|
@classmethod
|
|
def failed(
|
|
cls,
|
|
workflow_id: str,
|
|
error: str,
|
|
node_name: Optional[str] = None,
|
|
) -> "WorkflowEvent":
|
|
"""创建失败事件"""
|
|
return cls(
|
|
workflow_id=workflow_id,
|
|
event_type="failed",
|
|
node_name=node_name,
|
|
error=error,
|
|
)
|
|
|
|
|
|
class ErrorCode:
|
|
"""错误代码常量"""
|
|
|
|
SUCCESS = "success"
|
|
UNKNOWN_ERROR = "unknown_error"
|
|
INVALID_REQUEST = "invalid_request"
|
|
INVALID_WORKFLOW_TYPE = "invalid_workflow_type"
|
|
SQL_GENERATION_FAILED = "sql_generation_failed"
|
|
TOOL_NOT_FOUND = "tool_not_found"
|
|
TOOL_EXECUTION_FAILED = "tool_execution_failed"
|
|
INTERNAL_ERROR = "internal_error"
|
|
TIMEOUT = "timeout"
|
|
RATE_LIMITED = "rate_limited"
|
|
UNAUTHORIZED = "unauthorized"
|