186 lines
4.4 KiB
Python
186 lines
4.4 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""
|
||
面向 reflex.state 的类
|
||
"""
|
||
from enum import StrEnum
|
||
from typing import Any
|
||
from datetime import datetime
|
||
|
||
from pydantic import BaseModel, Field, TypeAdapter
|
||
from pydantic_ai import RunUsage
|
||
from pydantic_ai._uuid import uuid7
|
||
from pydantic_ai.usage import UsageLimits
|
||
|
||
|
||
class MessageType(StrEnum):
|
||
"""
|
||
消息类型枚举
|
||
"""
|
||
|
||
USER_PROMPT = "user_prompt"
|
||
THINKING = "thinking"
|
||
AGENT_RUN_RESULT_OUTPUT = "agent_run_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="消息内容")
|
||
is_expanded: bool = Field(
|
||
default=True, description="展开,True 表示展开,False 表示折叠"
|
||
)
|
||
|
||
|
||
class RunStatus(StrEnum):
|
||
"""
|
||
运行状态枚举
|
||
"""
|
||
|
||
RUNNING = "running"
|
||
FINISHED = "finished"
|
||
|
||
|
||
class Run(BaseModel):
|
||
"""
|
||
运行类
|
||
"""
|
||
|
||
id: str = Field(default_factory=lambda: str(uuid7()), description="运行唯一标识")
|
||
messages: dict[str, Message] = Field(default_factory=dict, description="消息字典")
|
||
status: RunStatus = Field(default=RunStatus.RUNNING, description="运行状态")
|
||
usage: RunUsage = Field(default=RunUsage(), description="运行使用量")
|
||
usage_limits: UsageLimits | None = Field(default=None, description="运行使用量限制")
|
||
is_thinking: bool = Field(
|
||
default=False, description="正在思考,True 表示正在思考,False 表示未正在思考"
|
||
)
|
||
|
||
|
||
class TaskType(StrEnum):
|
||
"""
|
||
任务类型枚举
|
||
"""
|
||
|
||
CHAT = "聊天"
|
||
BOOK_FLIGHT = "预定航班"
|
||
|
||
|
||
class TaskStatus(StrEnum):
|
||
"""
|
||
任务状态枚举
|
||
"""
|
||
|
||
NONE = "none"
|
||
|
||
|
||
class Deps(BaseModel):
|
||
"""
|
||
依赖项类
|
||
"""
|
||
|
||
...
|
||
|
||
|
||
class Task(BaseModel):
|
||
"""
|
||
任务类
|
||
"""
|
||
|
||
type: TaskType = Field(default=TaskType.CHAT, description="任务类型")
|
||
status: TaskStatus = Field(default=TaskStatus.NONE, description="任务状态")
|
||
deps: Deps | None = Field(default=None, description="任务依赖项")
|
||
usage: RunUsage = Field(default=RunUsage(), description="任务使用量")
|
||
usage_limits: UsageLimits | None = Field(
|
||
default=None,
|
||
description="任务使用量限制",
|
||
)
|
||
|
||
|
||
class Conversation(BaseModel):
|
||
"""
|
||
会话类
|
||
"""
|
||
|
||
id: str = Field(default_factory=lambda: str(uuid7()), description="会话唯一标识")
|
||
description: str = Field(default="新会话", description="会话描述")
|
||
usage: RunUsage = Field(default=RunUsage(), description="会话使用量")
|
||
messages: dict[str, Message] = Field(default_factory=dict, description="运行字典")
|
||
created_at: datetime = Field(..., description="会话创建日期时间")
|
||
|
||
|
||
class ConversationHistoryItem(BaseModel):
|
||
"""
|
||
会话历史项类
|
||
"""
|
||
|
||
description: str = Field(..., description="会话描述")
|
||
created_at: str = Field(..., description="会话创建日期时间")
|
||
|
||
|
||
# Dpes 适配器
|
||
DepsAdapter = TypeAdapter(Deps)
|
||
|
||
|
||
def deps_to_object(deps: dict[str, Any]) -> Deps | None:
|
||
"""
|
||
Deps 转为对象
|
||
"""
|
||
if not deps:
|
||
return None
|
||
return DepsAdapter.validate_python(deps)
|
||
|
||
|
||
def deps_to_dict(deps: Deps | None) -> dict[str, Any]:
|
||
"""
|
||
Deps 转为字典
|
||
"""
|
||
if not deps:
|
||
return {}
|
||
return DepsAdapter.dump_python(deps)
|
||
|
||
|
||
# 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_to_dict(usage: RunUsage) -> dict[str, Any]:
|
||
"""
|
||
Usage 转为字典
|
||
"""
|
||
return UsageAdapter.dump_python(usage)
|
||
|
||
|
||
# UsageLimits 适配器
|
||
UsageLimitsAdapter = TypeAdapter(UsageLimits)
|
||
|
||
|
||
def usage_limits_to_object(usage_limits: dict[str, Any]) -> UsageLimits | None:
|
||
"""
|
||
UsageLimits 转为对象
|
||
"""
|
||
if not usage_limits:
|
||
return None
|
||
return UsageLimitsAdapter.validate_python(usage_limits)
|
||
|
||
|
||
def usage_limits_to_dict(usage_limits: UsageLimits | None) -> dict[str, Any]:
|
||
"""
|
||
UsageLimits 转为字典
|
||
"""
|
||
if not usage_limits:
|
||
return {}
|
||
return UsageLimitsAdapter.dump_python(usage_limits)
|