206 lines
5.4 KiB
Python
206 lines
5.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"
|
||
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 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 WorkFlowType(StrEnum):
|
||
"""
|
||
工作流类型枚举
|
||
"""
|
||
|
||
BOOK_FLIGHT = "预定航班"
|
||
|
||
|
||
class Deps(BaseModel):
|
||
"""
|
||
依赖项类
|
||
"""
|
||
|
||
...
|
||
|
||
|
||
class WorkFlow(BaseModel):
|
||
"""
|
||
工作流类
|
||
"""
|
||
|
||
type: WorkFlowType = Field(..., 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(..., 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 表示隐藏"
|
||
)
|
||
|
||
|
||
# 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_dump_python(usage: RunUsage) -> dict[str, Any]:
|
||
"""
|
||
将 Usage 序列化
|
||
:param usage: 使用量
|
||
:return: python 字典
|
||
"""
|
||
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)
|