arbitration_session.py 33 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891
  1. """
  2. Arbitration Session Management for Trixy
  3. This module provides comprehensive session tracking and management for arbitration
  4. processes. It handles the lifecycle of arbitration sessions from wakeword detection
  5. through satellite selection and conversation handoff.
  6. The session system implements the 1-second collection window specified in CLAUDE.md
  7. and tracks all aspects of the arbitration process including:
  8. - Multiple satellite wakeword reports collection
  9. - Timing and timeout management
  10. - Selection algorithm execution
  11. - Result tracking and analytics
  12. - Session state management
  13. - Integration with conversation system
  14. Usage:
  15. from trixy_core.arbitration import ArbitrationSession, ArbitrationSessionManager
  16. # Create session manager
  17. session_manager = ArbitrationSessionManager(config, algorithms)
  18. # Start new arbitration session
  19. session = session_manager.start_session("session_123")
  20. # Add wakeword reports
  21. session.add_wakeword_report(report1)
  22. session.add_wakeword_report(report2)
  23. # Execute arbitration after collection window
  24. result = session.execute_arbitration()
  25. # Complete session
  26. session.complete(conversation_id="conv_456")
  27. """
  28. import asyncio
  29. import time
  30. import uuid
  31. from dataclasses import dataclass, field
  32. from datetime import datetime, timedelta
  33. from enum import Enum
  34. from typing import Dict, List, Optional, Any, Callable, Union
  35. import threading
  36. from concurrent.futures import ThreadPoolExecutor, Future
  37. from .arbitration_config import ArbitrationConfig, ArbitrationAlgorithm
  38. from .arbitration_algorithms import (
  39. ArbitrationAlgorithms, WakewordReport, SelectionResult,
  40. ArbitrationConfigError
  41. )
  42. def pprint(message: str) -> None:
  43. """Arbitration session logging function."""
  44. print(f"[ARBITRATION_SESSION] {message}")
  45. class SessionState(Enum):
  46. """States of an arbitration session."""
  47. CREATED = "created"
  48. COLLECTING = "collecting"
  49. READY_FOR_ARBITRATION = "ready_for_arbitration"
  50. ARBITRATING = "arbitrating"
  51. COMPLETED = "completed"
  52. CANCELLED = "cancelled"
  53. TIMED_OUT = "timed_out"
  54. ERROR = "error"
  55. class SessionPhase(Enum):
  56. """Phases within an arbitration session."""
  57. INITIALIZATION = "initialization"
  58. COLLECTION_WINDOW = "collection_window"
  59. SELECTION_PROCESS = "selection_process"
  60. RESULT_PROCESSING = "result_processing"
  61. HANDOFF = "handoff"
  62. CLEANUP = "cleanup"
  63. @dataclass
  64. class SessionTiming:
  65. """Timing information for an arbitration session."""
  66. created_at: datetime = field(default_factory=datetime.now)
  67. collection_started_at: Optional[datetime] = None
  68. collection_ended_at: Optional[datetime] = None
  69. arbitration_started_at: Optional[datetime] = None
  70. arbitration_ended_at: Optional[datetime] = None
  71. completed_at: Optional[datetime] = None
  72. def get_collection_duration(self) -> Optional[float]:
  73. """Get collection window duration in seconds."""
  74. if self.collection_started_at and self.collection_ended_at:
  75. return (self.collection_ended_at - self.collection_started_at).total_seconds()
  76. return None
  77. def get_arbitration_duration(self) -> Optional[float]:
  78. """Get arbitration processing duration in seconds."""
  79. if self.arbitration_started_at and self.arbitration_ended_at:
  80. return (self.arbitration_ended_at - self.arbitration_started_at).total_seconds()
  81. return None
  82. def get_total_duration(self) -> Optional[float]:
  83. """Get total session duration in seconds."""
  84. if self.completed_at:
  85. return (self.completed_at - self.created_at).total_seconds()
  86. return None
  87. def to_dict(self) -> Dict[str, Any]:
  88. """Convert timing to dictionary representation."""
  89. return {
  90. 'created_at': self.created_at.isoformat(),
  91. 'collection_started_at': self.collection_started_at.isoformat() if self.collection_started_at else None,
  92. 'collection_ended_at': self.collection_ended_at.isoformat() if self.collection_ended_at else None,
  93. 'arbitration_started_at': self.arbitration_started_at.isoformat() if self.arbitration_started_at else None,
  94. 'arbitration_ended_at': self.arbitration_ended_at.isoformat() if self.arbitration_ended_at else None,
  95. 'completed_at': self.completed_at.isoformat() if self.completed_at else None,
  96. 'collection_duration': self.get_collection_duration(),
  97. 'arbitration_duration': self.get_arbitration_duration(),
  98. 'total_duration': self.get_total_duration(),
  99. }
  100. @dataclass
  101. class SessionMetrics:
  102. """Metrics and analytics for an arbitration session."""
  103. total_reports_received: int = 0
  104. unique_satellites_count: int = 0
  105. unique_rooms_count: int = 0
  106. average_volume: float = 0.0
  107. average_confidence: float = 0.0
  108. volume_range: tuple = field(default_factory=lambda: (0.0, 0.0))
  109. confidence_range: tuple = field(default_factory=lambda: (0.0, 0.0))
  110. algorithms_attempted: List[str] = field(default_factory=list)
  111. selection_confidence: float = 0.0
  112. def update_from_reports(self, reports: List[WakewordReport]) -> None:
  113. """Update metrics from wakeword reports."""
  114. if not reports:
  115. return
  116. self.total_reports_received = len(reports)
  117. self.unique_satellites_count = len(set(r.satellite_id for r in reports))
  118. self.unique_rooms_count = len(set(r.room_id for r in reports))
  119. volumes = [r.volume for r in reports]
  120. confidences = [r.confidence for r in reports]
  121. self.average_volume = sum(volumes) / len(volumes)
  122. self.average_confidence = sum(confidences) / len(confidences)
  123. self.volume_range = (min(volumes), max(volumes))
  124. self.confidence_range = (min(confidences), max(confidences))
  125. def to_dict(self) -> Dict[str, Any]:
  126. """Convert metrics to dictionary representation."""
  127. return {
  128. 'total_reports_received': self.total_reports_received,
  129. 'unique_satellites_count': self.unique_satellites_count,
  130. 'unique_rooms_count': self.unique_rooms_count,
  131. 'average_volume': self.average_volume,
  132. 'average_confidence': self.average_confidence,
  133. 'volume_range': self.volume_range,
  134. 'confidence_range': self.confidence_range,
  135. 'algorithms_attempted': self.algorithms_attempted,
  136. 'selection_confidence': self.selection_confidence,
  137. }
  138. class SessionError(Exception):
  139. """Base exception for arbitration session errors."""
  140. pass
  141. class SessionTimeoutError(SessionError):
  142. """Raised when session times out."""
  143. pass
  144. class SessionStateError(SessionError):
  145. """Raised when session is in invalid state for operation."""
  146. pass
  147. class SessionConfigurationError(SessionError):
  148. """Raised when session configuration is invalid."""
  149. pass
  150. class ArbitrationSession:
  151. """
  152. Individual arbitration session for managing wakeword conflict resolution.
  153. This class handles a single arbitration session from creation through completion,
  154. implementing the 1-second collection window and satellite selection process
  155. as specified in CLAUDE.md.
  156. """
  157. def __init__(
  158. self,
  159. session_id: str,
  160. config: ArbitrationConfig,
  161. algorithms: ArbitrationAlgorithms,
  162. trigger_report: Optional[WakewordReport] = None,
  163. callback: Optional[Callable] = None
  164. ):
  165. """
  166. Initialize arbitration session.
  167. Args:
  168. session_id: Unique session identifier
  169. config: Arbitration configuration
  170. algorithms: Arbitration algorithms processor
  171. trigger_report: Initial wakeword report that triggered session
  172. callback: Optional callback for session events
  173. """
  174. self.session_id = session_id
  175. self.config = config
  176. self.algorithms = algorithms
  177. self.callback = callback
  178. # Session state
  179. self.state = SessionState.CREATED
  180. self.phase = SessionPhase.INITIALIZATION
  181. self.error_message: Optional[str] = None
  182. # Timing and metrics
  183. self.timing = SessionTiming()
  184. self.metrics = SessionMetrics()
  185. # Wakeword reports
  186. self.wakeword_reports: List[WakewordReport] = []
  187. if trigger_report:
  188. self.wakeword_reports.append(trigger_report)
  189. # Selection result
  190. self.selection_result: Optional[SelectionResult] = None
  191. self.selected_satellite_id: Optional[str] = None
  192. self.conversation_id: Optional[str] = None
  193. # Threading and timing
  194. self._lock = threading.RLock()
  195. self._collection_timer: Optional[threading.Timer] = None
  196. self._timeout_timer: Optional[threading.Timer] = None
  197. self._is_cancelled = False
  198. # Metadata
  199. self.metadata: Dict[str, Any] = {}
  200. pprint(f"Created arbitration session {session_id}")
  201. def start_collection(self) -> None:
  202. """
  203. Start the collection window for wakeword reports.
  204. This implements the 1-second collection window specified in CLAUDE.md.
  205. """
  206. with self._lock:
  207. if self.state != SessionState.CREATED:
  208. raise SessionStateError(f"Cannot start collection in state {self.state}")
  209. self.state = SessionState.COLLECTING
  210. self.phase = SessionPhase.COLLECTION_WINDOW
  211. self.timing.collection_started_at = datetime.now()
  212. # Start collection timer
  213. self._collection_timer = threading.Timer(
  214. self.config.collection_window_seconds,
  215. self._collection_timeout
  216. )
  217. self._collection_timer.start()
  218. # Start overall timeout timer
  219. self._timeout_timer = threading.Timer(
  220. self.config.arbitration_timeout_seconds,
  221. self._session_timeout
  222. )
  223. self._timeout_timer.start()
  224. pprint(f"Started collection window for session {self.session_id} "
  225. f"({self.config.collection_window_seconds}s)")
  226. self._notify_callback("collection_started")
  227. def add_wakeword_report(self, report: WakewordReport) -> bool:
  228. """
  229. Add a wakeword report to the session.
  230. Args:
  231. report: Wakeword detection report
  232. Returns:
  233. bool: True if report was added, False if rejected
  234. """
  235. with self._lock:
  236. if self.state not in [SessionState.CREATED, SessionState.COLLECTING]:
  237. pprint(f"Rejected report for session {self.session_id} - wrong state {self.state}")
  238. return False
  239. if self._is_cancelled:
  240. pprint(f"Rejected report for session {self.session_id} - session cancelled")
  241. return False
  242. # Validate report timestamp (should be recent)
  243. current_time = time.time()
  244. report_age = current_time - report.timestamp
  245. if report_age > self.config.max_wait_time:
  246. pprint(f"Rejected old report for session {self.session_id} - age {report_age:.2f}s")
  247. return False
  248. # Check for duplicate satellite reports
  249. existing_satellite_ids = [r.satellite_id for r in self.wakeword_reports]
  250. if report.satellite_id in existing_satellite_ids:
  251. pprint(f"Rejected duplicate report from {report.satellite_id} for session {self.session_id}")
  252. return False
  253. # Add report
  254. self.wakeword_reports.append(report)
  255. pprint(f"Added wakeword report from {report.satellite_id} to session {self.session_id} "
  256. f"(total: {len(self.wakeword_reports)})")
  257. # Start collection if this is the first report
  258. if len(self.wakeword_reports) == 1 and self.state == SessionState.CREATED:
  259. self.start_collection()
  260. self._notify_callback("report_added", report)
  261. return True
  262. def _collection_timeout(self) -> None:
  263. """Handle collection window timeout."""
  264. with self._lock:
  265. if self.state != SessionState.COLLECTING:
  266. return
  267. self.timing.collection_ended_at = datetime.now()
  268. self.state = SessionState.READY_FOR_ARBITRATION
  269. self.phase = SessionPhase.SELECTION_PROCESS
  270. pprint(f"Collection window ended for session {self.session_id} "
  271. f"with {len(self.wakeword_reports)} reports")
  272. self._notify_callback("collection_ended")
  273. # Automatically execute arbitration if we have reports
  274. if self.wakeword_reports:
  275. try:
  276. self.execute_arbitration()
  277. except Exception as e:
  278. self._handle_error(f"Auto-arbitration failed: {e}")
  279. else:
  280. self._handle_error("No wakeword reports received during collection window")
  281. def _session_timeout(self) -> None:
  282. """Handle overall session timeout."""
  283. with self._lock:
  284. if self.state in [SessionState.COMPLETED, SessionState.CANCELLED]:
  285. return
  286. self.state = SessionState.TIMED_OUT
  287. self.error_message = "Session timed out"
  288. self.timing.completed_at = datetime.now()
  289. pprint(f"Session {self.session_id} timed out")
  290. self._cleanup_timers()
  291. self._notify_callback("session_timeout")
  292. def execute_arbitration(self) -> SelectionResult:
  293. """
  294. Execute arbitration algorithm to select best satellite.
  295. Returns:
  296. SelectionResult: Result of satellite selection
  297. Raises:
  298. SessionStateError: If session is not ready for arbitration
  299. SessionError: If arbitration fails
  300. """
  301. with self._lock:
  302. if self.state != SessionState.READY_FOR_ARBITRATION:
  303. raise SessionStateError(f"Cannot execute arbitration in state {self.state}")
  304. if not self.wakeword_reports:
  305. raise SessionError("No wakeword reports available for arbitration")
  306. self.state = SessionState.ARBITRATING
  307. self.timing.arbitration_started_at = datetime.now()
  308. pprint(f"Executing arbitration for session {self.session_id} "
  309. f"with {len(self.wakeword_reports)} reports")
  310. try:
  311. # Update metrics from reports
  312. self.metrics.update_from_reports(self.wakeword_reports)
  313. # Execute primary algorithm
  314. primary_algorithm = self.config.primary_algorithm
  315. self.metrics.algorithms_attempted.append(primary_algorithm.value)
  316. self.selection_result = self.algorithms.select_satellite(
  317. self.wakeword_reports,
  318. primary_algorithm
  319. )
  320. self.selected_satellite_id = self.selection_result.selected_satellite_id
  321. self.metrics.selection_confidence = self.selection_result.confidence
  322. self.timing.arbitration_ended_at = datetime.now()
  323. self.state = SessionState.COMPLETED
  324. self.phase = SessionPhase.RESULT_PROCESSING
  325. pprint(f"Arbitration completed for session {self.session_id}: "
  326. f"selected {self.selected_satellite_id} "
  327. f"(score: {self.selection_result.selection_score:.3f})")
  328. self._cleanup_timers()
  329. self._notify_callback("arbitration_completed", self.selection_result)
  330. return self.selection_result
  331. except Exception as e:
  332. # Try fallback algorithm if primary fails
  333. if (self.config.fallback_algorithm != primary_algorithm and
  334. self.config.fallback_algorithm.value not in self.metrics.algorithms_attempted):
  335. try:
  336. pprint(f"Primary arbitration failed for session {self.session_id}, trying fallback")
  337. self.metrics.algorithms_attempted.append(self.config.fallback_algorithm.value)
  338. self.selection_result = self.algorithms.select_satellite(
  339. self.wakeword_reports,
  340. self.config.fallback_algorithm
  341. )
  342. self.selected_satellite_id = self.selection_result.selected_satellite_id
  343. self.metrics.selection_confidence = self.selection_result.confidence
  344. self.timing.arbitration_ended_at = datetime.now()
  345. self.state = SessionState.COMPLETED
  346. pprint(f"Fallback arbitration completed for session {self.session_id}: "
  347. f"selected {self.selected_satellite_id}")
  348. self._cleanup_timers()
  349. self._notify_callback("arbitration_completed", self.selection_result)
  350. return self.selection_result
  351. except Exception as fallback_error:
  352. self._handle_error(f"Both primary and fallback arbitration failed: {fallback_error}")
  353. raise SessionError(f"Arbitration failed: {fallback_error}")
  354. else:
  355. self._handle_error(f"Arbitration failed: {e}")
  356. raise SessionError(f"Arbitration failed: {e}")
  357. def cancel(self, reason: str = "Cancelled by user") -> None:
  358. """
  359. Cancel the arbitration session.
  360. Args:
  361. reason: Reason for cancellation
  362. """
  363. with self._lock:
  364. if self.state in [SessionState.COMPLETED, SessionState.CANCELLED]:
  365. return
  366. self._is_cancelled = True
  367. self.state = SessionState.CANCELLED
  368. self.error_message = reason
  369. self.timing.completed_at = datetime.now()
  370. pprint(f"Cancelled session {self.session_id}: {reason}")
  371. self._cleanup_timers()
  372. self._notify_callback("session_cancelled", reason)
  373. def complete(self, conversation_id: Optional[str] = None) -> None:
  374. """
  375. Mark session as completed and perform cleanup.
  376. Args:
  377. conversation_id: ID of conversation that was started from this session
  378. """
  379. with self._lock:
  380. if self.state == SessionState.COMPLETED:
  381. if conversation_id:
  382. self.conversation_id = conversation_id
  383. pprint(f"Updated session {self.session_id} with conversation ID {conversation_id}")
  384. return
  385. if self.state != SessionState.ARBITRATING:
  386. pprint(f"Warning: Completing session {self.session_id} in state {self.state}")
  387. self.conversation_id = conversation_id
  388. self.state = SessionState.COMPLETED
  389. self.phase = SessionPhase.CLEANUP
  390. self.timing.completed_at = datetime.now()
  391. pprint(f"Completed session {self.session_id} "
  392. f"(conversation: {conversation_id or 'none'})")
  393. self._cleanup_timers()
  394. self._notify_callback("session_completed", conversation_id)
  395. def get_session_summary(self) -> Dict[str, Any]:
  396. """
  397. Get comprehensive session summary.
  398. Returns:
  399. Dict[str, Any]: Session summary with all relevant information
  400. """
  401. with self._lock:
  402. return {
  403. 'session_id': self.session_id,
  404. 'state': self.state.value,
  405. 'phase': self.phase.value,
  406. 'error_message': self.error_message,
  407. 'timing': self.timing.to_dict(),
  408. 'metrics': self.metrics.to_dict(),
  409. 'wakeword_reports': [r.to_dict() for r in self.wakeword_reports],
  410. 'selection_result': self.selection_result.to_dict() if self.selection_result else None,
  411. 'selected_satellite_id': self.selected_satellite_id,
  412. 'conversation_id': self.conversation_id,
  413. 'metadata': self.metadata,
  414. }
  415. def _handle_error(self, error_message: str) -> None:
  416. """Handle session error."""
  417. self.state = SessionState.ERROR
  418. self.error_message = error_message
  419. self.timing.completed_at = datetime.now()
  420. pprint(f"Error in session {self.session_id}: {error_message}")
  421. self._cleanup_timers()
  422. self._notify_callback("session_error", error_message)
  423. def _cleanup_timers(self) -> None:
  424. """Clean up any active timers."""
  425. if self._collection_timer:
  426. self._collection_timer.cancel()
  427. self._collection_timer = None
  428. if self._timeout_timer:
  429. self._timeout_timer.cancel()
  430. self._timeout_timer = None
  431. def _notify_callback(self, event_type: str, data: Any = None) -> None:
  432. """Notify session callback of events."""
  433. if self.callback:
  434. try:
  435. self.callback(self, event_type, data)
  436. except Exception as e:
  437. pprint(f"Callback error for session {self.session_id}: {e}")
  438. def __del__(self):
  439. """Destructor to ensure cleanup."""
  440. self._cleanup_timers()
  441. class ArbitrationSessionManager:
  442. """
  443. Manager for arbitration sessions.
  444. This class handles creation, tracking, and lifecycle management of
  445. arbitration sessions. It provides the main interface for starting
  446. arbitration processes when wakeword conflicts occur.
  447. """
  448. def __init__(
  449. self,
  450. config: ArbitrationConfig,
  451. algorithms: ArbitrationAlgorithms,
  452. max_concurrent_sessions: Optional[int] = None
  453. ):
  454. """
  455. Initialize session manager.
  456. Args:
  457. config: Arbitration configuration
  458. algorithms: Arbitration algorithms processor
  459. max_concurrent_sessions: Maximum concurrent sessions (defaults to config)
  460. """
  461. self.config = config
  462. self.algorithms = algorithms
  463. self.max_concurrent_sessions = (
  464. max_concurrent_sessions or config.max_concurrent_arbitrations
  465. )
  466. # Session tracking
  467. self.active_sessions: Dict[str, ArbitrationSession] = {}
  468. self.completed_sessions: Dict[str, ArbitrationSession] = {}
  469. self.session_history: List[str] = []
  470. # Threading
  471. self._lock = threading.RLock()
  472. self.executor = ThreadPoolExecutor(max_workers=config.thread_pool_size)
  473. # Statistics
  474. self.stats = {
  475. 'sessions_created': 0,
  476. 'sessions_completed': 0,
  477. 'sessions_cancelled': 0,
  478. 'sessions_timed_out': 0,
  479. 'sessions_error': 0,
  480. 'total_reports_processed': 0,
  481. }
  482. pprint(f"Initialized arbitration session manager "
  483. f"(max concurrent: {self.max_concurrent_sessions})")
  484. def create_session(
  485. self,
  486. session_id: Optional[str] = None,
  487. trigger_report: Optional[WakewordReport] = None,
  488. callback: Optional[Callable] = None
  489. ) -> ArbitrationSession:
  490. """
  491. Create a new arbitration session.
  492. Args:
  493. session_id: Session ID (generated if None)
  494. trigger_report: Initial wakeword report
  495. callback: Session event callback
  496. Returns:
  497. ArbitrationSession: New arbitration session
  498. Raises:
  499. SessionConfigurationError: If session cannot be created
  500. """
  501. with self._lock:
  502. # Check concurrent session limit
  503. if len(self.active_sessions) >= self.max_concurrent_sessions:
  504. self._cleanup_completed_sessions()
  505. if len(self.active_sessions) >= self.max_concurrent_sessions:
  506. raise SessionConfigurationError(
  507. f"Maximum concurrent sessions ({self.max_concurrent_sessions}) reached"
  508. )
  509. # Generate session ID if needed
  510. if session_id is None:
  511. session_id = f"arb_{int(time.time() * 1000)}_{uuid.uuid4().hex[:8]}"
  512. if session_id in self.active_sessions:
  513. raise SessionConfigurationError(f"Session {session_id} already exists")
  514. # Create session
  515. session = ArbitrationSession(
  516. session_id=session_id,
  517. config=self.config,
  518. algorithms=self.algorithms,
  519. trigger_report=trigger_report,
  520. callback=self._session_callback
  521. )
  522. # Add to tracking
  523. self.active_sessions[session_id] = session
  524. self.session_history.append(session_id)
  525. self.stats['sessions_created'] += 1
  526. # Keep history within limits
  527. if len(self.session_history) > self.config.max_history_entries:
  528. old_session_id = self.session_history.pop(0)
  529. if old_session_id in self.completed_sessions:
  530. del self.completed_sessions[old_session_id]
  531. pprint(f"Created arbitration session {session_id} "
  532. f"({len(self.active_sessions)} active)")
  533. return session
  534. def start_arbitration(
  535. self,
  536. trigger_report: WakewordReport,
  537. session_id: Optional[str] = None
  538. ) -> ArbitrationSession:
  539. """
  540. Start arbitration process for a wakeword conflict.
  541. This is the main entry point for starting arbitration when
  542. multiple satellites detect the wakeword.
  543. Args:
  544. trigger_report: Initial wakeword report that triggered arbitration
  545. session_id: Optional session ID
  546. Returns:
  547. ArbitrationSession: Started arbitration session
  548. """
  549. session = self.create_session(
  550. session_id=session_id,
  551. trigger_report=trigger_report
  552. )
  553. # Start collection window
  554. session.start_collection()
  555. pprint(f"Started arbitration for wakeword from {trigger_report.satellite_id}")
  556. return session
  557. def add_wakeword_report(
  558. self,
  559. session_id: str,
  560. report: WakewordReport
  561. ) -> bool:
  562. """
  563. Add wakeword report to existing session.
  564. Args:
  565. session_id: Target session ID
  566. report: Wakeword report to add
  567. Returns:
  568. bool: True if report was added successfully
  569. """
  570. with self._lock:
  571. if session_id not in self.active_sessions:
  572. pprint(f"Cannot add report to unknown session {session_id}")
  573. return False
  574. session = self.active_sessions[session_id]
  575. success = session.add_wakeword_report(report)
  576. if success:
  577. self.stats['total_reports_processed'] += 1
  578. return success
  579. def get_session(self, session_id: str) -> Optional[ArbitrationSession]:
  580. """
  581. Get session by ID.
  582. Args:
  583. session_id: Session ID to retrieve
  584. Returns:
  585. ArbitrationSession: Session if found, None otherwise
  586. """
  587. with self._lock:
  588. return (self.active_sessions.get(session_id) or
  589. self.completed_sessions.get(session_id))
  590. def get_active_sessions(self) -> List[ArbitrationSession]:
  591. """Get list of all active sessions."""
  592. with self._lock:
  593. return list(self.active_sessions.values())
  594. def cancel_session(self, session_id: str, reason: str = "Cancelled") -> bool:
  595. """
  596. Cancel an arbitration session.
  597. Args:
  598. session_id: Session to cancel
  599. reason: Cancellation reason
  600. Returns:
  601. bool: True if session was cancelled
  602. """
  603. with self._lock:
  604. if session_id not in self.active_sessions:
  605. return False
  606. session = self.active_sessions[session_id]
  607. session.cancel(reason)
  608. return True
  609. def get_session_statistics(self) -> Dict[str, Any]:
  610. """
  611. Get comprehensive session statistics.
  612. Returns:
  613. Dict[str, Any]: Session statistics and metrics
  614. """
  615. with self._lock:
  616. active_count = len(self.active_sessions)
  617. completed_count = len(self.completed_sessions)
  618. # Calculate success rate
  619. total_finished = (self.stats['sessions_completed'] +
  620. self.stats['sessions_cancelled'] +
  621. self.stats['sessions_timed_out'] +
  622. self.stats['sessions_error'])
  623. success_rate = (
  624. self.stats['sessions_completed'] / total_finished
  625. if total_finished > 0 else 0.0
  626. )
  627. return {
  628. 'active_sessions': active_count,
  629. 'completed_sessions': completed_count,
  630. 'total_sessions_created': self.stats['sessions_created'],
  631. 'sessions_completed': self.stats['sessions_completed'],
  632. 'sessions_cancelled': self.stats['sessions_cancelled'],
  633. 'sessions_timed_out': self.stats['sessions_timed_out'],
  634. 'sessions_error': self.stats['sessions_error'],
  635. 'success_rate': success_rate,
  636. 'total_reports_processed': self.stats['total_reports_processed'],
  637. 'max_concurrent_sessions': self.max_concurrent_sessions,
  638. }
  639. def _session_callback(
  640. self,
  641. session: ArbitrationSession,
  642. event_type: str,
  643. data: Any = None
  644. ) -> None:
  645. """Handle session events."""
  646. session_id = session.session_id
  647. if event_type == "session_completed":
  648. self._move_to_completed(session_id)
  649. self.stats['sessions_completed'] += 1
  650. elif event_type == "session_cancelled":
  651. self._move_to_completed(session_id)
  652. self.stats['sessions_cancelled'] += 1
  653. elif event_type == "session_timeout":
  654. self._move_to_completed(session_id)
  655. self.stats['sessions_timed_out'] += 1
  656. elif event_type == "session_error":
  657. self._move_to_completed(session_id)
  658. self.stats['sessions_error'] += 1
  659. def _move_to_completed(self, session_id: str) -> None:
  660. """Move session from active to completed."""
  661. with self._lock:
  662. if session_id in self.active_sessions:
  663. session = self.active_sessions.pop(session_id)
  664. self.completed_sessions[session_id] = session
  665. pprint(f"Moved session {session_id} to completed "
  666. f"({len(self.active_sessions)} active remaining)")
  667. def _cleanup_completed_sessions(self) -> None:
  668. """Clean up old completed sessions."""
  669. with self._lock:
  670. # Remove sessions older than configured limit
  671. cutoff_time = datetime.now() - timedelta(hours=1) # Keep 1 hour of history
  672. to_remove = []
  673. for session_id, session in self.completed_sessions.items():
  674. if (session.timing.completed_at and
  675. session.timing.completed_at < cutoff_time):
  676. to_remove.append(session_id)
  677. for session_id in to_remove:
  678. del self.completed_sessions[session_id]
  679. if session_id in self.session_history:
  680. self.session_history.remove(session_id)
  681. if to_remove:
  682. pprint(f"Cleaned up {len(to_remove)} old completed sessions")
  683. def cleanup(self) -> None:
  684. """Clean up session manager resources."""
  685. with self._lock:
  686. # Cancel all active sessions
  687. for session in list(self.active_sessions.values()):
  688. session.cancel("Manager shutdown")
  689. # Clean up executor
  690. if self.executor:
  691. self.executor.shutdown(wait=True)
  692. pprint("Arbitration session manager cleaned up")
  693. def __del__(self):
  694. """Destructor to ensure cleanup."""
  695. self.cleanup()
  696. def create_session_manager(
  697. config: ArbitrationConfig,
  698. algorithms: ArbitrationAlgorithms
  699. ) -> ArbitrationSessionManager:
  700. """
  701. Create arbitration session manager with given configuration.
  702. Args:
  703. config: Arbitration configuration
  704. algorithms: Arbitration algorithms processor
  705. Returns:
  706. ArbitrationSessionManager: Configured session manager
  707. """
  708. return ArbitrationSessionManager(config, algorithms)