152 lines
3.9 KiB
Python
152 lines
3.9 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""
|
||
面向 reflex.state 的类
|
||
"""
|
||
from datetime import datetime
|
||
from typing import Any
|
||
from pydantic import BaseModel, Field, TypeAdapter
|
||
from pydantic_ai import RunUsage
|
||
from pydantic_ai._uuid import uuid7
|
||
from enum import StrEnum
|
||
from pydantic_ai.usage import UsageLimits
|
||
|
||
|
||
class Thought(BaseModel):
|
||
"""
|
||
思考类
|
||
"""
|
||
|
||
type: str = Field(..., description="思考类型")
|
||
content: str = Field(..., description="思考内容")
|
||
|
||
|
||
# 思考领域模型类型适配器
|
||
ThoughtsAdapter = TypeAdapter(dict[int, Thought])
|
||
|
||
|
||
def thoughts_to_dict(thoughts: dict[int, Thought]) -> dict:
|
||
"""
|
||
Thought 转为字典
|
||
"""
|
||
return ThoughtsAdapter.dump_python(thoughts)
|
||
|
||
|
||
class Dialog(BaseModel):
|
||
"""
|
||
对话类
|
||
"""
|
||
|
||
id: str = Field(default_factory=lambda: str(uuid7()), description="对话唯一标识")
|
||
user_prompt: str = Field(default="", description="用户提示词")
|
||
thoughts: dict[int, Thought] = Field(default_factory=dict, description="思考列表")
|
||
result_output: str = Field(default="", description="结果输出")
|
||
usage: dict[str, Any] = Field(default_factory=dict, description="对话使用量")
|
||
is_thinking: bool = Field(
|
||
default=False, description="正在思考,True 表示正在思考,False 表示未正在思考"
|
||
)
|
||
is_expanded: bool = Field(
|
||
default=False, description="思考折叠面板展开状态,True 表示展开,False 表示折叠"
|
||
)
|
||
|
||
|
||
# Usage 适配器
|
||
UsageAdapter = TypeAdapter(RunUsage)
|
||
|
||
|
||
def usage_to_object(usage: dict) -> RunUsage:
|
||
"""
|
||
Usage 转为对象
|
||
"""
|
||
return UsageAdapter.validate_python(usage) if usage else RunUsage()
|
||
|
||
|
||
def usage_to_dict(usage: RunUsage) -> dict:
|
||
"""
|
||
Usage 转为字典
|
||
"""
|
||
return UsageAdapter.dump_python(usage)
|
||
|
||
|
||
# UsageLimits 适配器
|
||
UsageLimitsAdapter = TypeAdapter(UsageLimits)
|
||
|
||
|
||
def usage_limits_to_object(usage_limits: dict) -> UsageLimits:
|
||
"""
|
||
UsageLimits 转为对象
|
||
"""
|
||
return (
|
||
UsageLimitsAdapter.validate_python(usage_limits)
|
||
if usage_limits
|
||
else UsageLimits()
|
||
)
|
||
|
||
|
||
def usage_limits_to_dict(usage_limits: UsageLimits) -> dict:
|
||
"""
|
||
UsageLimits 转为字典
|
||
"""
|
||
return UsageLimitsAdapter.dump_python(usage_limits)
|
||
|
||
|
||
class TaskType(StrEnum):
|
||
"""
|
||
任务类型枚举
|
||
"""
|
||
|
||
NONE = "none"
|
||
BOOK_FLIGHT = "book_flight"
|
||
|
||
|
||
class TaskNode(StrEnum):
|
||
"""
|
||
任务节点枚举
|
||
"""
|
||
|
||
EXECUTION = "execution"
|
||
PENDING_INPUT = "pending_input"
|
||
FINISH = "finish"
|
||
|
||
|
||
class Task(BaseModel):
|
||
"""
|
||
任务类
|
||
"""
|
||
|
||
id: str = Field(default_factory=lambda: str(uuid7()), description="任务唯一标识")
|
||
type: TaskType = Field(..., description="任务类型")
|
||
node: TaskNode = Field(default=TaskNode.EXECUTION, description="任务节点")
|
||
deps: Any = Field(default=None, description="任务依赖项")
|
||
usage: dict[str, Any] = Field(default_factory=dict, description="任务使用量")
|
||
usage_limits: dict[str, Any] = Field(
|
||
default_factory=dict, description="任务使用量限制"
|
||
)
|
||
|
||
|
||
class TaskResultEvent(BaseModel):
|
||
"""
|
||
任务结果事件类
|
||
"""
|
||
|
||
task: Task = Field(..., description="任务实例")
|
||
content: str = Field(default="", description="任务结果内容")
|
||
|
||
|
||
class Conversation(BaseModel):
|
||
"""
|
||
会话类
|
||
"""
|
||
|
||
id: str = Field(default_factory=lambda: str(uuid7()), description="会话唯一标识")
|
||
description: str = Field(default="新会话", description="会话描述")
|
||
is_running: bool = Field(
|
||
default=False,
|
||
description="会话正在运行,True 表示正在运行,False 表示未正在运行",
|
||
)
|
||
dialogs: dict[str, Dialog] = Field(default_factory=dict, description="对话字典")
|
||
task: Task | None = Field(default=None, description="会话任务")
|
||
created_at: str = Field(
|
||
...,
|
||
description="创建时间",
|
||
)
|