""" Arbitration Session Management for Trixy This module provides comprehensive session tracking and management for arbitration processes. It handles the lifecycle of arbitration sessions from wakeword detection through satellite selection and conversation handoff. The session system implements the 1-second collection window specified in CLAUDE.md and tracks all aspects of the arbitration process including: - Multiple satellite wakeword reports collection - Timing and timeout management - Selection algorithm execution - Result tracking and analytics - Session state management - Integration with conversation system Usage: from trixy_core.arbitration import ArbitrationSession, ArbitrationSessionManager # Create session manager session_manager = ArbitrationSessionManager(config, algorithms) # Start new arbitration session session = session_manager.start_session("session_123") # Add wakeword reports session.add_wakeword_report(report1) session.add_wakeword_report(report2) # Execute arbitration after collection window result = session.execute_arbitration() # Complete session session.complete(conversation_id="conv_456") """ import asyncio import time import uuid from dataclasses import dataclass, field from datetime import datetime, timedelta from enum import Enum from typing import Dict, List, Optional, Any, Callable, Union import threading from concurrent.futures import ThreadPoolExecutor, Future from .arbitration_config import ArbitrationConfig, ArbitrationAlgorithm from .arbitration_algorithms import ( ArbitrationAlgorithms, WakewordReport, SelectionResult, ArbitrationConfigError ) def pprint(message: str) -> None: """Arbitration session logging function.""" print(f"[ARBITRATION_SESSION] {message}") class SessionState(Enum): """States of an arbitration session.""" CREATED = "created" COLLECTING = "collecting" READY_FOR_ARBITRATION = "ready_for_arbitration" ARBITRATING = "arbitrating" COMPLETED = "completed" CANCELLED = "cancelled" TIMED_OUT = "timed_out" ERROR = "error" class SessionPhase(Enum): """Phases within an arbitration session.""" INITIALIZATION = "initialization" COLLECTION_WINDOW = "collection_window" SELECTION_PROCESS = "selection_process" RESULT_PROCESSING = "result_processing" HANDOFF = "handoff" CLEANUP = "cleanup" @dataclass class SessionTiming: """Timing information for an arbitration session.""" created_at: datetime = field(default_factory=datetime.now) collection_started_at: Optional[datetime] = None collection_ended_at: Optional[datetime] = None arbitration_started_at: Optional[datetime] = None arbitration_ended_at: Optional[datetime] = None completed_at: Optional[datetime] = None def get_collection_duration(self) -> Optional[float]: """Get collection window duration in seconds.""" if self.collection_started_at and self.collection_ended_at: return (self.collection_ended_at - self.collection_started_at).total_seconds() return None def get_arbitration_duration(self) -> Optional[float]: """Get arbitration processing duration in seconds.""" if self.arbitration_started_at and self.arbitration_ended_at: return (self.arbitration_ended_at - self.arbitration_started_at).total_seconds() return None def get_total_duration(self) -> Optional[float]: """Get total session duration in seconds.""" if self.completed_at: return (self.completed_at - self.created_at).total_seconds() return None def to_dict(self) -> Dict[str, Any]: """Convert timing to dictionary representation.""" return { 'created_at': self.created_at.isoformat(), 'collection_started_at': self.collection_started_at.isoformat() if self.collection_started_at else None, 'collection_ended_at': self.collection_ended_at.isoformat() if self.collection_ended_at else None, 'arbitration_started_at': self.arbitration_started_at.isoformat() if self.arbitration_started_at else None, 'arbitration_ended_at': self.arbitration_ended_at.isoformat() if self.arbitration_ended_at else None, 'completed_at': self.completed_at.isoformat() if self.completed_at else None, 'collection_duration': self.get_collection_duration(), 'arbitration_duration': self.get_arbitration_duration(), 'total_duration': self.get_total_duration(), } @dataclass class SessionMetrics: """Metrics and analytics for an arbitration session.""" total_reports_received: int = 0 unique_satellites_count: int = 0 unique_rooms_count: int = 0 average_volume: float = 0.0 average_confidence: float = 0.0 volume_range: tuple = field(default_factory=lambda: (0.0, 0.0)) confidence_range: tuple = field(default_factory=lambda: (0.0, 0.0)) algorithms_attempted: List[str] = field(default_factory=list) selection_confidence: float = 0.0 def update_from_reports(self, reports: List[WakewordReport]) -> None: """Update metrics from wakeword reports.""" if not reports: return self.total_reports_received = len(reports) self.unique_satellites_count = len(set(r.satellite_id for r in reports)) self.unique_rooms_count = len(set(r.room_id for r in reports)) volumes = [r.volume for r in reports] confidences = [r.confidence for r in reports] self.average_volume = sum(volumes) / len(volumes) self.average_confidence = sum(confidences) / len(confidences) self.volume_range = (min(volumes), max(volumes)) self.confidence_range = (min(confidences), max(confidences)) def to_dict(self) -> Dict[str, Any]: """Convert metrics to dictionary representation.""" return { 'total_reports_received': self.total_reports_received, 'unique_satellites_count': self.unique_satellites_count, 'unique_rooms_count': self.unique_rooms_count, 'average_volume': self.average_volume, 'average_confidence': self.average_confidence, 'volume_range': self.volume_range, 'confidence_range': self.confidence_range, 'algorithms_attempted': self.algorithms_attempted, 'selection_confidence': self.selection_confidence, } class SessionError(Exception): """Base exception for arbitration session errors.""" pass class SessionTimeoutError(SessionError): """Raised when session times out.""" pass class SessionStateError(SessionError): """Raised when session is in invalid state for operation.""" pass class SessionConfigurationError(SessionError): """Raised when session configuration is invalid.""" pass class ArbitrationSession: """ Individual arbitration session for managing wakeword conflict resolution. This class handles a single arbitration session from creation through completion, implementing the 1-second collection window and satellite selection process as specified in CLAUDE.md. """ def __init__( self, session_id: str, config: ArbitrationConfig, algorithms: ArbitrationAlgorithms, trigger_report: Optional[WakewordReport] = None, callback: Optional[Callable] = None ): """ Initialize arbitration session. Args: session_id: Unique session identifier config: Arbitration configuration algorithms: Arbitration algorithms processor trigger_report: Initial wakeword report that triggered session callback: Optional callback for session events """ self.session_id = session_id self.config = config self.algorithms = algorithms self.callback = callback # Session state self.state = SessionState.CREATED self.phase = SessionPhase.INITIALIZATION self.error_message: Optional[str] = None # Timing and metrics self.timing = SessionTiming() self.metrics = SessionMetrics() # Wakeword reports self.wakeword_reports: List[WakewordReport] = [] if trigger_report: self.wakeword_reports.append(trigger_report) # Selection result self.selection_result: Optional[SelectionResult] = None self.selected_satellite_id: Optional[str] = None self.conversation_id: Optional[str] = None # Threading and timing self._lock = threading.RLock() self._collection_timer: Optional[threading.Timer] = None self._timeout_timer: Optional[threading.Timer] = None self._is_cancelled = False # Metadata self.metadata: Dict[str, Any] = {} pprint(f"Created arbitration session {session_id}") def start_collection(self) -> None: """ Start the collection window for wakeword reports. This implements the 1-second collection window specified in CLAUDE.md. """ with self._lock: if self.state != SessionState.CREATED: raise SessionStateError(f"Cannot start collection in state {self.state}") self.state = SessionState.COLLECTING self.phase = SessionPhase.COLLECTION_WINDOW self.timing.collection_started_at = datetime.now() # Start collection timer self._collection_timer = threading.Timer( self.config.collection_window_seconds, self._collection_timeout ) self._collection_timer.start() # Start overall timeout timer self._timeout_timer = threading.Timer( self.config.arbitration_timeout_seconds, self._session_timeout ) self._timeout_timer.start() pprint(f"Started collection window for session {self.session_id} " f"({self.config.collection_window_seconds}s)") self._notify_callback("collection_started") def add_wakeword_report(self, report: WakewordReport) -> bool: """ Add a wakeword report to the session. Args: report: Wakeword detection report Returns: bool: True if report was added, False if rejected """ with self._lock: if self.state not in [SessionState.CREATED, SessionState.COLLECTING]: pprint(f"Rejected report for session {self.session_id} - wrong state {self.state}") return False if self._is_cancelled: pprint(f"Rejected report for session {self.session_id} - session cancelled") return False # Validate report timestamp (should be recent) current_time = time.time() report_age = current_time - report.timestamp if report_age > self.config.max_wait_time: pprint(f"Rejected old report for session {self.session_id} - age {report_age:.2f}s") return False # Check for duplicate satellite reports existing_satellite_ids = [r.satellite_id for r in self.wakeword_reports] if report.satellite_id in existing_satellite_ids: pprint(f"Rejected duplicate report from {report.satellite_id} for session {self.session_id}") return False # Add report self.wakeword_reports.append(report) pprint(f"Added wakeword report from {report.satellite_id} to session {self.session_id} " f"(total: {len(self.wakeword_reports)})") # Start collection if this is the first report if len(self.wakeword_reports) == 1 and self.state == SessionState.CREATED: self.start_collection() self._notify_callback("report_added", report) return True def _collection_timeout(self) -> None: """Handle collection window timeout.""" with self._lock: if self.state != SessionState.COLLECTING: return self.timing.collection_ended_at = datetime.now() self.state = SessionState.READY_FOR_ARBITRATION self.phase = SessionPhase.SELECTION_PROCESS pprint(f"Collection window ended for session {self.session_id} " f"with {len(self.wakeword_reports)} reports") self._notify_callback("collection_ended") # Automatically execute arbitration if we have reports if self.wakeword_reports: try: self.execute_arbitration() except Exception as e: self._handle_error(f"Auto-arbitration failed: {e}") else: self._handle_error("No wakeword reports received during collection window") def _session_timeout(self) -> None: """Handle overall session timeout.""" with self._lock: if self.state in [SessionState.COMPLETED, SessionState.CANCELLED]: return self.state = SessionState.TIMED_OUT self.error_message = "Session timed out" self.timing.completed_at = datetime.now() pprint(f"Session {self.session_id} timed out") self._cleanup_timers() self._notify_callback("session_timeout") def execute_arbitration(self) -> SelectionResult: """ Execute arbitration algorithm to select best satellite. Returns: SelectionResult: Result of satellite selection Raises: SessionStateError: If session is not ready for arbitration SessionError: If arbitration fails """ with self._lock: if self.state != SessionState.READY_FOR_ARBITRATION: raise SessionStateError(f"Cannot execute arbitration in state {self.state}") if not self.wakeword_reports: raise SessionError("No wakeword reports available for arbitration") self.state = SessionState.ARBITRATING self.timing.arbitration_started_at = datetime.now() pprint(f"Executing arbitration for session {self.session_id} " f"with {len(self.wakeword_reports)} reports") try: # Update metrics from reports self.metrics.update_from_reports(self.wakeword_reports) # Execute primary algorithm primary_algorithm = self.config.primary_algorithm self.metrics.algorithms_attempted.append(primary_algorithm.value) self.selection_result = self.algorithms.select_satellite( self.wakeword_reports, primary_algorithm ) self.selected_satellite_id = self.selection_result.selected_satellite_id self.metrics.selection_confidence = self.selection_result.confidence self.timing.arbitration_ended_at = datetime.now() self.state = SessionState.COMPLETED self.phase = SessionPhase.RESULT_PROCESSING pprint(f"Arbitration completed for session {self.session_id}: " f"selected {self.selected_satellite_id} " f"(score: {self.selection_result.selection_score:.3f})") self._cleanup_timers() self._notify_callback("arbitration_completed", self.selection_result) return self.selection_result except Exception as e: # Try fallback algorithm if primary fails if (self.config.fallback_algorithm != primary_algorithm and self.config.fallback_algorithm.value not in self.metrics.algorithms_attempted): try: pprint(f"Primary arbitration failed for session {self.session_id}, trying fallback") self.metrics.algorithms_attempted.append(self.config.fallback_algorithm.value) self.selection_result = self.algorithms.select_satellite( self.wakeword_reports, self.config.fallback_algorithm ) self.selected_satellite_id = self.selection_result.selected_satellite_id self.metrics.selection_confidence = self.selection_result.confidence self.timing.arbitration_ended_at = datetime.now() self.state = SessionState.COMPLETED pprint(f"Fallback arbitration completed for session {self.session_id}: " f"selected {self.selected_satellite_id}") self._cleanup_timers() self._notify_callback("arbitration_completed", self.selection_result) return self.selection_result except Exception as fallback_error: self._handle_error(f"Both primary and fallback arbitration failed: {fallback_error}") raise SessionError(f"Arbitration failed: {fallback_error}") else: self._handle_error(f"Arbitration failed: {e}") raise SessionError(f"Arbitration failed: {e}") def cancel(self, reason: str = "Cancelled by user") -> None: """ Cancel the arbitration session. Args: reason: Reason for cancellation """ with self._lock: if self.state in [SessionState.COMPLETED, SessionState.CANCELLED]: return self._is_cancelled = True self.state = SessionState.CANCELLED self.error_message = reason self.timing.completed_at = datetime.now() pprint(f"Cancelled session {self.session_id}: {reason}") self._cleanup_timers() self._notify_callback("session_cancelled", reason) def complete(self, conversation_id: Optional[str] = None) -> None: """ Mark session as completed and perform cleanup. Args: conversation_id: ID of conversation that was started from this session """ with self._lock: if self.state == SessionState.COMPLETED: if conversation_id: self.conversation_id = conversation_id pprint(f"Updated session {self.session_id} with conversation ID {conversation_id}") return if self.state != SessionState.ARBITRATING: pprint(f"Warning: Completing session {self.session_id} in state {self.state}") self.conversation_id = conversation_id self.state = SessionState.COMPLETED self.phase = SessionPhase.CLEANUP self.timing.completed_at = datetime.now() pprint(f"Completed session {self.session_id} " f"(conversation: {conversation_id or 'none'})") self._cleanup_timers() self._notify_callback("session_completed", conversation_id) def get_session_summary(self) -> Dict[str, Any]: """ Get comprehensive session summary. Returns: Dict[str, Any]: Session summary with all relevant information """ with self._lock: return { 'session_id': self.session_id, 'state': self.state.value, 'phase': self.phase.value, 'error_message': self.error_message, 'timing': self.timing.to_dict(), 'metrics': self.metrics.to_dict(), 'wakeword_reports': [r.to_dict() for r in self.wakeword_reports], 'selection_result': self.selection_result.to_dict() if self.selection_result else None, 'selected_satellite_id': self.selected_satellite_id, 'conversation_id': self.conversation_id, 'metadata': self.metadata, } def _handle_error(self, error_message: str) -> None: """Handle session error.""" self.state = SessionState.ERROR self.error_message = error_message self.timing.completed_at = datetime.now() pprint(f"Error in session {self.session_id}: {error_message}") self._cleanup_timers() self._notify_callback("session_error", error_message) def _cleanup_timers(self) -> None: """Clean up any active timers.""" if self._collection_timer: self._collection_timer.cancel() self._collection_timer = None if self._timeout_timer: self._timeout_timer.cancel() self._timeout_timer = None def _notify_callback(self, event_type: str, data: Any = None) -> None: """Notify session callback of events.""" if self.callback: try: self.callback(self, event_type, data) except Exception as e: pprint(f"Callback error for session {self.session_id}: {e}") def __del__(self): """Destructor to ensure cleanup.""" self._cleanup_timers() class ArbitrationSessionManager: """ Manager for arbitration sessions. This class handles creation, tracking, and lifecycle management of arbitration sessions. It provides the main interface for starting arbitration processes when wakeword conflicts occur. """ def __init__( self, config: ArbitrationConfig, algorithms: ArbitrationAlgorithms, max_concurrent_sessions: Optional[int] = None ): """ Initialize session manager. Args: config: Arbitration configuration algorithms: Arbitration algorithms processor max_concurrent_sessions: Maximum concurrent sessions (defaults to config) """ self.config = config self.algorithms = algorithms self.max_concurrent_sessions = ( max_concurrent_sessions or config.max_concurrent_arbitrations ) # Session tracking self.active_sessions: Dict[str, ArbitrationSession] = {} self.completed_sessions: Dict[str, ArbitrationSession] = {} self.session_history: List[str] = [] # Threading self._lock = threading.RLock() self.executor = ThreadPoolExecutor(max_workers=config.thread_pool_size) # Statistics self.stats = { 'sessions_created': 0, 'sessions_completed': 0, 'sessions_cancelled': 0, 'sessions_timed_out': 0, 'sessions_error': 0, 'total_reports_processed': 0, } pprint(f"Initialized arbitration session manager " f"(max concurrent: {self.max_concurrent_sessions})") def create_session( self, session_id: Optional[str] = None, trigger_report: Optional[WakewordReport] = None, callback: Optional[Callable] = None ) -> ArbitrationSession: """ Create a new arbitration session. Args: session_id: Session ID (generated if None) trigger_report: Initial wakeword report callback: Session event callback Returns: ArbitrationSession: New arbitration session Raises: SessionConfigurationError: If session cannot be created """ with self._lock: # Check concurrent session limit if len(self.active_sessions) >= self.max_concurrent_sessions: self._cleanup_completed_sessions() if len(self.active_sessions) >= self.max_concurrent_sessions: raise SessionConfigurationError( f"Maximum concurrent sessions ({self.max_concurrent_sessions}) reached" ) # Generate session ID if needed if session_id is None: session_id = f"arb_{int(time.time() * 1000)}_{uuid.uuid4().hex[:8]}" if session_id in self.active_sessions: raise SessionConfigurationError(f"Session {session_id} already exists") # Create session session = ArbitrationSession( session_id=session_id, config=self.config, algorithms=self.algorithms, trigger_report=trigger_report, callback=self._session_callback ) # Add to tracking self.active_sessions[session_id] = session self.session_history.append(session_id) self.stats['sessions_created'] += 1 # Keep history within limits if len(self.session_history) > self.config.max_history_entries: old_session_id = self.session_history.pop(0) if old_session_id in self.completed_sessions: del self.completed_sessions[old_session_id] pprint(f"Created arbitration session {session_id} " f"({len(self.active_sessions)} active)") return session def start_arbitration( self, trigger_report: WakewordReport, session_id: Optional[str] = None ) -> ArbitrationSession: """ Start arbitration process for a wakeword conflict. This is the main entry point for starting arbitration when multiple satellites detect the wakeword. Args: trigger_report: Initial wakeword report that triggered arbitration session_id: Optional session ID Returns: ArbitrationSession: Started arbitration session """ session = self.create_session( session_id=session_id, trigger_report=trigger_report ) # Start collection window session.start_collection() pprint(f"Started arbitration for wakeword from {trigger_report.satellite_id}") return session def add_wakeword_report( self, session_id: str, report: WakewordReport ) -> bool: """ Add wakeword report to existing session. Args: session_id: Target session ID report: Wakeword report to add Returns: bool: True if report was added successfully """ with self._lock: if session_id not in self.active_sessions: pprint(f"Cannot add report to unknown session {session_id}") return False session = self.active_sessions[session_id] success = session.add_wakeword_report(report) if success: self.stats['total_reports_processed'] += 1 return success def get_session(self, session_id: str) -> Optional[ArbitrationSession]: """ Get session by ID. Args: session_id: Session ID to retrieve Returns: ArbitrationSession: Session if found, None otherwise """ with self._lock: return (self.active_sessions.get(session_id) or self.completed_sessions.get(session_id)) def get_active_sessions(self) -> List[ArbitrationSession]: """Get list of all active sessions.""" with self._lock: return list(self.active_sessions.values()) def cancel_session(self, session_id: str, reason: str = "Cancelled") -> bool: """ Cancel an arbitration session. Args: session_id: Session to cancel reason: Cancellation reason Returns: bool: True if session was cancelled """ with self._lock: if session_id not in self.active_sessions: return False session = self.active_sessions[session_id] session.cancel(reason) return True def get_session_statistics(self) -> Dict[str, Any]: """ Get comprehensive session statistics. Returns: Dict[str, Any]: Session statistics and metrics """ with self._lock: active_count = len(self.active_sessions) completed_count = len(self.completed_sessions) # Calculate success rate total_finished = (self.stats['sessions_completed'] + self.stats['sessions_cancelled'] + self.stats['sessions_timed_out'] + self.stats['sessions_error']) success_rate = ( self.stats['sessions_completed'] / total_finished if total_finished > 0 else 0.0 ) return { 'active_sessions': active_count, 'completed_sessions': completed_count, 'total_sessions_created': self.stats['sessions_created'], 'sessions_completed': self.stats['sessions_completed'], 'sessions_cancelled': self.stats['sessions_cancelled'], 'sessions_timed_out': self.stats['sessions_timed_out'], 'sessions_error': self.stats['sessions_error'], 'success_rate': success_rate, 'total_reports_processed': self.stats['total_reports_processed'], 'max_concurrent_sessions': self.max_concurrent_sessions, } def _session_callback( self, session: ArbitrationSession, event_type: str, data: Any = None ) -> None: """Handle session events.""" session_id = session.session_id if event_type == "session_completed": self._move_to_completed(session_id) self.stats['sessions_completed'] += 1 elif event_type == "session_cancelled": self._move_to_completed(session_id) self.stats['sessions_cancelled'] += 1 elif event_type == "session_timeout": self._move_to_completed(session_id) self.stats['sessions_timed_out'] += 1 elif event_type == "session_error": self._move_to_completed(session_id) self.stats['sessions_error'] += 1 def _move_to_completed(self, session_id: str) -> None: """Move session from active to completed.""" with self._lock: if session_id in self.active_sessions: session = self.active_sessions.pop(session_id) self.completed_sessions[session_id] = session pprint(f"Moved session {session_id} to completed " f"({len(self.active_sessions)} active remaining)") def _cleanup_completed_sessions(self) -> None: """Clean up old completed sessions.""" with self._lock: # Remove sessions older than configured limit cutoff_time = datetime.now() - timedelta(hours=1) # Keep 1 hour of history to_remove = [] for session_id, session in self.completed_sessions.items(): if (session.timing.completed_at and session.timing.completed_at < cutoff_time): to_remove.append(session_id) for session_id in to_remove: del self.completed_sessions[session_id] if session_id in self.session_history: self.session_history.remove(session_id) if to_remove: pprint(f"Cleaned up {len(to_remove)} old completed sessions") def cleanup(self) -> None: """Clean up session manager resources.""" with self._lock: # Cancel all active sessions for session in list(self.active_sessions.values()): session.cancel("Manager shutdown") # Clean up executor if self.executor: self.executor.shutdown(wait=True) pprint("Arbitration session manager cleaned up") def __del__(self): """Destructor to ensure cleanup.""" self.cleanup() def create_session_manager( config: ArbitrationConfig, algorithms: ArbitrationAlgorithms ) -> ArbitrationSessionManager: """ Create arbitration session manager with given configuration. Args: config: Arbitration configuration algorithms: Arbitration algorithms processor Returns: ArbitrationSessionManager: Configured session manager """ return ArbitrationSessionManager(config, algorithms)