296 lines
7.9 KiB
Python
296 lines
7.9 KiB
Python
"""
|
||||
|
|
对话状态管理模块
|
|||
|
|
管理对话历史、上下文和当前任务状态
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
import logging
|
|||
|
|
from dataclasses import dataclass, field
|
|||
|
|
from datetime import datetime
|
|||
|
|
from enum import Enum
|
|||
|
|
from typing import Dict, Any, List, Optional, Union
|
|||
|
|
|
|||
|
|
|
|||
|
|
class TaskStatus(Enum):
|
|||
|
|
"""任务状态"""
|
|||
|
|
|
|||
|
|
IDLE = "idle" # 空闲
|
|||
|
|
PROCESSING = "processing" # 处理中
|
|||
|
|
COMPLETED = "completed" # 已完成
|
|||
|
|
FAILED = "failed" # 失败
|
|||
|
|
WAITING_INPUT = "waiting_input" # 等待用户输入
|
|||
|
|
|
|||
|
|
|
|||
|
|
@dataclass
|
|||
|
|
class Message:
|
|||
|
|
"""对话消息"""
|
|||
|
|
|
|||
|
|
role: str # "user" 或 "assistant"
|
|||
|
|
content: str
|
|||
|
|
timestamp: str = field(default_factory=lambda: datetime.now().isoformat())
|
|||
|
|
metadata: Dict[str, Any] = field(default_factory=dict)
|
|||
|
|
|
|||
|
|
def __post_init__(self) -> None:
|
|||
|
|
"""Validate message after initialization"""
|
|||
|
|
if self.role not in ["user", "assistant"]:
|
|||
|
|
raise ValueError("role must be 'user' or 'assistant'")
|
|||
|
|
if not self.content or not self.content.strip():
|
|||
|
|
raise ValueError("content cannot be empty")
|
|||
|
|
|
|||
|
|
|
|||
|
|
@dataclass
|
|||
|
|
class Task:
|
|||
|
|
"""当前任务"""
|
|||
|
|
|
|||
|
|
command: str
|
|||
|
|
parameters: Dict[str, Any]
|
|||
|
|
status: TaskStatus = TaskStatus.IDLE
|
|||
|
|
result: Optional[Any] = None
|
|||
|
|
error: Optional[str] = None
|
|||
|
|
started_at: Optional[str] = None
|
|||
|
|
completed_at: Optional[str] = None
|
|||
|
|
|
|||
|
|
def __post_init__(self) -> None:
|
|||
|
|
"""Validate task after initialization"""
|
|||
|
|
if not self.command or not self.command.strip():
|
|||
|
|
raise ValueError("command cannot be empty")
|
|||
|
|
if not isinstance(self.parameters, dict):
|
|||
|
|
raise ValueError("parameters must be a dictionary")
|
|||
|
|
|
|||
|
|
|
|||
|
|
class ConversationState:
|
|||
|
|
"""对话状态管理器"""
|
|||
|
|
|
|||
|
|
def __init__(self, max_history: int = 20) -> None:
|
|||
|
|
"""
|
|||
|
|
初始化对话状态
|
|||
|
|
|
|||
|
|
Args:
|
|||
|
|
max_history: 保留的最大历史消息数
|
|||
|
|
"""
|
|||
|
|
self.logger: logging.Logger = logging.getLogger("ConversationState")
|
|||
|
|
self.max_history: int = max_history
|
|||
|
|
|
|||
|
|
# 对话历史
|
|||
|
|
self.messages: List[Message] = []
|
|||
|
|
|
|||
|
|
# 当前任务
|
|||
|
|
self.current_task: Optional[Task] = None
|
|||
|
|
|
|||
|
|
# 用户偏好和上下文
|
|||
|
|
self.user_preferences: Dict[str, Any] = {}
|
|||
|
|
self.context: Dict[str, Any] = {}
|
|||
|
|
|
|||
|
|
# 统计信息
|
|||
|
|
self.stats: Dict[str, Union[int, str]] = {
|
|||
|
|
"total_messages": 0,
|
|||
|
|
"total_tasks": 0,
|
|||
|
|
"successful_tasks": 0,
|
|||
|
|
"failed_tasks": 0,
|
|||
|
|
"session_start": datetime.now().isoformat(),
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
def add_message(
|
|||
|
|
self, role: str, content: str, metadata: Optional[Dict[str, Any]] = None
|
|||
|
|
) -> Message:
|
|||
|
|
"""
|
|||
|
|
添加消息到历史
|
|||
|
|
|
|||
|
|
Args:
|
|||
|
|
role: 消息角色("user" 或 "assistant")
|
|||
|
|
content: 消息内容
|
|||
|
|
metadata: 消息元数据
|
|||
|
|
|
|||
|
|
Returns:
|
|||
|
|
Message: 添加的消息
|
|||
|
|
"""
|
|||
|
|
message: Message = Message(role=role, content=content, metadata=metadata or {})
|
|||
|
|
|
|||
|
|
self.messages.append(message)
|
|||
|
|
self.stats["total_messages"] += 1
|
|||
|
|
|
|||
|
|
# 保持历史长度在限制内
|
|||
|
|
if len(self.messages) > self.max_history:
|
|||
|
|
self.messages.pop(0)
|
|||
|
|
|
|||
|
|
self.logger.debug(f"添加消息: {role} - {content[:50]}...")
|
|||
|
|
|
|||
|
|
return message
|
|||
|
|
|
|||
|
|
def get_recent_messages(self, count: int = 5) -> List[Message]:
|
|||
|
|
"""
|
|||
|
|
获取最近的 N 条消息
|
|||
|
|
|
|||
|
|
Args:
|
|||
|
|
count: 消息数量
|
|||
|
|
|
|||
|
|
Returns:
|
|||
|
|
最近的消息列表
|
|||
|
|
"""
|
|||
|
|
return self.messages[-count:]
|
|||
|
|
|
|||
|
|
def get_conversation_history(self) -> List[Dict[str, str]]:
|
|||
|
|
"""
|
|||
|
|
获取对话历史(用于 Claude API)
|
|||
|
|
|
|||
|
|
Returns:
|
|||
|
|
对话历史列表
|
|||
|
|
"""
|
|||
|
|
return [{"role": msg.role, "content": msg.content} for msg in self.messages]
|
|||
|
|
|
|||
|
|
def start_task(self, command: str, parameters: Dict[str, Any]) -> Task:
|
|||
|
|
"""
|
|||
|
|
开始一个新任务
|
|||
|
|
|
|||
|
|
Args:
|
|||
|
|
command: 命令名称
|
|||
|
|
parameters: 命令参数
|
|||
|
|
|
|||
|
|
Returns:
|
|||
|
|
Task: 创建的任务
|
|||
|
|
"""
|
|||
|
|
self.current_task = Task(
|
|||
|
|
command=command,
|
|||
|
|
parameters=parameters,
|
|||
|
|
status=TaskStatus.PROCESSING,
|
|||
|
|
started_at=datetime.now().isoformat(),
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
self.stats["total_tasks"] += 1
|
|||
|
|
self.logger.info(f"开始任务: {command} - {parameters}")
|
|||
|
|
|
|||
|
|
return self.current_task
|
|||
|
|
|
|||
|
|
def complete_task(self, result: Any) -> Optional[Task]:
|
|||
|
|
"""
|
|||
|
|
完成当前任务
|
|||
|
|
|
|||
|
|
Args:
|
|||
|
|
result: 任务结果
|
|||
|
|
|
|||
|
|
Returns:
|
|||
|
|
Task: 完成的任务
|
|||
|
|
"""
|
|||
|
|
if not self.current_task:
|
|||
|
|
self.logger.warning("没有正在进行的任务")
|
|||
|
|
return None
|
|||
|
|
|
|||
|
|
self.current_task.status = TaskStatus.COMPLETED
|
|||
|
|
self.current_task.result = result
|
|||
|
|
self.current_task.completed_at = datetime.now().isoformat()
|
|||
|
|
self.stats["successful_tasks"] += 1
|
|||
|
|
|
|||
|
|
self.logger.info(f"任务完成: {self.current_task.command}")
|
|||
|
|
|
|||
|
|
return self.current_task
|
|||
|
|
|
|||
|
|
def fail_task(self, error: str) -> Optional[Task]:
|
|||
|
|
"""
|
|||
|
|
标记任务失败
|
|||
|
|
|
|||
|
|
Args:
|
|||
|
|
error: 错误信息
|
|||
|
|
|
|||
|
|
Returns:
|
|||
|
|
Task: 失败的任务
|
|||
|
|
"""
|
|||
|
|
if not self.current_task:
|
|||
|
|
self.logger.warning("没有正在进行的任务")
|
|||
|
|
return None
|
|||
|
|
|
|||
|
|
self.current_task.status = TaskStatus.FAILED
|
|||
|
|
self.current_task.error = error
|
|||
|
|
self.current_task.completed_at = datetime.now().isoformat()
|
|||
|
|
self.stats["failed_tasks"] += 1
|
|||
|
|
|
|||
|
|
self.logger.error(f"任务失败: {self.current_task.command} - {error}")
|
|||
|
|
|
|||
|
|
return self.current_task
|
|||
|
|
|
|||
|
|
def set_context(self, key: str, value: Any) -> None:
|
|||
|
|
"""
|
|||
|
|
设置上下文信息
|
|||
|
|
|
|||
|
|
Args:
|
|||
|
|
key: 上下文键
|
|||
|
|
value: 上下文值
|
|||
|
|
"""
|
|||
|
|
self.context[key] = value
|
|||
|
|
self.logger.debug(f"设置上下文: {key} = {value}")
|
|||
|
|
|
|||
|
|
def get_context(self, key: str, default: Any = None) -> Any:
|
|||
|
|
"""
|
|||
|
|
获取上下文信息
|
|||
|
|
|
|||
|
|
Args:
|
|||
|
|
key: 上下文键
|
|||
|
|
default: 默认值
|
|||
|
|
|
|||
|
|
Returns:
|
|||
|
|
上下文值
|
|||
|
|
"""
|
|||
|
|
return self.context.get(key, default)
|
|||
|
|
|
|||
|
|
def set_preference(self, key: str, value: Any) -> None:
|
|||
|
|
"""
|
|||
|
|
设置用户偏好
|
|||
|
|
|
|||
|
|
Args:
|
|||
|
|
key: 偏好键
|
|||
|
|
value: 偏好值
|
|||
|
|
"""
|
|||
|
|
self.user_preferences[key] = value
|
|||
|
|
self.logger.debug(f"设置偏好: {key} = {value}")
|
|||
|
|
|
|||
|
|
def get_preference(self, key: str, default: Any = None) -> Any:
|
|||
|
|
"""
|
|||
|
|
获取用户偏好
|
|||
|
|
|
|||
|
|
Args:
|
|||
|
|
key: 偏好键
|
|||
|
|
default: 默认值
|
|||
|
|
|
|||
|
|
Returns:
|
|||
|
|
偏好值
|
|||
|
|
"""
|
|||
|
|
return self.user_preferences.get(key, default)
|
|||
|
|
|
|||
|
|
def get_summary(self) -> Dict[str, Any]:
|
|||
|
|
"""
|
|||
|
|
获取对话状态摘要
|
|||
|
|
|
|||
|
|
Returns:
|
|||
|
|
状态摘要字典
|
|||
|
|
"""
|
|||
|
|
return {
|
|||
|
|
"total_messages": len(self.messages),
|
|||
|
|
"recent_messages": [
|
|||
|
|
{
|
|||
|
|
"role": msg.role,
|
|||
|
|
"content": msg.content[:100],
|
|||
|
|
"timestamp": msg.timestamp,
|
|||
|
|
}
|
|||
|
|
for msg in self.get_recent_messages(3)
|
|||
|
|
],
|
|||
|
|
"current_task": {
|
|||
|
|
"command": self.current_task.command,
|
|||
|
|
"status": self.current_task.status.value,
|
|||
|
|
"started_at": self.current_task.started_at,
|
|||
|
|
}
|
|||
|
|
if self.current_task
|
|||
|
|
else None,
|
|||
|
|
"stats": self.stats,
|
|||
|
|
"preferences": self.user_preferences,
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
def clear_history(self) -> None:
|
|||
|
|
"""清除对话历史"""
|
|||
|
|
self.messages.clear()
|
|||
|
|
self.logger.info("对话历史已清除")
|
|||
|
|
|
|||
|
|
def reset(self) -> None:
|
|||
|
|
"""重置对话状态"""
|
|||
|
|
self.messages.clear()
|
|||
|
|
self.current_task = None
|
|||
|
|
self.context.clear()
|
|||
|
|
self.logger.info("对话状态已重置")
|