forked from xiaohei/taiji-AI-PAD
574 lines
14 KiB
Python
574 lines
14 KiB
Python
"""
|
||
Pydantic schemas for API requests and responses
|
||
"""
|
||
|
||
from datetime import datetime
|
||
from typing import Any, Dict, List, Optional, Union, TypeVar, Generic
|
||
from pydantic import BaseModel, Field, validator
|
||
import uuid
|
||
|
||
T = TypeVar('T')
|
||
|
||
|
||
class BaseSchema(BaseModel):
|
||
"""基础schema类"""
|
||
class Config:
|
||
from_attributes = True
|
||
json_encoders = {
|
||
datetime: lambda v: v.isoformat(),
|
||
uuid.UUID: lambda v: str(v),
|
||
}
|
||
|
||
@staticmethod
|
||
def _convert_uuid(value):
|
||
"""转换asyncpg的UUID为标准UUID"""
|
||
if value is None:
|
||
return None
|
||
try:
|
||
from asyncpg.pgproto.pgproto import UUID as AsyncUUID
|
||
if isinstance(value, AsyncUUID):
|
||
return uuid.UUID(str(value))
|
||
except (ImportError, AttributeError):
|
||
pass
|
||
if isinstance(value, uuid.UUID):
|
||
return value
|
||
try:
|
||
return uuid.UUID(str(value))
|
||
except (ValueError, TypeError):
|
||
return value
|
||
|
||
|
||
# ========== MCP协议相关 ==========
|
||
|
||
class MCPRequest(BaseModel):
|
||
"""MCP请求模型"""
|
||
jsonrpc: str = "2.0"
|
||
id: Union[str, int] = Field(default_factory=lambda: str(uuid.uuid4()))
|
||
method: str
|
||
params: Optional[Dict[str, Any]] = None
|
||
|
||
|
||
class MCPResponse(BaseModel):
|
||
"""MCP响应模型"""
|
||
jsonrpc: str = "2.0"
|
||
id: Union[str, int]
|
||
result: Optional[Any] = None
|
||
error: Optional[Dict[str, Any]] = None
|
||
|
||
|
||
class MCPError(BaseModel):
|
||
"""MCP错误模型"""
|
||
code: int
|
||
message: str
|
||
data: Optional[Any] = None
|
||
|
||
|
||
# ========== 工具相关 ==========
|
||
|
||
class ToolParameter(BaseModel):
|
||
"""工具参数定义"""
|
||
name: str
|
||
type: str
|
||
description: Optional[str] = None
|
||
required: bool = True
|
||
default: Optional[Any] = None
|
||
enum: Optional[List[Any]] = None
|
||
|
||
|
||
class ToolDefinition(BaseSchema):
|
||
"""工具定义"""
|
||
name: str
|
||
description: str
|
||
category: Optional[str] = None
|
||
parameters: List[ToolParameter] = []
|
||
returns: Optional[Dict[str, Any]] = None
|
||
|
||
# API相关
|
||
endpoint: Optional[str] = None
|
||
method: str = "POST"
|
||
headers: Optional[Dict[str, str]] = None
|
||
|
||
# 限制信息
|
||
rate_limit: int = 100
|
||
timeout: int = 30
|
||
cost_per_call: float = 0.0
|
||
|
||
|
||
class ToolExecution(BaseModel):
|
||
"""工具执行请求"""
|
||
tool_name: str
|
||
parameters: Dict[str, Any]
|
||
timeout: Optional[int] = None
|
||
|
||
|
||
class ToolResult(BaseSchema):
|
||
"""工具执行结果"""
|
||
success: bool
|
||
result: Optional[Any] = None
|
||
error: Optional[str] = None
|
||
execution_time: float = 0.0
|
||
cost: float = 0.0
|
||
|
||
|
||
class ToolCreate(BaseModel):
|
||
"""创建工具请求"""
|
||
name: str = Field(..., min_length=1, max_length=100)
|
||
description: Optional[str] = None
|
||
category: Optional[str] = None
|
||
|
||
# 工具定义
|
||
schema: Dict[str, Any] = Field(..., description="OpenAPI or Pydantic schema")
|
||
endpoint: Optional[str] = None
|
||
method: str = Field(default="POST", pattern="^(GET|POST|PUT|DELETE|PATCH)$")
|
||
|
||
# 认证信息
|
||
auth_type: Optional[str] = None
|
||
auth_config: Dict[str, Any] = Field(default_factory=dict)
|
||
|
||
# 限制和配额
|
||
rate_limit: int = Field(default=100, ge=1, le=10000)
|
||
cost_per_call: float = Field(default=0.0, ge=0.0)
|
||
timeout: int = Field(default=30, ge=1, le=300)
|
||
|
||
# 状态信息
|
||
is_active: bool = True
|
||
is_public: bool = False
|
||
|
||
|
||
class ToolUpdate(BaseModel):
|
||
"""更新工具请求"""
|
||
name: Optional[str] = Field(None, min_length=1, max_length=100)
|
||
description: Optional[str] = None
|
||
category: Optional[str] = None
|
||
|
||
schema: Optional[Dict[str, Any]] = None
|
||
endpoint: Optional[str] = None
|
||
method: Optional[str] = Field(None, pattern="^(GET|POST|PUT|DELETE|PATCH)$")
|
||
|
||
auth_type: Optional[str] = None
|
||
auth_config: Optional[Dict[str, Any]] = None
|
||
|
||
rate_limit: Optional[int] = Field(None, ge=1, le=10000)
|
||
cost_per_call: Optional[float] = Field(None, ge=0.0)
|
||
timeout: Optional[int] = Field(None, ge=1, le=300)
|
||
|
||
is_active: Optional[bool] = None
|
||
is_public: Optional[bool] = None
|
||
|
||
|
||
class ToolResponse(BaseSchema):
|
||
"""工具响应"""
|
||
id: uuid.UUID
|
||
name: str
|
||
description: Optional[str]
|
||
category: Optional[str]
|
||
|
||
schema: Dict[str, Any]
|
||
endpoint: Optional[str]
|
||
method: str
|
||
|
||
auth_type: Optional[str]
|
||
# Note: auth_config is excluded from response for security
|
||
|
||
rate_limit: int
|
||
cost_per_call: float
|
||
timeout: int
|
||
|
||
is_active: bool
|
||
is_public: bool
|
||
|
||
total_calls: int
|
||
success_rate: float
|
||
avg_response_time: float
|
||
|
||
owner_id: Optional[uuid.UUID] = None # 允许为空
|
||
created_at: datetime
|
||
updated_at: datetime
|
||
|
||
@validator('id', 'owner_id', pre=True)
|
||
def validate_uuid(cls, v):
|
||
"""转换为标准UUID"""
|
||
if v is None:
|
||
return None
|
||
return uuid.UUID(str(v))
|
||
|
||
|
||
# ========== Agent相关 ==========
|
||
|
||
class K8sResourceConfig(BaseModel):
|
||
"""Kubernetes 资源配置"""
|
||
cpu_request: str = Field(default="100m", description="CPU 请求量(如 100m, 1)")
|
||
cpu_limit: str = Field(default="500m", description="CPU 限制量(如 500m, 2)")
|
||
memory_request: str = Field(default="128Mi", description="内存请求量(如 128Mi, 1Gi)")
|
||
memory_limit: str = Field(default="512Mi", description="内存限制量(如 512Mi, 2Gi)")
|
||
replicas: int = Field(default=1, ge=1, le=10, description="副本数量")
|
||
env: Dict[str, str] = Field(default_factory=dict, description="环境变量")
|
||
|
||
|
||
class AgentCreateRequest(BaseModel):
|
||
"""创建Agent请求"""
|
||
name: str = Field(..., min_length=1, max_length=63, description="Agent 名称(小写字母、数字、连字符)")
|
||
description: Optional[str] = None
|
||
role: str = "general-purpose agent"
|
||
goal: str = "Handle generic MCP tasks and routing"
|
||
|
||
# 模板配置(用于 K8s 部署)
|
||
template: Optional[str] = Field(None, description="Agent 模板类型(如 echo_agent, jina_search_agent)")
|
||
|
||
tools: List[str] = [] # 工具名称列表
|
||
config: Dict[str, Any] = {}
|
||
capabilities: List[str] = []
|
||
|
||
# K8s 资源配置
|
||
resource_config: Optional[K8sResourceConfig] = Field(None, description="Kubernetes 资源配置")
|
||
|
||
owner_id: Optional[Union[uuid.UUID, str]] = None
|
||
|
||
@validator('name')
|
||
def validate_name(cls, v):
|
||
"""验证名称符合 K8s 命名规范"""
|
||
import re
|
||
if not re.match(r'^[a-z0-9]([-a-z0-9]*[a-z0-9])?$', v):
|
||
raise ValueError('名称必须是小写字母、数字和连字符,且以字母或数字开头和结尾')
|
||
return v
|
||
|
||
|
||
class AgentUpdateRequest(BaseModel):
|
||
"""更新Agent请求"""
|
||
name: Optional[str] = None
|
||
description: Optional[str] = None
|
||
role: Optional[str] = None
|
||
goal: Optional[str] = None
|
||
tools: Optional[List[str]] = None
|
||
config: Optional[Dict[str, Any]] = None
|
||
capabilities: Optional[List[str]] = None
|
||
|
||
# K8s 资源配置更新
|
||
resource_config: Optional[K8sResourceConfig] = None
|
||
|
||
|
||
class AgentCard(BaseSchema):
|
||
"""Agent卡片信息"""
|
||
id: uuid.UUID
|
||
name: str
|
||
description: Optional[str] = None
|
||
role: Optional[str] = "general-purpose agent"
|
||
goal: Optional[str] = "Handle generic MCP tasks"
|
||
|
||
tools: List[str] = []
|
||
capabilities: List[str] = []
|
||
|
||
# 端点信息
|
||
endpoints: Dict[str, str] = {}
|
||
|
||
# 状态信息
|
||
status: str = "active"
|
||
version: str = "1.0.0"
|
||
|
||
# K8s 相关信息
|
||
template: Optional[str] = None
|
||
pod_name: Optional[str] = None
|
||
pod_ip: Optional[str] = None
|
||
k8s_status: Optional[str] = None
|
||
service_port: Optional[int] = None
|
||
access_url: Optional[str] = None
|
||
|
||
# 资源配置
|
||
cpu_request: Optional[str] = None
|
||
cpu_limit: Optional[str] = None
|
||
memory_request: Optional[str] = None
|
||
memory_limit: Optional[str] = None
|
||
|
||
# 统计信息
|
||
total_executions: int = 0
|
||
success_rate: float = 0.0
|
||
avg_execution_time: float = 0.0
|
||
|
||
# 时间信息
|
||
created_at: datetime
|
||
updated_at: datetime
|
||
pod_created_at: Optional[datetime] = None
|
||
|
||
@validator('id', pre=True)
|
||
def validate_id(cls, v):
|
||
"""转换id为标准UUID"""
|
||
if v is None:
|
||
return None
|
||
# 先转换为字符串,避免asyncpg UUID对象的问题
|
||
v_str = str(v)
|
||
# 然后转换为标准UUID
|
||
return uuid.UUID(v_str)
|
||
|
||
|
||
class AgentStatusResponse(BaseSchema):
|
||
"""Agent 状态响应"""
|
||
id: uuid.UUID
|
||
name: str
|
||
status: str # active, inactive, error
|
||
k8s_status: str # Pending, Running, Succeeded, Failed, Unknown
|
||
|
||
# Pod 信息
|
||
pod_name: Optional[str] = None
|
||
pod_ip: Optional[str] = None
|
||
node: Optional[str] = None
|
||
|
||
# 访问信息
|
||
service_port: Optional[int] = None
|
||
access_url: Optional[str] = None
|
||
endpoints: Dict[str, str] = {}
|
||
|
||
# 资源配置
|
||
cpu_request: Optional[str] = None
|
||
cpu_limit: Optional[str] = None
|
||
memory_request: Optional[str] = None
|
||
memory_limit: Optional[str] = None
|
||
|
||
# 时间信息
|
||
created_at: datetime
|
||
pod_created_at: Optional[datetime] = None
|
||
|
||
# 条件信息
|
||
conditions: Optional[List[Dict[str, Any]]] = None
|
||
|
||
|
||
class AgentMetricsResponse(BaseSchema):
|
||
"""Agent 资源使用响应"""
|
||
id: uuid.UUID
|
||
name: str
|
||
|
||
# 资源请求
|
||
requests: Dict[str, str] = {} # {"cpu": "100m", "memory": "128Mi"}
|
||
|
||
# 资源限制
|
||
limits: Dict[str, str] = {} # {"cpu": "500m", "memory": "512Mi"}
|
||
|
||
# 实际使用(如果可用)
|
||
usage: Optional[Dict[str, str]] = None
|
||
|
||
|
||
class TemplateInfo(BaseModel):
|
||
"""模板信息"""
|
||
template: str
|
||
port: Optional[int] = None
|
||
env_info: Dict[str, Any] = {}
|
||
|
||
|
||
class TemplateListResponse(BaseModel):
|
||
"""模板列表响应"""
|
||
templates: List[TemplateInfo]
|
||
count: int
|
||
type: Optional[str] = None # platform, custom, 或 None(表示所有类型)
|
||
|
||
|
||
class AgentExecution(BaseModel):
|
||
"""Agent执行请求"""
|
||
method: str
|
||
params: Optional[Dict[str, Any]] = None
|
||
timeout: Optional[int] = None
|
||
session_id: Optional[str] = None
|
||
|
||
|
||
class ExecutionResult(BaseSchema):
|
||
"""执行结果"""
|
||
execution_id: str
|
||
success: bool
|
||
result: Optional[Any] = None
|
||
error: Optional[str] = None
|
||
|
||
# 性能指标
|
||
execution_time: float = 0.0
|
||
cpu_usage: float = 0.0
|
||
memory_usage: float = 0.0
|
||
network_io: float = 0.0
|
||
|
||
# 成本信息
|
||
eu_consumed: float = 0.0
|
||
cost: float = 0.0
|
||
|
||
# 时间戳
|
||
started_at: datetime
|
||
completed_at: Optional[datetime] = None
|
||
|
||
|
||
# ========== 用户相关 ==========
|
||
|
||
class UserCreate(BaseModel):
|
||
"""创建用户请求"""
|
||
username: str = Field(..., min_length=3, max_length=50)
|
||
email: str = Field(..., pattern=r'^[^@]+@[^@]+\.[^@]+$')
|
||
password: str = Field(..., min_length=8)
|
||
full_name: Optional[str] = None
|
||
verification_code: str = Field(..., min_length=6, max_length=6, description="邮箱验证码")
|
||
|
||
|
||
class UserUpdate(BaseModel):
|
||
"""更新用户请求"""
|
||
email: Optional[str] = None
|
||
full_name: Optional[str] = None
|
||
is_active: Optional[bool] = None
|
||
|
||
|
||
class UserResponse(BaseSchema):
|
||
"""用户响应"""
|
||
id: uuid.UUID
|
||
username: str
|
||
email: str
|
||
full_name: Optional[str]
|
||
is_active: bool
|
||
is_admin: bool
|
||
created_at: datetime
|
||
updated_at: datetime
|
||
|
||
|
||
class UserLogin(BaseModel):
|
||
"""用户登录请求"""
|
||
username: str
|
||
password: str
|
||
|
||
|
||
class Token(BaseModel):
|
||
"""访问令牌"""
|
||
access_token: str
|
||
token_type: str = "bearer"
|
||
expires_in: int
|
||
|
||
|
||
# ========== 会话相关 ==========
|
||
|
||
class SessionCreate(BaseModel):
|
||
"""创建会话请求"""
|
||
context: Optional[Dict[str, Any]] = None
|
||
metadata: Optional[Dict[str, Any]] = None
|
||
|
||
|
||
class SessionResponse(BaseSchema):
|
||
"""会话响应"""
|
||
id: uuid.UUID
|
||
session_id: str
|
||
status: str
|
||
context: Dict[str, Any]
|
||
metadata: Dict[str, Any]
|
||
created_at: datetime
|
||
updated_at: datetime
|
||
|
||
|
||
# ========== 计费相关 ==========
|
||
|
||
class BillingRecord(BaseSchema):
|
||
"""计费记录"""
|
||
id: uuid.UUID
|
||
execution_id: uuid.UUID
|
||
eu_consumed: float
|
||
cost: float
|
||
currency: str
|
||
|
||
# 资源详情
|
||
cpu_time: float
|
||
memory_max: float
|
||
network_io: float
|
||
storage_io: float
|
||
|
||
created_at: datetime
|
||
|
||
|
||
class BillingSummary(BaseModel):
|
||
"""计费汇总"""
|
||
user_id: uuid.UUID
|
||
period_start: datetime
|
||
period_end: datetime
|
||
|
||
total_executions: int
|
||
total_eu_consumed: float
|
||
total_cost: float
|
||
currency: str
|
||
|
||
# 按服务分类
|
||
breakdown_by_agent: Dict[str, float] = {}
|
||
breakdown_by_tool: Dict[str, float] = {}
|
||
|
||
|
||
# ========== API响应包装 ==========
|
||
|
||
class APIResponse(BaseModel):
|
||
"""API响应包装"""
|
||
success: bool
|
||
message: str = ""
|
||
data: Optional[Any] = None
|
||
timestamp: datetime = Field(default_factory=datetime.utcnow)
|
||
|
||
|
||
class PaginatedResponse(BaseModel, Generic[T]):
|
||
"""分页响应"""
|
||
items: List[T]
|
||
total: int
|
||
page: int
|
||
page_size: int
|
||
pages: int
|
||
|
||
@property
|
||
def has_next(self) -> bool:
|
||
return self.page < self.pages
|
||
|
||
@property
|
||
def has_prev(self) -> bool:
|
||
return self.page > 1
|
||
|
||
|
||
# ========== 系统状态 ==========
|
||
|
||
class ServiceStatus(BaseModel):
|
||
"""服务状态"""
|
||
status: str
|
||
latency: int = 0
|
||
error: Optional[str] = None
|
||
code: Optional[int] = None
|
||
|
||
|
||
class HealthCheck(BaseModel):
|
||
"""健康检查响应"""
|
||
status: str
|
||
timestamp: datetime
|
||
services: Dict[str, ServiceStatus]
|
||
version: str = "1.0.0"
|
||
score: Optional[int] = None
|
||
uptime_seconds: Optional[float] = None
|
||
|
||
|
||
class SystemMetrics(BaseModel):
|
||
"""系统指标"""
|
||
timestamp: datetime
|
||
|
||
# 服务指标
|
||
active_agents: int
|
||
total_executions: int
|
||
success_rate: float
|
||
avg_response_time: float
|
||
|
||
# 资源指标
|
||
cpu_usage: float
|
||
memory_usage: float
|
||
disk_usage: float
|
||
|
||
# 业务指标
|
||
daily_active_users: int
|
||
total_eu_consumed: float
|
||
total_cost: float
|
||
|
||
|
||
# ========== 错误响应 ==========
|
||
|
||
class ErrorResponse(BaseModel):
|
||
"""错误响应"""
|
||
error: str
|
||
message: str
|
||
details: Optional[Dict[str, Any]] = None
|
||
timestamp: datetime = Field(default_factory=datetime.utcnow)
|
||
|
||
|
||
class ValidationError(BaseModel):
|
||
"""验证错误"""
|
||
field: str
|
||
message: str
|
||
invalid_value: Optional[Any] = None
|
||
|