145 lines
3.9 KiB
Python
145 lines
3.9 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 MessageType(StrEnum):
|
||
"""
|
||
消息类型枚举
|
||
"""
|
||
|
||
USER_PROMPT = "user_prompt"
|
||
THINKING = "thinking"
|
||
TOOL_CALL = "tool_call"
|
||
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="消息内容")
|
||
is_running: bool = Field(
|
||
default=False, description="正在运行,True 表示正在运行,False 表示运行完成"
|
||
)
|
||
is_shown: bool = Field(
|
||
default=False, description="展示组件,True 表示展示,False 表示隐藏"
|
||
)
|
||
|
||
|
||
class WorkFlowType(StrEnum):
|
||
"""
|
||
工作流类型枚举
|
||
"""
|
||
|
||
BOOK_FLIGHT = "预定航班"
|
||
|
||
|
||
class Deps(BaseModel):
|
||
"""
|
||
依赖项类
|
||
"""
|
||
|
||
pass
|
||
|
||
|
||
class WorkFlow(BaseModel):
|
||
"""
|
||
工作流类
|
||
"""
|
||
|
||
type: WorkFlowType = 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_flow: WorkFlow | 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(BaseModel):
|
||
"""
|
||
消息历史项类
|
||
"""
|
||
|
||
id: str = Field(..., description="消息唯一标识")
|
||
type: MessageType = Field(..., description="消息类型")
|
||
title: str = Field(default="", description="消息标题")
|
||
content: str = Field(default="", description="消息内容")
|
||
is_running: bool = Field(
|
||
default=False, description="正在运行,True 表示正在运行,False 表示运行完成"
|
||
)
|
||
is_shown: bool = Field(
|
||
default=False, description="展示组件,True 表示展示,False 表示隐藏"
|
||
)
|
||
|
||
|
||
# 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)
|