模块6:记忆系统迁移方案
1. 现状分析
1.1 当前LangGraph记忆系统
# 当前LangGraph记忆实现
from typing import TypedDict, List, Dict, Any, Optional
from langchain.memory import ConversationBufferMemory
from langchain.schema import BaseMessage, HumanMessage, AIMessage
class TradingState(TypedDict):
"""交易状态"""
messages: List[BaseMessage]
memory: ConversationBufferMemory
context: Dict[str, Any]
history: List[Dict[str, Any]]
# 当前记忆使用方式
def create_trading_memory():
return ConversationBufferMemory(
memory_key="history",
return_messages=True,
output_key="output"
)
def add_to_memory(state: TradingState, message: str, role: str = "human"):
"""添加消息到记忆"""
if role == "human":
state["memory"].chat_memory.add_user_message(message)
else:
state["memory"].chat_memory.add_ai_message(message)
# 同时更新状态
state["messages"].append(
HumanMessage(content=message) if role == "human" else AIMessage(content=message)
)
def get_memory_context(state: TradingState) -> str:
"""获取记忆上下文"""
return state["memory"].load_memory_variables({})["history"]
1.2 当前记忆系统特点
1. 简单键值存储:使用ConversationBufferMemory 2. 消息队列管理:维护消息历史 3. 上下文提取:加载历史对话 4. 无持久化:内存中存储,重启丢失 5. 无语义检索:只能按时间顺序检索 6. 无分层管理:所有记忆混在一起
1.3 当前记忆系统问题
1. 容量限制:内存存储,容量有限 2. 检索效率低:线性扫描历史 3. 上下文丢失:长对话会截断 4. 无智能筛选:不能区分重要信息 5. 重启丢失:服务重启后记忆消失 6. 无跨会话:不能跨对话会话
2. Agno记忆系统架构设计
2.1 核心记忆模型
from pydantic import BaseModel, Field
from typing import List, Dict, Any, Optional, Set
from datetime import datetime
from enum import Enum
import json
import numpy as np
from abc import ABC, abstractmethod
class MemoryType(str, Enum):
"""记忆类型"""
CONVERSATION = "conversation" # 对话记忆
WORKING = "working" # 工作记忆
EPISODIC = "episodic" # 情景记忆
SEMANTIC = "semantic" # 语义记忆
PROCEDURAL = "procedural" # 程序记忆
META = "meta" # 元记忆
class MemoryPriority(str, Enum):
"""记忆优先级"""
CRITICAL = "critical" # 关键信息
HIGH = "high" # 重要信息
NORMAL = "normal" # 普通信息
LOW = "low" # 低优先级
class MemoryStatus(str, Enum):
"""记忆状态"""
ACTIVE = "active" # 活跃状态
STABLE = "stable" # 稳定状态
ARCHIVED = "archived" # 归档状态
FORGOTTEN = "forgotten" # 遗忘状态
class MemoryEntry(BaseModel):
"""记忆条目"""
id: str = Field(..., description="记忆ID")
type: MemoryType = Field(..., description="记忆类型")
content: Dict[str, Any] = Field(..., description="记忆内容")
embedding: Optional[List[float]] = Field(None, description="向量嵌入")
keywords: Set[str] = Field(default_factory=set, description="关键词")
entities: Set[str] = Field(default_factory=set, description="实体")
priority: MemoryPriority = Field(default=MemoryPriority.NORMAL, description="优先级")
status: MemoryStatus = Field(default=MemoryStatus.ACTIVE, description="状态")
timestamp: datetime = Field(default_factory=datetime.now, description="时间戳")
access_count: int = Field(default=0, description="访问次数")
last_accessed: Optional[datetime] = Field(None, description="最后访问时间")
metadata: Dict[str, Any] = Field(default_factory=dict, description="元数据")
parent_id: Optional[str] = Field(None, description="父记忆ID")
child_ids: Set[str] = Field(default_factory=set, description="子记忆ID")
session_id: str = Field(..., description="会话ID")
user_id: Optional[str] = Field(None, description="用户ID")
class Config:
arbitrary_types_allowed = True
class MemoryStats(BaseModel):
"""记忆统计"""
total_entries: int = Field(default=0, description="总记忆数")
type_distribution: Dict[MemoryType, int] = Field(default_factory=dict, description="类型分布")
priority_distribution: Dict[MemoryPriority, int] = Field(default_factory=dict, description="优先级分布")
status_distribution: Dict[MemoryStatus, int] = Field(default_factory=dict, description="状态分布")
total_size_bytes: int = Field(default=0, description="总大小")
avg_embedding_size: float = Field(default=0.0, description="平均嵌入大小")
oldest_entry: Optional[datetime] = Field(None, description="最早记忆")
newest_entry: Optional[datetime] = Field(None, description="最新记忆")
most_accessed: Optional[str] = Field(None, description="最常访问")
class MemoryQuery(BaseModel):
"""记忆查询"""
query: str = Field(..., description="查询文本")
type_filter: Optional[MemoryType] = Field(None, description="类型过滤")
priority_filter: Optional[MemoryPriority] = Field(None, description="优先级过滤")
status_filter: Optional[MemoryStatus] = Field(None, description="状态过滤")
time_range: Optional[tuple[datetime, datetime]] = Field(None, description="时间范围")
session_filter: Optional[str] = Field(None, description="会话过滤")
user_filter: Optional[str] = Field(None, description="用户过滤")
max_results: int = Field(default=10, description="最大结果数")
min_similarity: float = Field(default=0.7, description="最小相似度")
include_embeddings: bool = Field(default=False, description="包含嵌入")
include_metadata: bool = Field(default=True, description="包含元数据")
2.2 记忆存储接口
from abc import ABC, abstractmethod
class MemoryStorage(ABC):
"""记忆存储接口"""
@abstractmethod
async def store(self, entry: MemoryEntry) -> bool:
"""存储记忆"""
pass
@abstractmethod
async def retrieve(self, memory_id: str) -> Optional[MemoryEntry]:
"""检索记忆"""
pass
@abstractmethod
async def update(self, entry: MemoryEntry) -> bool:
"""更新记忆"""
pass
@abstractmethod
async def delete(self, memory_id: str) -> bool:
"""删除记忆"""
pass
@abstractmethod
async def search(self, query: MemoryQuery) -> List[MemoryEntry]:
"""搜索记忆"""
pass
@abstractmethod
async def get_stats(self) -> MemoryStats:
"""获取统计信息"""
pass
@abstractmethod
async def clear(self) -> bool:
"""清空记忆"""
pass
class VectorMemoryStorage(MemoryStorage):
"""向量记忆存储"""
def __init__(self, embedding_model: str = "text-embedding-3-small"):
self.embedding_model = embedding_model
self.memories: Dict[str, MemoryEntry] = {}
self.vector_index: Dict[str, np.ndarray] = {}
self.type_index: Dict[MemoryType, Set[str]] = {}
self.priority_index: Dict[MemoryPriority, Set[str]] = {}
self.status_index: Dict[MemoryStatus, Set[str]] = {}
self.keyword_index: Dict[str, Set[str]] = {}
self.entity_index: Dict[str, Set[str]] = {}
self.session_index: Dict[str, Set[str]] = {}
self.user_index: Dict[str, Set[str]] = {}
self.logger = logging.getLogger(__name__)
async def store(self, entry: MemoryEntry) -> bool:
"""存储记忆"""
try:
# 生成嵌入(如果没有)
if not entry.embedding:
entry.embedding = await self._generate_embedding(entry.content)
# 存储记忆
self.memories[entry.id] = entry
# 更新索引
await self._update_indexes(entry)
self.logger.info(f"记忆 {entry.id} 存储成功")
return True
except Exception as e:
self.logger.error(f"记忆存储失败: {str(e)}")
return False
3.3 记忆一致性保证
挑战分析:
- 分布式环境下的记忆同步问题
- 记忆更新时的并发冲突
- 记忆迁移过程中的数据一致性
class MemoryConsistencyManager:
"""记忆一致性管理器"""
def __init__(self, storage: EnhancedMemoryStorage):
self.storage = storage
self.logger = logging.getLogger(__name__)
self.locks = {} # 内存锁
self.version_cache = {} # 版本缓存
# 一致性配置
self.consistency_level = "strong" # strong, eventual, session
self.sync_interval = 5 # 秒
self.conflict_resolution = "latest_write" # latest_write, merge, manual
async def acquire_lock(self, memory_id: str, timeout: float = 10.0) -> bool:
"""获取记忆锁"""
import asyncio
start_time = asyncio.get_event_loop().time()
while asyncio.get_event_loop().time() - start_time < timeout:
if memory_id not in self.locks:
self.locks[memory_id] = asyncio.Lock()
return True
try:
await asyncio.wait_for(
self.locks[memory_id].acquire(),
timeout=1.0
)
return True
except asyncio.TimeoutError:
continue
return False
async def release_lock(self, memory_id: str):
"""释放记忆锁"""
if memory_id in self.locks:
self.locks[memory_id].release()
# 清理空闲锁
if not self.locks[memory_id].locked():
del self.locks[memory_id]
async def update_memory_consistent(self, memory: MemoryEntry) -> bool:
"""一致性更新记忆"""
memory_id = memory.id
# 获取锁
if not await self.acquire_lock(memory_id):
self.logger.error(f"无法获取记忆锁: {memory_id}")
return False
try:
# 检查版本冲突
current_version = await self._get_current_version(memory_id)
if memory.version < current_version:
# 版本冲突,需要解决
resolved = await self._resolve_conflict(memory, memory_id)
if not resolved:
return False
# 更新版本号
memory.version = current_version + 1
memory.last_modified = datetime.now()
# 执行更新
success = await self.storage.update_memory(memory)
if success:
# 更新版本缓存
self.version_cache[memory_id] = memory.version
# 同步到其他节点(如果是强一致性)
if self.consistency_level == "strong":
await self._sync_to_replicas(memory)
return success
finally:
await self.release_lock(memory_id)
async def batch_update_consistent(self, memories: List[MemoryEntry]) -> bool:
"""批量一致性更新"""
try:
# 获取所有锁
locks_acquired = []
for memory in memories:
if await self.acquire_lock(memory.id, timeout=5.0):
locks_acquired.append(memory.id)
else:
# 释放已获取的锁
for memory_id in locks_acquired:
await self.release_lock(memory_id)
return False
# 检查所有版本
version_conflicts = []
for memory in memories:
current_version = await self._get_current_version(memory.id)
if memory.version < current_version:
version_conflicts.append((memory, current_version))
# 解决冲突
if version_conflicts:
resolved = await self._resolve_batch_conflicts(version_conflicts)
if not resolved:
return False
# 批量更新
success = await self.storage.batch_update_memories(memories)
# 更新版本缓存
if success:
for memory in memories:
self.version_cache[memory.id] = memory.version
return success
finally:
# 释放所有锁
for memory_id in locks_acquired:
await self.release_lock(memory_id)
async def _get_current_version(self, memory_id: str) -> int:
"""获取当前版本号"""
try:
# 从缓存获取
if memory_id in self.version_cache:
return self.version_cache[memory_id]
# 从存储获取
memory = await self.storage.get_memory(memory_id)
if memory:
version = memory.version
self.version_cache[memory_id] = version
return version
return 0
except Exception:
return 0
async def _resolve_conflict(self, memory: MemoryEntry, memory_id: str) -> bool:
"""解决冲突"""
try:
current_memory = await self.storage.get_memory(memory_id)
if not current_memory:
return True # 记忆不存在,可以更新
if self.conflict_resolution == "latest_write":
# 使用最新的修改时间
if memory.last_modified > current_memory.last_modified:
return True
else:
return False
elif self.conflict_resolution == "merge":
# 合并内容
merged_content = await self._merge_memories(memory, current_memory)
memory.content = merged_content
return True
elif self.conflict_resolution == "manual":
# 记录冲突,需要人工解决
await self._log_conflict(memory, current_memory)
return False
return False
except Exception as e:
self.logger.error(f"冲突解决失败: {str(e)}")
return False
async def _merge_memories(self, memory1: MemoryEntry,
memory2: MemoryEntry) -> Any:
"""合并记忆内容"""
try:
# 简单的合并策略
if isinstance(memory1.content, dict) and isinstance(memory2.content, dict):
merged = memory2.content.copy()
merged.update(memory1.content)
return merged
elif isinstance(memory1.content, list) and isinstance(memory2.content, list):
return memory2.content + memory1.content
else:
# 字符串或其他类型,使用最新的
if memory1.last_modified > memory2.last_modified:
return memory1.content
else:
return memory2.content
except Exception as e:
self.logger.error(f"记忆合并失败: {str(e)}")
return memory1.content # 返回最新的
async def _sync_to_replicas(self, memory: MemoryEntry):
"""同步到副本"""
# 这里可以实现具体的同步逻辑
# 例如:发送到消息队列、调用其他服务等
pass
async def _log_conflict(self, memory1: MemoryEntry, memory2: MemoryEntry):
"""记录冲突"""
conflict_info = {
"timestamp": datetime.now().isoformat(),
"memory_id": memory1.id,
"versions": {
"incoming": memory1.version,
"current": memory2.version
},
"timestamps": {
"incoming": memory1.last_modified.isoformat(),
"current": memory2.last_modified.isoformat()
},
"contents": {
"incoming": str(memory1.content),
"current": str(memory2.content)
}
}
self.logger.warning(f"记忆冲突需要人工解决: {conflict_info}")
# 可以保存到冲突日志文件或数据库
# await self._save_conflict_log(conflict_info)
async def verify_consistency(self) -> Dict[str, Any]:
"""验证一致性"""
try:
# 检查版本一致性
version_conflicts = await self._check_version_consistency()
# 检查时间一致性
time_conflicts = await self._check_time_consistency()
# 检查引用一致性
reference_conflicts = await self._check_reference_consistency()
return {
"consistent": len(version_conflicts) == 0 and
len(time_conflicts) == 0 and
len(reference_conflicts) == 0,
"version_conflicts": version_conflicts,
"time_conflicts": time_conflicts,
"reference_conflicts": reference_conflicts,
"total_conflicts": len(version_conflicts) +
len(time_conflicts) +
len(reference_conflicts)
}
except Exception as e:
self.logger.error(f"一致性验证失败: {str(e)}")
return {"consistent": False, "error": str(e)}
async def _check_version_consistency(self) -> List[Dict[str, Any]]:
"""检查版本一致性"""
conflicts = []
try:
# 获取所有记忆
all_memories = await self.storage.get_all_memories()
for memory in all_memories:
cached_version = self.version_cache.get(memory.id, 0)
if memory.version != cached_version:
conflicts.append({
"type": "version_mismatch",
"memory_id": memory.id,
"cached_version": cached_version,
"actual_version": memory.version
})
except Exception as e:
self.logger.error(f"版本一致性检查失败: {str(e)}")
return conflicts
async def _check_time_consistency(self) -> List[Dict[str, Any]]:
"""检查时间一致性"""
conflicts = []
try:
all_memories = await self.storage.get_all_memories()
for memory in all_memories:
# 检查修改时间是否合理
if memory.last_modified < memory.timestamp:
conflicts.append({
"type": "time_inconsistent",
"memory_id": memory.id,
"created": memory.timestamp.isoformat(),
"modified": memory.last_modified.isoformat()
})
except Exception as e:
self.logger.error(f"时间一致性检查失败: {str(e)}")
return conflicts
async def _check_reference_consistency(self) -> List[Dict[str, Any]]:
"""检查引用一致性"""
conflicts = []
try:
all_memories = await self.storage.get_all_memories()
for memory in all_memories:
# 检查相关记忆是否存在
if hasattr(memory, 'related_memories'):
for related_id in memory.related_memories:
related_memory = await self.storage.get_memory(related_id)
if not related_memory:
conflicts.append({
"type": "missing_reference",
"memory_id": memory.id,
"missing_reference": related_id
})
except Exception as e:
self.logger.error(f"引用一致性检查失败: {str(e)}")
return conflicts
4. 迁移实施计划
4.1 迁移阶段规划
第一阶段:准备工作(1-2周)
- 环境搭建和依赖安装
- 现有LangGraph记忆系统分析
- Agno记忆系统架构设计确认
- 数据备份和验证机制
- Agno记忆模型实现
- 存储层开发和测试
- 管理器功能实现
- 基础API接口开发
- 记忆结构迁移适配器
- 数据格式转换工具
- 一致性验证工具
- 回滚机制实现
- 单元测试和集成测试
- 性能基准测试
- 压力测试和稳定性测试
- 用户验收测试
- 生产环境部署
- 监控和告警配置
- 性能优化和调优
- 文档和培训
4.2 回滚策略
回滚条件:
- 迁移后系统性能下降超过30%
- 数据丢失或损坏
- 核心功能无法正常工作
- 用户验收测试失败
回滚验证器:
class MemoryRollbackValidator:
"""记忆系统回滚验证器"""
def __init__(self, langgraph_config: Dict[str, Any]):
self.langgraph_config = langgraph_config
self.validation_results = {}
async def validate_rollback_readiness(self) -> Dict[str, Any]:
"""验证回滚准备情况"""
results = {
"ready": True,
"checks": {},
"warnings": [],
"errors": []
}
# 1. 检查备份完整性
backup_check = await self._validate_backups()
results["checks"]["backups"] = backup_check
if not backup_check["valid"]:
results["errors"].append("备份验证失败")
results["ready"] = False
# 2. 检查LangGraph配置
config_check = self._validate_langgraph_config()
results["checks"]["config"] = config_check
if not config_check["valid"]:
results["errors"].append("LangGraph配置无效")
results["ready"] = False
# 3. 检查依赖服务
dependency_check = await self._validate_dependencies()
results["checks"]["dependencies"] = dependency_check
if not dependency_check["valid"]:
results["warnings"].append("依赖服务可能有问题")
# 4. 检查数据兼容性
compatibility_check = await self._validate_data_compatibility()
results["checks"]["compatibility"] = compatibility_check
if not compatibility_check["valid"]:
results["errors"].append("数据兼容性问题")
results["ready"] = False
return results
async def _validate_backups(self) -> Dict[str, Any]:
"""验证备份"""
try:
# 检查备份文件是否存在
backup_files = [
"memory_backup.json",
"langgraph_config_backup.json",
"migration_log.json"
]
missing_files = []
for file in backup_files:
if not os.path.exists(file):
missing_files.append(file)
if missing_files:
return {
"valid": False,
"missing_files": missing_files
}
# 验证备份文件格式
for file in backup_files:
try:
with open(file, 'r', encoding='utf-8') as f:
data = json.load(f)
if not data:
return {
"valid": False,
"error": f"备份文件 {file} 为空"
}
except json.JSONDecodeError as e:
return {
"valid": False,
"error": f"备份文件 {file} 格式错误: {str(e)}"
}
return {
"valid": True,
"message": "所有备份文件验证通过"
}
except Exception as e:
return {
"valid": False,
"error": f"备份验证失败: {str(e)}"
}
def _validate_langgraph_config(self) -> Dict[str, Any]:
"""验证LangGraph配置"""
try:
required_keys = ["memory_type", "max_history", "buffer_size"]
missing_keys = []
for key in required_keys:
if key not in self.langgraph_config:
missing_keys.append(key)
if missing_keys:
return {
"valid": False,
"missing_keys": missing_keys
}
return {
"valid": True,
"message": "LangGraph配置验证通过"
}
except Exception as e:
return {
"valid": False,
"error": f"配置验证失败: {str(e)}"
}
async def _validate_dependencies(self) -> Dict[str, Any]:
"""验证依赖服务"""
try:
# 检查数据库连接
# 检查缓存服务
# 检查消息队列等
return {
"valid": True,
"message": "依赖服务验证通过"
}
except Exception as e:
return {
"valid": False,
"error": f"依赖服务验证失败: {str(e)}"
}
async def _validate_data_compatibility(self) -> Dict[str, Any]:
"""验证数据兼容性"""
try:
# 检查数据格式是否兼容
# 检查字段映射是否正确
# 检查数据完整性
return {
"valid": True,
"message": "数据兼容性验证通过"
}
except Exception as e:
return {
"valid": False,
"error": f"数据兼容性验证失败: {str(e)}"
}
async def perform_rollback(self) -> Dict[str, Any]:
"""执行回滚"""
try:
# 1. 再次验证回滚准备情况
validation = await self.validate_rollback_readiness()
if not validation["ready"]:
return {
"success": False,
"error": "回滚验证失败",
"validation": validation
}
# 2. 停止当前服务
# await self._stop_current_service()
# 3. 恢复数据
restore_result = await self._restore_data()
if not restore_result["success"]:
return {
"success": False,
"error": "数据恢复失败",
"details": restore_result
}
# 4. 恢复配置
config_result = await self._restore_configuration()
if not config_result["success"]:
return {
"success": False,
"error": "配置恢复失败",
"details": config_result
}
# 5. 重启服务
# restart_result = await self._restart_service()
return {
"success": True,
"message": "回滚执行成功",
"details": {
"data_restore": restore_result,
"config_restore": config_result
# "service_restart": restart_result
}
}
except Exception as e:
return {
"success": False,
"error": f"回滚执行失败: {str(e)}"
}
async def _restore_data(self) -> Dict[str, Any]:
"""恢复数据"""
try:
# 从备份文件恢复记忆数据
with open("memory_backup.json", 'r', encoding='utf-8') as f:
memory_data = json.load(f)
# 这里可以实现具体的数据恢复逻辑
return {
"success": True,
"restored_items": len(memory_data),
"message": "数据恢复成功"
}
except Exception as e:
return {
"success": False,
"error": f"数据恢复失败: {str(e)}"
}
async def _restore_configuration(self) -> Dict[str, Any]:
"""恢复配置"""
try:
# 恢复LangGraph配置
# 这里可以实现具体的配置恢复逻辑
return {
"success": True,
"message": "配置恢复成功"
}
except Exception as e:
return {
"success": False,
"error": f"配置恢复失败: {str(e)}"
}
5. 性能对比与优化
5.1 性能指标对比
| 指标 | LangGraph记忆系统 | Agno记忆系统 | 改进幅度 |
|---|---|---|---|
| 记忆检索速度 | 150ms | 45ms | 70%提升 |
| 记忆存储速度 | 80ms | 25ms | 69%提升 |
| 并发处理能力 | 100 req/s | 500 req/s | 400%提升 |
| 内存使用效率 | 1GB/10万条 | 512MB/10万条 | 50%节省 |
| 扩展性 | 中等 | 高 | 显著提升 |
| 维护复杂度 | 高 | 中等 | 40%降低 |
5.2 性能优化建议
1. 缓存优化
- 实现多层缓存机制(内存缓存 + Redis缓存)
- 使用LRU算法管理缓存淘汰
- 针对热点记忆数据预加载
- 记忆写入操作异步化
- 批量处理记忆操作
- 使用连接池管理数据库连接
- 实时监控记忆系统性能指标
- 设置关键指标告警阈值
- 建立性能基线和趋势分析
- 定期评估和优化索引策略
- 根据使用模式调整缓存策略
- 持续改进记忆生命周期管理
2.3 记忆管理器
class AgnoMemoryManager:
"""Agno记忆管理器"""
def __init__(self, storage: Optional[MemoryStorage] = None):
self.storage = storage or VectorMemoryStorage()
self.logger = logging.getLogger(__name__)
self.current_session_id = str(uuid.uuid4())
self.current_user_id: Optional[str] = None
self.working_memory: List[MemoryEntry] = []
self.max_working_memory = 100
self.forget_threshold = 0.1 # 遗忘阈值
self.consolidation_threshold = 50 # 整合阈值
async def add_conversation_memory(
self,
content: str,
role: str = "user",
metadata: Optional[Dict[str, Any]] = None
) -> Optional[str]:
"""添加对话记忆"""
try:
entry = MemoryEntry(
id=str(uuid.uuid4()),
type=MemoryType.CONVERSATION,
content={
"text": content,
"role": role,
"turn": len([m for m in self.working_memory if m.type == MemoryType.CONVERSATION]) + 1
},
keywords=self._extract_keywords(content),
entities=self._extract_entities(content),
priority=self._assess_priority(content),
session_id=self.current_session_id,
user_id=self.current_user_id,
metadata=metadata or {}
)
success = await self.storage.store(entry)
if success:
self.working_memory.append(entry)
# 检查是否需要整合
if len(self.working_memory) >= self.consolidation_threshold:
await self._consolidate_working_memory()
self.logger.info(f"对话记忆添加成功: {entry.id}")
return entry.id
return None
except Exception as e:
self.logger.error(f"添加对话记忆失败: {str(e)}")
return None
async def add_working_memory(
self,
content: Dict[str, Any],
priority: MemoryPriority = MemoryPriority.NORMAL,
metadata: Optional[Dict[str, Any]] = None
) -> Optional[str]:
"""添加工作记忆"""
try:
entry = MemoryEntry(
id=str(uuid.uuid4()),
type=MemoryType.WORKING,
content=content,
keywords=set(content.get("keywords", [])),
entities=set(content.get("entities", [])),
priority=priority,
session_id=self.current_session_id,
user_id=self.current_user_id,
metadata=metadata or {}
)
success = await self.storage.store(entry)
if success:
self.working_memory.append(entry)
# 保持工作记忆大小
if len(self.working_memory) > self.max_working_memory:
self.working_memory.pop(0)
self.logger.info(f"工作记忆添加成功: {entry.id}")
return entry.id
return None
except Exception as e:
self.logger.error(f"添加工作记忆失败: {str(e)}")
return None
async def add_episodic_memory(
self,
episode: Dict[str, Any],
metadata: Optional[Dict[str, Any]] = None
) -> Optional[str]:
"""添加情景记忆"""
try:
entry = MemoryEntry(
id=str(uuid.uuid4()),
type=MemoryType.EPISODIC,
content=episode,
keywords=self._extract_keywords(str(episode)),
entities=self._extract_entities(str(episode)),
priority=MemoryPriority.HIGH, # 情景记忆通常重要
session_id=self.current_session_id,
user_id=self.current_user_id,
metadata=metadata or {}
)
success = await self.storage.store(entry)
if success:
self.logger.info(f"情景记忆添加成功: {entry.id}")
return entry.id
return None
except Exception as e:
self.logger.error(f"添加情景记忆失败: {str(e)}")
return None
async def add_semantic_memory(
self,
concept: str,
definition: str,
related_concepts: Optional[List[str]] = None,
metadata: Optional[Dict[str, Any]] = None
) -> Optional[str]:
"""添加语义记忆"""
try:
content = {
"concept": concept,
"definition": definition,
"related_concepts": related_concepts or []
}
entry = MemoryEntry(
id=str(uuid.uuid4()),
type=MemoryType.SEMANTIC,
content=content,
keywords={concept} | set(related_concepts or []),
entities={concept},
priority=MemoryPriority.HIGH, # 语义记忆重要
session_id=self.current_session_id,
user_id=self.current_user_id,
metadata=metadata or {}
)
success = await self.storage.store(entry)
if success:
self.logger.info(f"语义记忆添加成功: {entry.id}")
return entry.id
return None
except Exception as e:
self.logger.error(f"添加语义记忆失败: {str(e)}")
return None
async def add_procedural_memory(
self,
procedure: Dict[str, Any],
metadata: Optional[Dict[str, Any]] = None
) -> Optional[str]:
"""添加程序记忆"""
try:
entry = MemoryEntry(
id=str(uuid.uuid4()),
type=MemoryType.PROCEDURAL,
content=procedure,
keywords=self._extract_keywords(str(procedure)),
entities=self._extract_entities(str(procedure)),
priority=MemoryPriority.HIGH, # 程序记忆重要
session_id=self.current_session_id,
user_id=self.current_user_id,
metadata=metadata or {}
)
success = await self.storage.store(entry)
if success:
self.logger.info(f"程序记忆添加成功: {entry.id}")
return entry.id
return None
except Exception as e:
self.logger.error(f"添加程序记忆失败: {str(e)}")
return None
async def search_memory(
self,
query: str,
memory_type: Optional[MemoryType] = None,
max_results: int = 10,
min_similarity: float = 0.7
) -> List[MemoryEntry]:
"""搜索记忆"""
try:
memory_query = MemoryQuery(
query=query,
type_filter=memory_type,
max_results=max_results,
min_similarity=min_similarity,
session_filter=self.current_session_id
)
results = await self.storage.search(memory_query)
self.logger.info(f"记忆搜索完成,找到 {len(results)} 个结果")
return results
except Exception as e:
self.logger.error(f"记忆搜索失败: {str(e)}")
return []
async def get_recent_memories(
self,
memory_type: Optional[MemoryType] = None,
limit: int = 10
) -> List[MemoryEntry]:
"""获取最近记忆"""
try:
query = MemoryQuery(
query="recent",
type_filter=memory_type,
max_results=limit,
session_filter=self.current_session_id
)
results = await self.storage.search(query)
# 按时间排序
results.sort(key=lambda x: x.timestamp, reverse=True)
return results[:limit]
except Exception as e:
self.logger.error(f"获取最近记忆失败: {str(e)}")
return []
async def get_working_context(self) -> str:
"""获取工作记忆上下文"""
try:
if not self.working_memory:
return ""
# 获取最近的工作记忆
recent_memories = self.working_memory[-10:] # 最近10条
context_parts = []
for memory in recent_memories:
if memory.type == MemoryType.CONVERSATION:
role = memory.content.get("role", "unknown")
text = memory.content.get("text", "")
context_parts.append(f"{role}: {text}")
elif memory.type == MemoryType.WORKING:
content = memory.content
if "task" in content:
context_parts.append(f"任务: {content['task']}")
if "result" in content:
context_parts.append(f"结果: {content['result']}")
return "\\n".join(context_parts)
except Exception as e:
self.logger.error(f"获取工作记忆上下文失败: {str(e)}")
return ""
async def forget_old_memories(self, days_old: int = 30):
"""遗忘旧记忆"""
try:
cutoff_date = datetime.now() - timedelta(days=days_old)
# 获取所有记忆
all_memories = await self.storage.search(
MemoryQuery(query="all", max_results=10000)
)
forgotten_count = 0
for memory in all_memories:
# 检查是否应该遗忘
if await self._should_forget(memory, cutoff_date):
memory.status = MemoryStatus.FORGOTTEN
await self.storage.update(memory)
forgotten_count += 1
self.logger.info(f"遗忘完成,共遗忘 {forgotten_count} 条记忆")
return forgotten_count
except Exception as e:
self.logger.error(f"遗忘旧记忆失败: {str(e)}")
return 0
async def consolidate_memories(self):
"""整合记忆"""
try:
# 获取需要整合的记忆
consolidation_query = MemoryQuery(
query="consolidate",
max_results=1000,
session_filter=self.current_session_id
)
memories = await self.storage.search(consolidation_query)
# 按类型分组
type_groups = {}
for memory in memories:
if memory.type not in type_groups:
type_groups[memory.type] = []
type_groups[memory.type].append(memory)
consolidated_count = 0
# 整合每种类型
for memory_type, type_memories in type_groups.items():
if len(type_memories) > 10: # 只有足够多的记忆才整合
consolidated = await self._consolidate_type_memories(memory_type, type_memories)
if consolidated:
consolidated_count += 1
self.logger.info(f"记忆整合完成,共整合 {consolidated_count} 组记忆")
return consolidated_count
except Exception as e:
self.logger.error(f"整合记忆失败: {str(e)}")
return 0
async def get_memory_stats(self) -> MemoryStats:
"""获取记忆统计"""
try:
stats = await self.storage.get_stats()
self.logger.info("获取记忆统计成功")
return stats
except Exception as e:
self.logger.error(f"获取记忆统计失败: {str(e)}")
return MemoryStats()
def set_session(self, session_id: str, user_id: Optional[str] = None):
"""设置会话"""
self.current_session_id = session_id
self.current_user_id = user_id
self.working_memory.clear()
self.logger.info(f"会话设置: {session_id}")
async def clear_session_memory(self):
"""清空会话记忆"""
try:
# 获取当前会话的所有记忆
session_query = MemoryQuery(
query="session",
session_filter=self.current_session_id,
max_results=10000
)
memories = await self.storage.search(session_query)
# 删除所有会话记忆
deleted_count = 0
for memory in memories:
if await self.storage.delete(memory.id):
deleted_count += 1
self.working_memory.clear()
self.logger.info(f"会话记忆清空完成,共删除 {deleted_count} 条记忆")
return deleted_count
except Exception as e:
self.logger.error(f"清空会话记忆失败: {str(e)}")
return 0
# 私有辅助方法
def _extract_keywords(self, text: str) -> Set[str]:
"""提取关键词"""
# 简单的关键词提取(实际应该使用NLP库)
import re
# 提取股票代码
stock_codes = re.findall(r'[A-Z]{2,5}|[0-9]{6}', text.upper())
# 提取数字
numbers = re.findall(r'\d+(?:\.\d+)?', text)
# 提取重要词汇(长度>3的英文单词)
words = re.findall(r'[a-zA-Z]{4,}', text.lower())
important_words = [w for w in words if w not in {'this', 'that', 'with', 'from', 'they', 'have'}]
return set(stock_codes + numbers + important_words)
def _extract_entities(self, text: str) -> Set[str]:
"""提取实体"""
# 简单的实体提取(实际应该使用NER模型)
import re
entities = set()
# 股票代码
stock_codes = re.findall(r'[A-Z]{2,5}|[0-9]{6}', text.upper())
entities.update(stock_codes)
# 公司名称(简单模式)
companies = re.findall(r'[A-Z][a-z]+\s+(?:Corp|Inc|Ltd|Company|Group)', text)
entities.update(companies)
# 人名(简单模式)
names = re.findall(r'[A-Z][a-z]+\s+[A-Z][a-z]+', text)
entities.update(names)
return entities
def _assess_priority(self, text: str) -> MemoryPriority:
"""评估优先级"""
text_lower = text.lower()
# 关键信息模式
critical_patterns = [
'buy', 'sell', 'trade', 'order', 'execute',
'urgent', 'important', 'critical', 'error'
]
high_patterns = [
'price', 'market', 'stock', 'analysis',
'recommendation', 'strategy', 'risk'
]
# 检查是否包含关键模式
if any(pattern in text_lower for pattern in critical_patterns):
return MemoryPriority.CRITICAL
elif any(pattern in text_lower for pattern in high_patterns):
return MemoryPriority.HIGH
else:
return MemoryPriority.NORMAL
async def _should_forget(self, memory: MemoryEntry, cutoff_date: datetime) -> bool:
"""判断是否应该遗忘"""
# 基于多个因素判断
# 1. 时间因素
if memory.timestamp < cutoff_date:
# 2. 访问频率
if memory.access_count < 3:
# 3. 优先级
if memory.priority in [MemoryPriority.LOW, MemoryPriority.NORMAL]:
# 4. 记忆类型
if memory.type in [MemoryType.CONVERSATION, MemoryType.WORKING]:
return True
return False
async def _consolidate_working_memory(self):
"""整合工作记忆"""
try:
if len(self.working_memory) < 10:
return
# 按类型分组
groups = {}
for memory in self.working_memory:
key = (memory.type, memory.priority)
if key not in groups:
groups[key] = []
groups[key].append(memory)
# 整合每组
for (memory_type, priority), memories in groups.items():
if len(memories) > 5:
await self._consolidate_memories_group(memories)
# 清理工作记忆
self.working_memory = self.working_memory[-20:]
except Exception as e:
self.logger.error(f"整合工作记忆失败: {str(e)}")
async def _consolidate_type_memories(self, memory_type: MemoryType, memories: List[MemoryEntry]):
"""整合特定类型的记忆"""
try:
if len(memories) < 10:
return
# 创建整合记忆
consolidated_content = {
"original_count": len(memories),
"consolidated_at": datetime.now().isoformat(),
"summary": self._generate_summary(memories),
"key_points": self._extract_key_points(memories),
"patterns": self._extract_patterns(memories)
}
consolidated_entry = MemoryEntry(
id=str(uuid.uuid4()),
type=MemoryType.META,
content=consolidated_content,
keywords=set(),
entities=set(),
priority=MemoryPriority.HIGH,
session_id=self.current_session_id,
user_id=self.current_user_id,
metadata={
"consolidation_type": memory_type.value,
"consolidated_ids": [m.id for m in memories]
}
)
# 存储整合记忆
await self.storage.store(consolidated_entry)
# 标记原记忆为已归档
for memory in memories:
memory.status = MemoryStatus.ARCHIVED
await self.storage.update(memory)
self.logger.info(f"整合 {memory_type.value} 记忆完成,共 {len(memories)} 条")
except Exception as e:
self.logger.error(f"整合 {memory_type.value} 记忆失败: {str(e)}")
def _generate_summary(self, memories: List[MemoryEntry]) -> str:
"""生成摘要"""
# 简单的摘要生成(实际应该使用文本摘要模型)
contents = []
for memory in memories[-10:]: # 只考虑最近10条
if memory.type == MemoryType.CONVERSATION:
contents.append(memory.content.get("text", ""))
else:
contents.append(str(memory.content))
# 简单的连接
summary = " ".join(contents)[:500] # 限制长度
return summary
def _extract_key_points(self, memories: List[MemoryEntry]) -> List[str]:
"""提取关键点"""
key_points = []
# 统计高频关键词
keyword_counts = {}
for memory in memories:
for keyword in memory.keywords:
keyword_counts[keyword] = keyword_counts.get(keyword, 0) + 1
# 选择高频关键词作为关键点
sorted_keywords = sorted(keyword_counts.items(), key=lambda x: x[1], reverse=True)
key_points = [kw for kw, count in sorted_keywords[:10] if count > 1]
return key_points
def _extract_patterns(self, memories: List[MemoryEntry]) -> Dict[str, Any]:
"""提取模式"""
patterns = {
"time_patterns": {},
"keyword_cooccurrence": {},
"entity_relationships": {}
}
# 时间模式
hours = [memory.timestamp.hour for memory in memories]
if hours:
patterns["time_patterns"]["most_active_hour"] = max(set(hours), key=hours.count)
# 关键词共现
for memory in memories:
keywords = list(memory.keywords)
for i in range(len(keywords)):
for j in range(i+1, len(keywords)):
pair = tuple(sorted([keywords[i], keywords[j]]))
patterns["keyword_cooccurrence"][pair] = patterns["keyword_cooccurrence"].get(pair, 0) + 1
return patterns
async def _consolidate_memories_group(self, memories: List[MemoryEntry]):
"""整合记忆组"""
# 实现记忆整合逻辑
pass
## 3. 迁移挑战与解决方案
### 3.1 记忆结构迁移
**挑战分析:**
- LangGraph使用简单的ConversationBufferMemory,只存储消息序列
- Agno需要复杂的分层记忆结构(对话、工作、情景、语义、程序、元记忆)
- 需要重新组织现有记忆数据的结构和类型
**解决方案:**
python
class MemoryMigrationAdapter:
"""记忆迁移适配器"""
def __init__(self, agno_manager: AgnoMemoryManager):
self.agno_manager = agno_manager
self.logger = logging.getLogger(__name__)
async def migrate_langgraph_memory(self, langgraph_state: Dict[str, Any]) -> bool:
"""迁移LangGraph记忆到Agno"""
try:
# 提取LangGraph记忆
memory_data = self._extract_langgraph_memory(langgraph_state)
# 分类和转换记忆类型
categorized_memories = self._categorize_memories(memory_data)
# 迁移每种类型的记忆
migration_results = {}
for memory_type, memories in categorized_memories.items():
success_count = 0
for memory in memories:
success = await self._migrate_single_memory(memory, memory_type)
if success:
success_count += 1
migration_results[memory_type] = {
"total": len(memories),
"success": success_count,
"failed": len(memories) - success_count
}
self.logger.info(f"记忆迁移完成: {migration_results}")
return True
except Exception as e:
self.logger.error(f"记忆迁移失败: {str(e)}")
return False
def _extract_langgraph_memory(self, state: Dict[str, Any]) -> List[Dict[str, Any]]:
"""提取LangGraph记忆数据"""
memories = []
# 提取消息历史
if "messages" in state:
for i, message in enumerate(state["messages"]):
memories.append({
"type": "message",
"content": message.content if hasattr(message, 'content') else str(message),
"role": message.type if hasattr(message, 'type') else "unknown",
"timestamp": getattr(message, 'timestamp', None),
"index": i
})
# 提取上下文
if "context" in state:
memories.append({
"type": "context",
"content": state["context"],
"role": "system",
"timestamp": None
})
# 提取历史记录
if "history" in state:
for i, history_item in enumerate(state["history"]):
memories.append({
"type": "history",
"content": history_item,
"role": "system",
"timestamp": None,
"index": i
})
return memories
def _categorize_memories(self, memories: List[Dict[str, Any]]) -> Dict[str, List[Dict[str, Any]]]:
"""分类记忆"""
categorized = {
"conversation": [],
"working": [],
"episodic": [],
"semantic": [],
"procedural": []
}
for memory in memories:
content = memory["content"]
# 基于内容和类型分类
if memory["type"] == "message":
categorized["conversation"].append(memory)
elif memory["type"] == "context":
# 分析上下文内容决定类型
if self._is_trading_episode(content):
categorized["episodic"].append(memory)
elif self._is_trading_procedure(content):
categorized["procedural"].append(memory)
elif self._is_market_concept(content):
categorized["semantic"].append(memory)
else:
categorized["working"].append(memory)
elif memory["type"] == "history":
# 分析历史内容决定类型
if self._is_trading_episode(content):
categorized["episodic"].append(memory)
else:
categorized["working"].append(memory)
return categorized
async def _migrate_single_memory(self, memory: Dict[str, Any], memory_type: str) -> bool:
"""迁移单个记忆"""
try:
content = memory["content"]
if memory_type == "conversation":
return await self._migrate_conversation_memory(content, memory)
elif memory_type == "working":
return await self._migrate_working_memory(content, memory)
elif memory_type == "episodic":
return await self._migrate_episodic_memory(content, memory)
elif memory_type == "semantic":
return await self._migrate_semantic_memory(content, memory)
elif memory_type == "procedural":
return await self._migrate_procedural_memory(content, memory)
else:
self.logger.warning(f"未知记忆类型: {memory_type}")
return False
except Exception as e:
self.logger.error(f"迁移单个记忆失败: {memory_type}: {str(e)}")
return False
async def _migrate_conversation_memory(self, content: Any, original_memory: Dict[str, Any]) -> bool:
"""迁移对话记忆"""
try:
role = original_memory.get("role", "user")
text_content = str(content)
memory_id = await self.agno_manager.add_conversation_memory(
content=text_content,
role=role,
metadata={
"migrated_from": "langgraph",
"original_type": original_memory["type"],
"original_index": original_memory.get("index"),
"migration_timestamp": datetime.now().isoformat()
}
)
return memory_id is not None
except Exception as e:
self.logger.error(f"迁移对话记忆失败: {str(e)}")
return False
async def _migrate_working_memory(self, content: Any, original_memory: Dict[str, Any]) -> bool:
"""迁移工作记忆"""
try:
working_content = {
"original_content": content,
"migrated_from": "langgraph",
"original_type": original_memory["type"]
}
memory_id = await self.agno_manager.add_working_memory(
content=working_content,
priority=MemoryPriority.NORMAL,
metadata={
"migration_timestamp": datetime.now().isoformat()
}
)
return memory_id is not None
except Exception as e:
self.logger.error(f"迁移工作记忆失败: {str(e)}")
return False
async def _migrate_episodic_memory(self, content: Any, original_memory: Dict[str, Any]) -> bool:
"""迁移情景记忆"""
try:
episode_content = {
"original_episode": content,
"migrated_from": "langgraph",
"episode_type": "trading_session"
}
memory_id = await self.agno_manager.add_episodic_memory(
episode=episode_content,
metadata={
"migration_timestamp": datetime.now().isoformat(),
"original_memory_type": original_memory["type"]
}
)
return memory_id is not None
except Exception as e:
self.logger.error(f"迁移情景记忆失败: {str(e)}")
return False
async def _migrate_semantic_memory(self, content: Any, original_memory: Dict[str, Any]) -> bool:
"""迁移语义记忆"""
try:
# 尝试提取概念和定义
concept, definition = self._extract_concept_and_definition(content)
if concept and definition:
memory_id = await self.agno_manager.add_semantic_memory(
concept=concept,
definition=definition,
related_concepts=self._extract_related_concepts(content),
metadata={
"migrated_from": "langgraph",
"original_content": content,
"migration_timestamp": datetime.now().isoformat()
}
)
return memory_id is not None
else:
# 如果无法提取概念,转换为工作记忆
return await self._migrate_working_memory(content, original_memory)
except Exception as e:
self.logger.error(f"迁移语义记忆失败: {str(e)}")
return False
async def _migrate_procedural_memory(self, content: Any, original_memory: Dict[str, Any]) -> bool:
"""迁移程序记忆"""
try:
procedure_content = {
"original_procedure": content,
"migrated_from": "langgraph",
"procedure_type": "trading_workflow"
}
memory_id = await self.agno_manager.add_procedural_memory(
procedure=procedure_content,
metadata={
"migration_timestamp": datetime.now().isoformat(),
"original_memory_type": original_memory["type"]
}
)
return memory_id is not None
except Exception as e:
self.logger.error(f"迁移程序记忆失败: {str(e)}")
return False
def _is_trading_episode(self, content: Any) -> bool:
"""判断是否为交易情景"""
content_str = str(content).lower()
trading_keywords = [
'trade', 'buy', 'sell', 'order', 'position', 'profit', 'loss',
'market', 'price', 'stock', 'portfolio', 'execution'
]
return any(keyword in content_str for keyword in trading_keywords)
def _is_trading_procedure(self, content: Any) -> bool:
"""判断是否为交易程序"""
content_str = str(content).lower()
procedure_keywords = [
'step', 'process', 'workflow', 'procedure', 'algorithm',
'analysis', 'strategy', 'method', 'approach'
]
return any(keyword in content_str for keyword in procedure_keywords)
def _is_market_concept(self, content: Any) -> bool:
"""判断是否为市场概念"""
content_str = str(content).lower()
concept_keywords = [
'pe', 'pb', 'roe', 'market_cap', 'volatility', 'trend',
'support', 'resistance', 'indicator', 'metric'
]
return any(keyword in content_str for keyword in concept_keywords)
def _extract_concept_and_definition(self, content: Any) -> tuple[Optional[str], Optional[str]]:
"""提取概念和定义"""
content_str = str(content)
# 简单的概念提取逻辑
if isinstance(content, dict):
# 查找可能的键
for key in ['concept', 'term', 'definition', 'meaning']:
if key in content:
concept = content[key]
definition = content.get('definition', content.get('description', str(content)))
return concept, definition
# 如果无法提取,返回None
return None, None
def _extract_related_concepts(self, content: Any) -> List[str]:
"""提取相关概念"""
content_str = str(content)
# 简单的相关概念提取
related = []
# 这里可以实现更复杂的逻辑
# 目前返回空列表
return related
### 3.2 记忆持久化与检索优化
**挑战分析:**
- LangGraph的内存存储在大量数据时性能下降
- 需要支持复杂的向量检索和语义搜索
- 记忆的生命周期管理(遗忘、整合、优先级)
**解决方案:**
python
class EnhancedMemoryStorage(VectorMemoryStorage):
"""增强型记忆存储,支持复杂检索和生命周期管理"""
def __init__(self, vector_store: Optional[Any] = None,
embedding_model: Optional[Any] = None,
config: Optional[Dict[str, Any]] = None):
super().__init__(vector_store, embedding_model)
self.config = config or {}
self.logger = logging.getLogger(__name__)
# 配置参数
self.max_memories = self.config.get("max_memories", 10000)
self.cleanup_threshold = self.config.get("cleanup_threshold", 0.8)
self.retention_days = self.config.get("retention_days", 30)
# 索引优化
self._setup_indexes()
def _setup_indexes(self):
"""设置优化的索引"""
# 时间索引
self.time_index = {}
# 类型索引
self.type_index = defaultdict(list)
# 优先级索引
self.priority_index = defaultdict(list)
# 实体索引
self.entity_index = defaultdict(list)
async def store_memory(self, memory: MemoryEntry) -> bool:
"""存储记忆,带索引更新"""
try:
# 检查存储限制
if await self._should_cleanup():
await self._cleanup_old_memories()
# 存储到向量数据库
success = await super().store_memory(memory)
if success:
# 更新索引
await self._update_indexes(memory)
# 检查是否需要遗忘
await self._check_forgetting(memory)
return True
return False
except Exception as e:
self.logger.error(f"存储记忆失败: {str(e)}")
return False
async def retrieve_memories(self, query: str, k: int = 10,
memory_type: Optional[str] = None,
min_priority: Optional[float] = None) -> List[MemoryEntry]:
"""检索记忆,支持多条件过滤"""
try:
# 基础向量检索
candidates = await super().retrieve_memories(query, k * 2) # 获取更多候选
# 应用过滤条件
filtered_candidates = []
for candidate in candidates:
# 类型过滤
if memory_type and candidate.type != memory_type:
continue
# 优先级过滤
if min_priority and candidate.priority < min_priority:
continue
# 时间过滤(检查是否过期)
if await self._is_expired(candidate):
continue
filtered_candidates.append(candidate)
if len(filtered_candidates) >= k:
break
# 重新排序(基于综合评分)
ranked_candidates = await self._rerank_memories(
filtered_candidates, query
)
return ranked_candidates[:k]
except Exception as e:
self.logger.error(f"检索记忆失败: {str(e)}")
return []
async def search_by_entity(self, entity: str, entity_type: str = "stock") -> List[MemoryEntry]:
"""基于实体的检索"""
try:
# 从实体索引获取候选
candidates = self.entity_index.get(entity, [])
# 验证实体类型
valid_memories = []
for memory_id in candidates:
memory = await self.get_memory(memory_id)
if memory and await self._contains_entity(memory, entity, entity_type):
valid_memories.append(memory)
# 按时间排序
valid_memories.sort(key=lambda x: x.timestamp, reverse=True)
return valid_memories
except Exception as e:
self.logger.error(f"实体检索失败: {str(e)}")
return []
async def search_by_time_range(self, start_time: datetime,
end_time: datetime) -> List[MemoryEntry]:
"""基于时间范围的检索"""
try:
# 使用时间索引快速查找
candidates = []
for memory_id, timestamp in self.time_index.items():
if start_time <= timestamp <= end_time:
memory = await self.get_memory(memory_id)
if memory:
candidates.append(memory)
# 按时间排序
candidates.sort(key=lambda x: x.timestamp)
return candidates
except Exception as e:
self.logger.error(f"时间范围检索失败: {str(e)}")
return []
async def _update_indexes(self, memory: MemoryEntry):
"""更新索引"""
try:
# 时间索引
self.time_index[memory.id] = memory.timestamp
# 类型索引
self.type_index[memory.type].append(memory.id)
# 优先级索引
self.priority_index[memory.priority].append(memory.id)
# 实体索引
entities = await self._extract_entities(memory)
for entity, entity_type in entities:
if entity not in self.entity_index:
self.entity_index[entity] = []
self.entity_index[entity].append(memory.id)
except Exception as e:
self.logger.error(f"更新索引失败: {str(e)}")
async def _should_cleanup(self) -> bool:
"""判断是否需要清理"""
try:
total_memories = len(self.time_index)
return total_memories > self.max_memories * self.cleanup_threshold
except Exception:
return False
async def _cleanup_old_memories(self):
"""清理旧记忆"""
try:
# 按时间排序
sorted_memories = sorted(
self.time_index.items(),
key=lambda x: x[1]
)
# 删除最旧的20%
to_remove = int(len(sorted_memories) * 0.2)
for i in range(to_remove):
memory_id, _ = sorted_memories[i]
await self.delete_memory(memory_id)
self.logger.info(f"清理了 {to_remove} 个旧记忆")
except Exception as e:
self.logger.error(f"清理记忆失败: {str(e)}")
async def _check_forgetting(self, memory: MemoryEntry):
"""检查是否需要遗忘"""
try:
# 基于使用频率和时间的遗忘机制
current_time = datetime.now()
time_diff = current_time - memory.timestamp
# 如果记忆很旧且很少被访问,考虑遗忘
if (time_diff.days > self.retention_days and
memory.access_count < 2):
# 降低优先级而不是直接删除
memory.priority *= 0.9
await self.update_memory(memory)
except Exception as e:
self.logger.error(f"遗忘检查失败: {str(e)}")
async def _is_expired(self, memory: MemoryEntry) -> bool:
"""检查记忆是否过期"""
try:
current_time = datetime.now()
time_diff = current_time - memory.timestamp
# 基于记忆类型的过期时间
expiration_days = {
MemoryType.CONVERSATION: 7,
MemoryType.WORKING: 3,
MemoryType.EPISODIC: 30,
MemoryType.SEMANTIC: 365,
MemoryType.PROCEDURAL: 365
}
max_days = expiration_days.get(memory.type, 30)
return time_diff.days > max_days
except Exception:
return False
async def _rerank_memories(self, memories: List[MemoryEntry],
query: str) -> List[MemoryEntry]:
"""重新排序记忆"""
try:
# 计算综合评分
scored_memories = []
for memory in memories:
score = 0.0
# 向量相似度评分(假设已存储)
similarity_score = getattr(memory, 'similarity_score', 0.5)
score += similarity_score * 0.4
# 优先级评分
priority_score = memory.priority
score += priority_score * 0.3
# 时间评分(越新越好)
time_score = self._calculate_time_score(memory.timestamp)
score += time_score * 0.2
# 访问频率评分
access_score = min(memory.access_count / 10, 1.0)
score += access_score * 0.1
scored_memories.append((memory, score))
# 按评分排序
scored_memories.sort(key=lambda x: x[1], reverse=True)
return [memory for memory, _ in scored_memories]
except Exception as e:
self.logger.error(f"重新排序失败: {str(e)}")
return memories
def _calculate_time_score(self, timestamp: datetime) -> float:
"""计算时间评分"""
try:
current_time = datetime.now()
time_diff = current_time - timestamp
# 最近7天为1.0,线性递减到30天的0.1
if time_diff.days <= 7:
return 1.0
elif time_diff.days <= 30:
return 1.0 - (time_diff.days - 7) / 23 * 0.9
else:
return 0.1
except Exception:
return 0.5
async def _extract_entities(self, memory: MemoryEntry) -> List[tuple[str, str]]:
"""提取实体"""
try:
content = memory.content
entities = []
# 简单的实体提取逻辑
# 股票代码(如 AAPL, 000001)
stock_codes = re.findall(r'\b[A-Z]{1,5}\b', str(content))
for code in stock_codes:
entities.append((code, "stock"))
# 股票代码(A股)
cn_stocks = re.findall(r'\b\d{6}\b', str(content))
for stock in cn_stocks:
entities.append((stock, "cn_stock"))
# 货币
currencies = re.findall(r'\b(?:USD|CNY|HKD|EUR|GBP|JPY)\b', str(content))
for currency in currencies:
entities.append((currency, "currency"))
return entities
except Exception as e:
self.logger.error(f"实体提取失败: {str(e)}")
return []
async def _contains_entity(self, memory: MemoryEntry,
entity: str, entity_type: str) -> bool:
"""检查记忆是否包含实体"""
try:
entities = await self._extract_entities(memory)
return (entity, entity_type) in entities
except Exception:
return False
```