- Add core agent architecture with Command + Skill pattern - Implement Claude API integration for content analysis - Add Obsidian REST API integration for vault operations - Create conversational interface (v2.0) with natural language processing - Add comprehensive configuration management and validation - Include project documentation and developer guides - Set up testing framework with unit, integration, and property tests - Add Kiro specs for Claude API configuration and code quality improvements - Configure project steering files for development guidelines
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("对话状态已重置")
|