""" 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" ]