163 lines
4.2 KiB
Python
163 lines
4.2 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""
|
||
面向 reflex.state 的类
|
||
"""
|
||
from datetime import datetime
|
||
from enum import StrEnum
|
||
from typing import Any
|
||
|
||
from pydantic import BaseModel, Field, TypeAdapter
|
||
from pydantic_ai import RunUsage
|
||
from pydantic_ai._uuid import uuid7
|
||
from pydantic_ai.usage import UsageLimits
|
||
|
||
|
||
class TaskStatus(StrEnum):
|
||
"""
|
||
任务状态枚举
|
||
"""
|
||
|
||
RUNNING = "running"
|
||
DONE = "done"
|
||
PENDING_APPROVAL = "pending_approval"
|
||
ERROR = "error"
|
||
|
||
|
||
class Task(BaseModel):
|
||
"""
|
||
任务类
|
||
"""
|
||
|
||
id: str = Field(default_factory=lambda: str(uuid7()), description="任务唯一标识")
|
||
tool_name: str = Field(..., description="工具名称")
|
||
tool_call_id: str = Field(..., description="工具调用唯一标识")
|
||
status: TaskStatus = Field(default=TaskStatus.RUNNING, description="任务状态")
|
||
title: str = Field(default="", description="任务标题")
|
||
content: str = Field(default="", description="任务内容")
|
||
|
||
|
||
class MessageType(StrEnum):
|
||
"""
|
||
消息类型枚举
|
||
"""
|
||
|
||
USER_PROMPT = "user_prompt"
|
||
THINKING = "thinking"
|
||
WORK = "work"
|
||
RESULT_OUTPUT = "result_output"
|
||
TEXT = "text"
|
||
|
||
|
||
class Message(BaseModel):
|
||
"""
|
||
消息类
|
||
"""
|
||
|
||
id: str = Field(default_factory=lambda: str(uuid7()), description="消息唯一标识")
|
||
type: MessageType = Field(..., description="消息类型")
|
||
title: str = Field(default="", description="消息标题")
|
||
content: str = Field(default="", description="消息内容")
|
||
tasks: list[Task] = Field(default_factory=list, description="任务列表")
|
||
is_running: bool = Field(
|
||
default=False, description="正在运行,True 表示正在运行, False 表示运行结束"
|
||
)
|
||
is_shown: bool = Field(
|
||
default=False, description="展示组件,True 表示展示,False 表示隐藏"
|
||
)
|
||
|
||
|
||
class AgentRunWarmUpEvent(BaseModel):
|
||
"""
|
||
智能体运行预热事件类
|
||
"""
|
||
|
||
content: str = Field(
|
||
default="",
|
||
description="预热内容",
|
||
)
|
||
|
||
|
||
class WorkType(StrEnum):
|
||
"""
|
||
工作类型枚举
|
||
"""
|
||
|
||
BOOK_FLIGHT = "预定航班"
|
||
|
||
|
||
class Work(BaseModel):
|
||
"""
|
||
工作流类
|
||
"""
|
||
|
||
type: WorkType = Field(..., description="工作类型")
|
||
deps: Any = Field(default=None, description="工作依赖项")
|
||
usage: RunUsage = Field(default=RunUsage(), description="工作使用量")
|
||
usage_limits: UsageLimits | None = Field(
|
||
default=None,
|
||
description="工作使用量限制",
|
||
)
|
||
|
||
|
||
class Conversation(BaseModel):
|
||
"""
|
||
会话类
|
||
"""
|
||
|
||
id: str = Field(..., description="会话唯一标识")
|
||
description: str = Field(..., description="会话描述")
|
||
work: Work | None = Field(default=None, description="工作")
|
||
user_prompt: str = Field(default="", description="用户提示词")
|
||
usage: dict[str, Any] = Field(..., description="会话使用量")
|
||
messages: dict[str, Message] = Field(..., description="消息字典")
|
||
created_at: datetime = Field(..., description="会话创建日期时间")
|
||
is_running: bool = Field(
|
||
default=False, description="正在运行,True 表示正在运行, False 表示运行结束"
|
||
)
|
||
awaiting_stream: bool = Field(
|
||
default=False,
|
||
description="等待流式输出,True 表示等待流式输出,False 表示已开始流式输出或已完成",
|
||
)
|
||
|
||
|
||
class ConversationHistoryItem(BaseModel):
|
||
"""
|
||
会话历史项类
|
||
"""
|
||
|
||
id: str = Field(..., description="会话唯一标识")
|
||
description: str = Field(..., description="会话描述")
|
||
created_at: str = Field(..., description="会话创建日期时间")
|
||
|
||
|
||
class MessageHistoryItem(Message):
|
||
"""
|
||
消息历史项类
|
||
"""
|
||
|
||
|
||
# Usage 适配器
|
||
UsageAdapter = TypeAdapter(RunUsage)
|
||
|
||
|
||
def usage_validate_python(usage: dict[str, Any]) -> RunUsage:
|
||
"""
|
||
将 Usage 反序列化
|
||
"""
|
||
if not usage:
|
||
return RunUsage()
|
||
return UsageAdapter.validate_python(usage)
|
||
|
||
|
||
def usage_dump_python(usage: RunUsage) -> dict[str, Any]:
|
||
"""
|
||
将 Usage 序列化
|
||
:param usage: 使用量
|
||
:return: python 字典
|
||
"""
|
||
return UsageAdapter.dump_python(usage)
|
||
|
||
|
||
# UsageLimits 适配器
|
||
UsageLimitsAdapter = TypeAdapter(UsageLimits)
|