| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891 |
- """
- 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)
|