agAdvisor / src /evaluators /memory_manager.py
tirtho149's picture
Deploy AgAdvisor
b30f068 verified
Raw
History Blame Contribute Delete
10.5 kB
"""
Memory management system for storing and retrieving learned patterns.
"""
import json
import time
from typing import Dict, Any, List, Optional
from pathlib import Path
from dataclasses import dataclass, asdict
from src.utils.logging_config import logger
@dataclass
class MemoryEntry:
"""Single memory entry storing learned patterns."""
id: str
timestamp: float
query_pattern: Dict[str, Any]
result_pattern: Dict[str, Any]
quality_score: float
improvements: List[str]
context: Dict[str, Any]
class MemoryManager:
"""Manages persistent memory for the agentic AI system."""
def __init__(self, memory_file: str = "memory/system_memory.json"):
"""Initialize memory manager."""
self.memory_file = Path(memory_file)
self.memory_file.parent.mkdir(parents=True, exist_ok=True)
self.memory_entries: List[MemoryEntry] = []
self.max_entries = 1000
# Load existing memory
self._load_memory()
logger.info(f"Memory manager initialized with {len(self.memory_entries)} entries")
def store_memory(
self,
query: str,
results: Dict[str, Any],
quality_score: float,
improvements: List[str],
context: Dict[str, Any] = None
) -> str:
"""Store a new memory entry."""
entry_id = f"mem_{int(time.time())}_{len(self.memory_entries)}"
entry = MemoryEntry(
id=entry_id,
timestamp=time.time(),
query_pattern=self._extract_query_pattern(query),
result_pattern=self._extract_result_pattern(results),
quality_score=quality_score,
improvements=improvements,
context=context or {}
)
self.memory_entries.append(entry)
# Maintain memory size limit
if len(self.memory_entries) > self.max_entries:
self.memory_entries = self.memory_entries[-self.max_entries:]
# Save to disk
self._save_memory()
logger.debug(f"Stored memory entry: {entry_id}")
return entry_id
def retrieve_similar_memories(
self,
query: str,
limit: int = 5,
min_similarity: float = 0.3
) -> List[MemoryEntry]:
"""Retrieve memories similar to the given query."""
query_pattern = self._extract_query_pattern(query)
similar_memories = []
for entry in self.memory_entries:
similarity = self._calculate_similarity(query_pattern, entry.query_pattern)
if similarity >= min_similarity:
similar_memories.append((entry, similarity))
# Sort by similarity (descending) and return top results
similar_memories.sort(key=lambda x: x[1], reverse=True)
return [entry for entry, _ in similar_memories[:limit]]
def get_successful_patterns(self, min_quality: float = 0.7) -> List[MemoryEntry]:
"""Get memories with high quality scores."""
return [
entry for entry in self.memory_entries
if entry.quality_score >= min_quality
]
def get_failure_patterns(self, max_quality: float = 0.5) -> List[MemoryEntry]:
"""Get memories with low quality scores for learning from failures."""
return [
entry for entry in self.memory_entries
if entry.quality_score <= max_quality
]
def get_improvement_suggestions(self, query: str) -> List[str]:
"""Get improvement suggestions based on similar past experiences."""
similar_memories = self.retrieve_similar_memories(query)
all_improvements = []
for memory in similar_memories:
all_improvements.extend(memory.improvements)
# Count frequency and return most common suggestions
improvement_counts = {}
for improvement in all_improvements:
improvement_counts[improvement] = improvement_counts.get(improvement, 0) + 1
# Sort by frequency
sorted_improvements = sorted(
improvement_counts.items(),
key=lambda x: x[1],
reverse=True
)
return [improvement for improvement, _ in sorted_improvements[:5]]
def get_memory_stats(self) -> Dict[str, Any]:
"""Get statistics about stored memories."""
if not self.memory_entries:
return {'total_entries': 0}
quality_scores = [entry.quality_score for entry in self.memory_entries]
return {
'total_entries': len(self.memory_entries),
'avg_quality': sum(quality_scores) / len(quality_scores),
'high_quality_count': len([s for s in quality_scores if s >= 0.7]),
'low_quality_count': len([s for s in quality_scores if s <= 0.5]),
'oldest_entry': min(entry.timestamp for entry in self.memory_entries),
'newest_entry': max(entry.timestamp for entry in self.memory_entries)
}
def clear_old_memories(self, days_old: int = 30) -> int:
"""Clear memories older than specified days."""
cutoff_time = time.time() - (days_old * 24 * 60 * 60)
old_count = len(self.memory_entries)
self.memory_entries = [
entry for entry in self.memory_entries
if entry.timestamp > cutoff_time
]
removed_count = old_count - len(self.memory_entries)
if removed_count > 0:
self._save_memory()
logger.info(f"Cleared {removed_count} old memory entries")
return removed_count
def _extract_query_pattern(self, query: str) -> Dict[str, Any]:
"""Extract pattern features from a query."""
words = query.lower().split()
return {
'length': len(query),
'word_count': len(words),
'keywords': words,
'has_api_terms': any(term in query.lower() for term in ['api', 'endpoint', 'service']),
'has_data_terms': any(term in query.lower() for term in ['data', 'dataset', 'information']),
'has_action_terms': any(term in query.lower() for term in ['get', 'find', 'search', 'create', 'update']),
'complexity_score': len([w for w in words if len(w) > 6]) / len(words) if words else 0
}
def _extract_result_pattern(self, results: Dict[str, Any]) -> Dict[str, Any]:
"""Extract pattern features from results."""
api_matches = results.get('api_matches', {})
api_results = results.get('results', {})
return {
'match_count': api_matches.get('count', 0),
'executed_count': api_results.get('executed_count', 0),
'successful_count': api_results.get('successful_calls', 0),
'success_rate': (
api_results.get('successful_calls', 0) / api_results.get('executed_count', 1)
if api_results.get('executed_count', 0) > 0 else 0
),
'has_errors': any(
not r.get('success', True)
for r in api_results.get('data', [])
)
}
def _calculate_similarity(self, pattern1: Dict[str, Any], pattern2: Dict[str, Any]) -> float:
"""Calculate similarity between two query patterns."""
similarity_factors = []
# Keyword overlap
keywords1 = set(pattern1.get('keywords', []))
keywords2 = set(pattern2.get('keywords', []))
if keywords1 and keywords2:
overlap = len(keywords1.intersection(keywords2))
union = len(keywords1.union(keywords2))
keyword_similarity = overlap / union if union > 0 else 0
similarity_factors.append(keyword_similarity * 0.4)
# Length similarity
len1 = pattern1.get('length', 0)
len2 = pattern2.get('length', 0)
if len1 > 0 and len2 > 0:
length_similarity = 1 - abs(len1 - len2) / max(len1, len2)
similarity_factors.append(length_similarity * 0.2)
# Feature similarity (boolean features)
features = ['has_api_terms', 'has_data_terms', 'has_action_terms']
feature_matches = sum(
1 for feature in features
if pattern1.get(feature, False) == pattern2.get(feature, False)
)
feature_similarity = feature_matches / len(features)
similarity_factors.append(feature_similarity * 0.3)
# Complexity similarity
comp1 = pattern1.get('complexity_score', 0)
comp2 = pattern2.get('complexity_score', 0)
complexity_similarity = 1 - abs(comp1 - comp2)
similarity_factors.append(complexity_similarity * 0.1)
return sum(similarity_factors) if similarity_factors else 0.0
def _load_memory(self) -> None:
"""Load memory from disk."""
if not self.memory_file.exists():
return
try:
with open(self.memory_file, 'r') as f:
data = json.load(f)
self.memory_entries = [
MemoryEntry(**entry_data)
for entry_data in data.get('entries', [])
]
logger.info(f"Loaded {len(self.memory_entries)} memory entries from disk")
except Exception as e:
logger.error(f"Failed to load memory from disk: {e}")
self.memory_entries = []
def _save_memory(self) -> None:
"""Save memory to disk."""
try:
data = {
'version': '1.0',
'timestamp': time.time(),
'entries': [asdict(entry) for entry in self.memory_entries]
}
with open(self.memory_file, 'w') as f:
json.dump(data, f, indent=2)
logger.debug(f"Saved {len(self.memory_entries)} memory entries to disk")
except Exception as e:
logger.error(f"Failed to save memory to disk: {e}")