| """
|
| 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)
|
|
|
|
|
| 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', [])
|
|
|
|
|
| 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'])
|
| }
|
| )
|
|
|
|
|
| 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."""
|
|
|
|
|
| observations = self._observe_execution(query, results, trajectory)
|
|
|
|
|
| quality_score = self._evaluate_quality(query, results, observations)
|
|
|
|
|
| reflection_insights = self._reflect_on_performance(query, results, observations, quality_score)
|
|
|
|
|
| memory_updates = self._generate_memory_updates(query, results, observations, reflection_insights)
|
|
|
|
|
| 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 = []
|
|
|
|
|
| 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))
|
|
|
|
|
| api_matches = results.get('api_matches', {})
|
| match_count = api_matches.get('count', 0)
|
|
|
| if match_count > 0:
|
|
|
| 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))
|
|
|
|
|
| efficiency_score = observations.get('execution_efficiency', 0.5)
|
| quality_factors.append(('efficiency', efficiency_score, 0.2))
|
|
|
|
|
| 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))
|
|
|
|
|
| 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
|
|
|
|
|
| if quality_score < 0.8:
|
|
|
| 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")
|
|
|
|
|
| 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")
|
|
|
|
|
| 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}")
|
|
|
|
|
| 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])
|
|
|
| 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)
|
| })
|
|
|
|
|
| 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:]
|
|
|
|
|
| def _assess_query_complexity(self, query: str) -> float:
|
| """Assess the complexity of the input query."""
|
| factors = [
|
| len(query.split()) / 20.0,
|
| len([w for w in query.split() if len(w) > 6]) / 10.0,
|
| query.count('and') + query.count('or') + query.count('but'),
|
| ]
|
| 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 = max(0.0, 1.0 - (avg_time / 5.0))
|
| 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', []))
|
|
|
|
|
| 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."""
|
|
|
| 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]
|
|
|
| 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"
|
| ]
|
|
|