agAdvisor / src /evaluators /reflection_evaluator.py
tirtho149's picture
Deploy AgAdvisor
b30f068 verified
Raw
History Blame Contribute Delete
18.4 kB
"""
Reflection-style self-critique evaluator implementing the plan → act → observe → evaluate → reflect → update memory → replan cycle.
"""
import time
from typing import Dict, Any, List, Optional, Tuple
from dataclasses import dataclass
from enum import Enum
import json
from src.agents.base_agent import BaseAgent, AgentResult, AgentState
from src.utils.logging_config import logger
from config.settings import settings
class ReflectionStage(Enum):
"""Stages in the reflection cycle."""
PLAN = "plan"
ACT = "act"
OBSERVE = "observe"
EVALUATE = "evaluate"
REFLECT = "reflect"
UPDATE_MEMORY = "update_memory"
REPLAN = "replan"
@dataclass
class ReflectionResult:
"""Result of a reflection evaluation."""
quality_score: float
improvement_suggestions: List[str]
should_retry: bool
confidence: float
reasoning: str
next_actions: List[str]
memory_updates: Dict[str, Any]
@dataclass
class TrajectoryStep:
"""Single step in the execution trajectory."""
stage: ReflectionStage
input_data: Dict[str, Any]
output_data: Dict[str, Any]
success: bool
execution_time: float
errors: List[str]
timestamp: float
class ReflectionEvaluator(BaseAgent):
"""
Evaluator agent that implements reflection-style self-critique.
Implements the cycle: plan → act → observe → evaluate → reflect → update memory → replan
"""
def _initialize(self) -> None:
"""Initialize the reflection evaluator."""
self.max_iterations = self.config.get('max_iterations', settings.reflection_max_iterations)
self.improvement_threshold = self.config.get('improvement_threshold', settings.reflection_improvement_threshold)
self.quality_threshold = self.config.get('quality_threshold', 0.7)
# Memory storage for learning
self.memory = {
'successful_patterns': [],
'failure_patterns': [],
'improvement_history': [],
'quality_trends': []
}
logger.info("Reflection evaluator initialized")
def execute(self, input_data: Dict[str, Any]) -> AgentResult:
"""
Execute reflection evaluation on a completed workflow.
Args:
input_data: Contains query, results, workflow_state, and trajectory
Returns:
AgentResult with reflection analysis and recommendations
"""
start_time = time.time()
self.set_state(AgentState.RUNNING)
try:
query = input_data.get('query', '')
results = input_data.get('results', {})
workflow_state = input_data.get('workflow_state')
trajectory = input_data.get('trajectory', [])
# Perform reflection evaluation
reflection_result = self._perform_reflection_cycle(query, results, workflow_state, trajectory)
execution_time = time.time() - start_time
self.set_state(AgentState.COMPLETED)
result = AgentResult(
agent_id=self.agent_id,
success=True,
data={
'quality_score': reflection_result.quality_score,
'improvement_suggestions': reflection_result.improvement_suggestions,
'should_retry': reflection_result.should_retry,
'confidence': reflection_result.confidence,
'reasoning': reflection_result.reasoning,
'next_actions': reflection_result.next_actions,
'memory_updates': reflection_result.memory_updates,
'reflection_stage': 'completed'
},
execution_time=execution_time,
metadata={
'reflection_method': 'comprehensive',
'trajectory_length': len(trajectory),
'memory_size': len(self.memory['successful_patterns']) + len(self.memory['failure_patterns'])
}
)
# Update memory based on reflection
self._update_memory(reflection_result)
self.log_execution(result)
return result
except Exception as e:
execution_time = time.time() - start_time
self.set_state(AgentState.FAILED)
result = AgentResult(
agent_id=self.agent_id,
success=False,
data=None,
error_message=str(e),
execution_time=execution_time
)
self.log_execution(result)
return result
def _perform_reflection_cycle(
self,
query: str,
results: Dict[str, Any],
workflow_state: Any,
trajectory: List[Dict[str, Any]]
) -> ReflectionResult:
"""Perform the complete reflection cycle."""
# 1. OBSERVE: Analyze what happened
observations = self._observe_execution(query, results, trajectory)
# 2. EVALUATE: Score the outcomes
quality_score = self._evaluate_quality(query, results, observations)
# 3. REFLECT: Generate insights and improvements
reflection_insights = self._reflect_on_performance(query, results, observations, quality_score)
# 4. UPDATE MEMORY: Learn from this execution
memory_updates = self._generate_memory_updates(query, results, observations, reflection_insights)
# 5. REPLAN: Determine next actions
next_actions = self._generate_next_actions(quality_score, reflection_insights)
return ReflectionResult(
quality_score=quality_score,
improvement_suggestions=reflection_insights['improvements'],
should_retry=quality_score < self.quality_threshold and reflection_insights['can_improve'],
confidence=reflection_insights['confidence'],
reasoning=reflection_insights['reasoning'],
next_actions=next_actions,
memory_updates=memory_updates
)
def _observe_execution(self, query: str, results: Dict[str, Any], trajectory: List[Dict[str, Any]]) -> Dict[str, Any]:
"""Observe and analyze the execution."""
observations = {
'query_complexity': self._assess_query_complexity(query),
'result_quality': self._assess_result_quality(results),
'execution_efficiency': self._assess_execution_efficiency(trajectory),
'error_patterns': self._identify_error_patterns(trajectory),
'success_patterns': self._identify_success_patterns(trajectory)
}
logger.debug(f"Execution observations: {observations}")
return observations
def _evaluate_quality(self, query: str, results: Dict[str, Any], observations: Dict[str, Any]) -> float:
"""Evaluate the overall quality of the execution."""
quality_factors = []
# Factor 1: Result completeness (40%)
api_results = results.get('results', {})
executed_count = api_results.get('executed_count', 0)
successful_count = api_results.get('successful_calls', 0)
if executed_count > 0:
completeness_score = successful_count / executed_count
else:
completeness_score = 0.0
quality_factors.append(('completeness', completeness_score, 0.4))
# Factor 2: Relevance of API matches (30%)
api_matches = results.get('api_matches', {})
match_count = api_matches.get('count', 0)
if match_count > 0:
# Check confidence scores of matches
matches = api_matches.get('matches', [])
if matches:
avg_confidence = sum(m.get('confidence', 0) for m in matches) / len(matches)
relevance_score = min(1.0, avg_confidence)
else:
relevance_score = 0.5
else:
relevance_score = 0.0
quality_factors.append(('relevance', relevance_score, 0.3))
# Factor 3: Execution efficiency (20%)
efficiency_score = observations.get('execution_efficiency', 0.5)
quality_factors.append(('efficiency', efficiency_score, 0.2))
# Factor 4: Error handling (10%)
error_patterns = observations.get('error_patterns', [])
error_score = 1.0 - min(1.0, len(error_patterns) * 0.2)
quality_factors.append(('error_handling', error_score, 0.1))
# Calculate weighted average
total_score = sum(score * weight for _, score, weight in quality_factors)
logger.info(f"Quality evaluation: {total_score:.3f} (factors: {quality_factors})")
return total_score
def _reflect_on_performance(
self,
query: str,
results: Dict[str, Any],
observations: Dict[str, Any],
quality_score: float
) -> Dict[str, Any]:
"""Generate reflection insights and improvement suggestions."""
improvements = []
reasoning_parts = []
can_improve = False
# Analyze each aspect
if quality_score < 0.8:
# Check API matching
api_matches = results.get('api_matches', {})
if api_matches.get('count', 0) == 0:
improvements.append("Improve keyword extraction to find more relevant terms")
improvements.append("Expand API operation database with more diverse operations")
can_improve = True
reasoning_parts.append("No API matches found - keyword extraction needs improvement")
# Check execution success
api_results = results.get('results', {})
executed_count = api_results.get('executed_count', 0)
successful_count = api_results.get('successful_calls', 0)
if executed_count > 0 and successful_count < executed_count:
improvements.append("Enhance error handling and retry mechanisms for API calls")
improvements.append("Implement better API health checking before execution")
can_improve = True
reasoning_parts.append(f"Only {successful_count}/{executed_count} API calls succeeded")
# Check query understanding
query_info = results.get('query', {})
confidence = query_info.get('confidence', 0)
if confidence < 0.7:
improvements.append("Improve natural language understanding for complex queries")
improvements.append("Add more training patterns for query intent classification")
can_improve = True
reasoning_parts.append(f"Low query understanding confidence: {confidence:.2f}")
# Check for patterns in memory
similar_patterns = self._find_similar_patterns(query, results)
if similar_patterns:
improvements.append("Apply lessons learned from similar previous queries")
reasoning_parts.append("Found similar patterns in execution history")
reasoning = "; ".join(reasoning_parts) if reasoning_parts else "Execution completed successfully"
return {
'improvements': improvements,
'reasoning': reasoning,
'can_improve': can_improve,
'confidence': min(1.0, quality_score + 0.1),
'similar_patterns': similar_patterns
}
def _generate_memory_updates(
self,
query: str,
results: Dict[str, Any],
observations: Dict[str, Any],
insights: Dict[str, Any]
) -> Dict[str, Any]:
"""Generate updates for the memory system."""
updates = {
'timestamp': time.time(),
'query_pattern': self._extract_query_pattern(query),
'result_pattern': self._extract_result_pattern(results),
'quality_score': observations.get('result_quality', 0),
'insights': insights['improvements']
}
return updates
def _generate_next_actions(self, quality_score: float, insights: Dict[str, Any]) -> List[str]:
"""Generate recommended next actions."""
actions = []
if quality_score >= self.quality_threshold:
actions.append("Continue with current approach - quality threshold met")
else:
actions.append("Consider retry with improved parameters")
actions.extend(insights['improvements'][:3]) # Top 3 improvements
if insights['similar_patterns']:
actions.append("Review similar patterns for additional optimization opportunities")
return actions
def _update_memory(self, reflection_result: ReflectionResult) -> None:
"""Update the memory system with new learnings."""
if reflection_result.quality_score >= self.quality_threshold:
self.memory['successful_patterns'].append(reflection_result.memory_updates)
else:
self.memory['failure_patterns'].append(reflection_result.memory_updates)
self.memory['quality_trends'].append({
'timestamp': time.time(),
'quality_score': reflection_result.quality_score,
'improvements_suggested': len(reflection_result.improvement_suggestions)
})
# Keep memory size manageable
max_patterns = 100
if len(self.memory['successful_patterns']) > max_patterns:
self.memory['successful_patterns'] = self.memory['successful_patterns'][-max_patterns:]
if len(self.memory['failure_patterns']) > max_patterns:
self.memory['failure_patterns'] = self.memory['failure_patterns'][-max_patterns:]
# Helper methods for analysis
def _assess_query_complexity(self, query: str) -> float:
"""Assess the complexity of the input query."""
factors = [
len(query.split()) / 20.0, # Word count factor
len([w for w in query.split() if len(w) > 6]) / 10.0, # Complex words
query.count('and') + query.count('or') + query.count('but'), # Logical operators
]
return min(1.0, sum(factors) / len(factors))
def _assess_result_quality(self, results: Dict[str, Any]) -> float:
"""Assess the quality of results produced."""
api_results = results.get('results', {})
executed = api_results.get('executed_count', 0)
successful = api_results.get('successful_calls', 0)
if executed == 0:
return 0.0
return successful / executed
def _assess_execution_efficiency(self, trajectory: List[Dict[str, Any]]) -> float:
"""Assess the efficiency of execution."""
if not trajectory:
return 0.5
total_time = sum(step.get('execution_time', 0) for step in trajectory)
avg_time = total_time / len(trajectory) if trajectory else 0
# Efficiency based on average step time (lower is better)
efficiency = max(0.0, 1.0 - (avg_time / 5.0)) # 5 seconds as baseline
return efficiency
def _identify_error_patterns(self, trajectory: List[Dict[str, Any]]) -> List[str]:
"""Identify patterns in errors."""
errors = []
for step in trajectory:
if not step.get('success', True):
errors.extend(step.get('errors', []))
# Group similar errors
error_patterns = list(set(errors))
return error_patterns
def _identify_success_patterns(self, trajectory: List[Dict[str, Any]]) -> List[str]:
"""Identify patterns in successful executions."""
successes = []
for step in trajectory:
if step.get('success', False):
successes.append(step.get('stage', 'unknown'))
return list(set(successes))
def _find_similar_patterns(self, query: str, results: Dict[str, Any]) -> List[Dict[str, Any]]:
"""Find similar patterns in memory."""
# Simple similarity based on query keywords
query_words = set(query.lower().split())
similar = []
for pattern in self.memory['successful_patterns']:
pattern_query = pattern.get('query_pattern', {})
pattern_words = set(pattern_query.get('keywords', []))
if query_words.intersection(pattern_words):
similar.append(pattern)
return similar[:3] # Return top 3 similar patterns
def _extract_query_pattern(self, query: str) -> Dict[str, Any]:
"""Extract pattern from query for memory storage."""
return {
'length': len(query),
'keywords': query.lower().split(),
'complexity': self._assess_query_complexity(query)
}
def _extract_result_pattern(self, results: Dict[str, Any]) -> Dict[str, Any]:
"""Extract pattern from results for memory storage."""
return {
'api_matches': results.get('api_matches', {}).get('count', 0),
'executed_apis': results.get('results', {}).get('executed_count', 0),
'success_rate': self._assess_result_quality(results)
}
def get_capabilities(self) -> List[str]:
"""Get evaluator capabilities."""
return [
"reflection_evaluation",
"quality_assessment",
"improvement_suggestions",
"memory_management",
"pattern_recognition",
"self_critique"
]