Python/agent/application/states/models.py

150 lines
3.9 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

# -*- 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"
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"
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 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)