arbitration_algorithms.py 38 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495969798991001011021031041051061071081091101111121131141151161171181191201211221231241251261271281291301311321331341351361371381391401411421431441451461471481491501511521531541551561571581591601611621631641651661671681691701711721731741751761771781791801811821831841851861871881891901911921931941951961971981992002012022032042052062072082092102112122132142152162172182192202212222232242252262272282292302312322332342352362372382392402412422432442452462472482492502512522532542552562572582592602612622632642652662672682692702712722732742752762772782792802812822832842852862872882892902912922932942952962972982993003013023033043053063073083093103113123133143153163173183193203213223233243253263273283293303313323333343353363373383393403413423433443453463473483493503513523533543553563573583593603613623633643653663673683693703713723733743753763773783793803813823833843853863873883893903913923933943953963973983994004014024034044054064074084094104114124134144154164174184194204214224234244254264274284294304314324334344354364374384394404414424434444454464474484494504514524534544554564574584594604614624634644654664674684694704714724734744754764774784794804814824834844854864874884894904914924934944954964974984995005015025035045055065075085095105115125135145155165175185195205215225235245255265275285295305315325335345355365375385395405415425435445455465475485495505515525535545555565575585595605615625635645655665675685695705715725735745755765775785795805815825835845855865875885895905915925935945955965975985996006016026036046056066076086096106116126136146156166176186196206216226236246256266276286296306316326336346356366376386396406416426436446456466476486496506516526536546556566576586596606616626636646656666676686696706716726736746756766776786796806816826836846856866876886896906916926936946956966976986997007017027037047057067077087097107117127137147157167177187197207217227237247257267277287297307317327337347357367377387397407417427437447457467477487497507517527537547557567577587597607617627637647657667677687697707717727737747757767777787797807817827837847857867877887897907917927937947957967977987998008018028038048058068078088098108118128138148158168178188198208218228238248258268278288298308318328338348358368378388398408418428438448458468478488498508518528538548558568578588598608618628638648658668678688698708718728738748758768778788798808818828838848858868878888898908918928938948958968978988999009019029039049059069079089099109119129139149159169179189199209219229239249259269279289299309319329339349359369379389399409419429439449459469479489499509519529539549559569579589599609619629639649659669679689699709719729739749759769779789799809819829839849859869879889899909919929939949959969979989991000100110021003100410051006100710081009101010111012101310141015101610171018101910201021
  1. """
  2. Arbitration Algorithms for Trixy Satellite Selection
  3. This module implements various algorithms for selecting the optimal satellite
  4. when multiple satellites detect the wakeword simultaneously. Each algorithm
  5. uses different criteria and strategies for making the selection decision.
  6. Supported Algorithms:
  7. - Volume-based: Select satellite with highest wakeword volume
  8. - Distance-based: Select satellite closest to speaker (estimated from audio)
  9. - Priority-based: Select based on room/satellite priority settings
  10. - Round-robin: Fair rotation among available satellites
  11. - Hybrid: Weighted combination of multiple algorithms
  12. - Custom: Plugin-based custom algorithms
  13. The algorithms implement the core arbitration logic as specified in CLAUDE.md,
  14. particularly the volume-based selection which is the primary algorithm.
  15. Usage:
  16. from trixy_core.arbitration import ArbitrationAlgorithms, ArbitrationAlgorithm
  17. # Create algorithms processor
  18. algorithms = ArbitrationAlgorithms(config)
  19. # Get satellite reports
  20. reports = [...] # List of WakewordReport objects
  21. # Select best satellite using volume algorithm
  22. selected = algorithms.select_satellite(
  23. reports, ArbitrationAlgorithm.VOLUME_BASED
  24. )
  25. """
  26. import asyncio
  27. import time
  28. import math
  29. import random
  30. from abc import ABC, abstractmethod
  31. from dataclasses import dataclass, field
  32. from typing import Dict, List, Optional, Any, Tuple, Union, Callable
  33. from enum import Enum
  34. import statistics
  35. from concurrent.futures import ThreadPoolExecutor
  36. from .arbitration_config import (
  37. ArbitrationConfig, ArbitrationAlgorithm, AlgorithmWeights,
  38. RoomSettings, ArbitrationConfigError
  39. )
  40. def pprint(message: str) -> None:
  41. """Arbitration algorithms logging function."""
  42. print(f"[ARBITRATION_ALGORITHMS] {message}")
  43. class SelectionCriteria(Enum):
  44. """Criteria used for satellite selection."""
  45. VOLUME = "volume"
  46. DISTANCE = "distance"
  47. PRIORITY = "priority"
  48. CONFIDENCE = "confidence"
  49. HISTORY = "history"
  50. ROUND_ROBIN = "round_robin"
  51. HYBRID = "hybrid"
  52. @dataclass
  53. class WakewordReport:
  54. """
  55. Report of wakeword detection from a satellite.
  56. This class encapsulates all information about a wakeword detection
  57. event from a specific satellite, including audio characteristics,
  58. satellite information, and timing data.
  59. """
  60. satellite_id: str
  61. wakeword_id: str
  62. volume: float
  63. confidence: float
  64. timestamp: float
  65. # Satellite information
  66. room_id: str = ""
  67. satellite_alias: str = ""
  68. mac_address: str = ""
  69. # Speaker information
  70. speaker_id: str = ""
  71. speaker_name: str = ""
  72. speaker_confidence: float = 0.0
  73. # Audio characteristics
  74. audio_buffer_length: float = 0.0
  75. sample_rate: int = 16000
  76. signal_to_noise_ratio: Optional[float] = None
  77. frequency_analysis: Dict[str, float] = field(default_factory=dict)
  78. # Distance estimation (if available)
  79. estimated_distance: Optional[float] = None
  80. distance_confidence: Optional[float] = None
  81. # Processing metadata
  82. processing_time_ms: float = 0.0
  83. model_version: str = ""
  84. def __post_init__(self):
  85. """Post-initialization processing."""
  86. # Ensure volume and confidence are within valid ranges
  87. self.volume = max(0.0, min(1.0, self.volume))
  88. self.confidence = max(0.0, min(1.0, self.confidence))
  89. self.speaker_confidence = max(0.0, min(1.0, self.speaker_confidence))
  90. def get_weighted_score(self, weights: AlgorithmWeights) -> float:
  91. """
  92. Calculate weighted score for hybrid algorithm.
  93. Args:
  94. weights: Algorithm weights configuration
  95. Returns:
  96. float: Weighted score for this report
  97. """
  98. score = 0.0
  99. # Volume component
  100. score += self.volume * weights.volume_weight
  101. # Distance component (inverted - closer is better)
  102. if self.estimated_distance is not None:
  103. distance_score = max(0.0, 1.0 - (self.estimated_distance / 10.0))
  104. score += distance_score * weights.distance_weight
  105. # Confidence component
  106. score += self.confidence * weights.confidence_weight
  107. return score
  108. def to_dict(self) -> Dict[str, Any]:
  109. """Convert report to dictionary representation."""
  110. return {
  111. 'satellite_id': self.satellite_id,
  112. 'wakeword_id': self.wakeword_id,
  113. 'volume': self.volume,
  114. 'confidence': self.confidence,
  115. 'timestamp': self.timestamp,
  116. 'room_id': self.room_id,
  117. 'satellite_alias': self.satellite_alias,
  118. 'mac_address': self.mac_address,
  119. 'speaker_id': self.speaker_id,
  120. 'speaker_name': self.speaker_name,
  121. 'speaker_confidence': self.speaker_confidence,
  122. 'audio_buffer_length': self.audio_buffer_length,
  123. 'sample_rate': self.sample_rate,
  124. 'signal_to_noise_ratio': self.signal_to_noise_ratio,
  125. 'frequency_analysis': self.frequency_analysis,
  126. 'estimated_distance': self.estimated_distance,
  127. 'distance_confidence': self.distance_confidence,
  128. 'processing_time_ms': self.processing_time_ms,
  129. 'model_version': self.model_version,
  130. }
  131. @dataclass
  132. class SelectionResult:
  133. """
  134. Result of satellite selection algorithm.
  135. Contains the selected satellite and metadata about the selection process.
  136. """
  137. selected_satellite_id: str
  138. selection_criteria: SelectionCriteria
  139. selection_score: float
  140. algorithm_used: ArbitrationAlgorithm
  141. processing_time_ms: float
  142. # Selection metadata
  143. total_candidates: int = 0
  144. tied_candidates: List[str] = field(default_factory=list)
  145. selection_reason: str = ""
  146. confidence: float = 1.0
  147. # Alternative selections (for debugging/analysis)
  148. alternative_selections: Dict[str, float] = field(default_factory=dict)
  149. def to_dict(self) -> Dict[str, Any]:
  150. """Convert result to dictionary representation."""
  151. return {
  152. 'selected_satellite_id': self.selected_satellite_id,
  153. 'selection_criteria': self.selection_criteria.value,
  154. 'selection_score': self.selection_score,
  155. 'algorithm_used': self.algorithm_used.value,
  156. 'processing_time_ms': self.processing_time_ms,
  157. 'total_candidates': self.total_candidates,
  158. 'tied_candidates': self.tied_candidates,
  159. 'selection_reason': self.selection_reason,
  160. 'confidence': self.confidence,
  161. 'alternative_selections': self.alternative_selections,
  162. }
  163. class BaseAlgorithm(ABC):
  164. """Base class for arbitration algorithms."""
  165. def __init__(self, config: ArbitrationConfig):
  166. """
  167. Initialize algorithm with configuration.
  168. Args:
  169. config: Arbitration configuration
  170. """
  171. self.config = config
  172. @abstractmethod
  173. def select_satellite(self, reports: List[WakewordReport]) -> SelectionResult:
  174. """
  175. Select the best satellite from available reports.
  176. Args:
  177. reports: List of wakeword reports from satellites
  178. Returns:
  179. SelectionResult: Selection result with chosen satellite
  180. """
  181. pass
  182. def _filter_valid_reports(self, reports: List[WakewordReport]) -> List[WakewordReport]:
  183. """Filter out invalid or disabled satellite reports."""
  184. valid_reports = []
  185. for report in reports:
  186. # Check if room is enabled
  187. if not self.config.is_room_enabled(report.room_id):
  188. continue
  189. # Check volume threshold
  190. if report.volume < self.config.volume_threshold:
  191. continue
  192. # Check confidence threshold (if applicable)
  193. if hasattr(self.config, 'confidence_threshold'):
  194. if report.confidence < getattr(self.config, 'confidence_threshold', 0.0):
  195. continue
  196. valid_reports.append(report)
  197. return valid_reports
  198. def _break_tie(self, tied_reports: List[WakewordReport]) -> WakewordReport:
  199. """
  200. Break tie between reports with equal scores.
  201. Args:
  202. tied_reports: Reports with tied scores
  203. Returns:
  204. WakewordReport: Selected report from tied candidates
  205. """
  206. if not tied_reports:
  207. raise ValueError("No tied reports provided")
  208. if len(tied_reports) == 1:
  209. return tied_reports[0]
  210. # Use secondary criteria for tie-breaking
  211. # 1. Highest confidence
  212. max_confidence = max(report.confidence for report in tied_reports)
  213. confidence_winners = [r for r in tied_reports if r.confidence == max_confidence]
  214. if len(confidence_winners) == 1:
  215. return confidence_winners[0]
  216. # 2. Most recent timestamp
  217. latest_timestamp = max(report.timestamp for report in confidence_winners)
  218. timestamp_winners = [r for r in confidence_winners if r.timestamp == latest_timestamp]
  219. if len(timestamp_winners) == 1:
  220. return timestamp_winners[0]
  221. # 3. Random selection as final fallback
  222. return random.choice(timestamp_winners)
  223. class VolumeBasedAlgorithm(BaseAlgorithm):
  224. """
  225. Volume-based satellite selection algorithm.
  226. Selects the satellite with the highest wakeword detection volume.
  227. This is the primary algorithm specified in CLAUDE.md.
  228. """
  229. def select_satellite(self, reports: List[WakewordReport]) -> SelectionResult:
  230. """Select satellite with highest volume."""
  231. start_time = time.time()
  232. valid_reports = self._filter_valid_reports(reports)
  233. if not valid_reports:
  234. raise ArbitrationConfigError("No valid satellite reports for volume-based selection")
  235. # Apply volume smoothing if configured
  236. if self.config.volume_smoothing_factor > 0:
  237. valid_reports = self._apply_volume_smoothing(valid_reports)
  238. # Find highest volume
  239. max_volume = max(report.volume for report in valid_reports)
  240. # Apply hysteresis to prevent selection bouncing
  241. hysteresis_threshold = max_volume - self.config.volume_hysteresis
  242. candidates = [r for r in valid_reports if r.volume >= hysteresis_threshold]
  243. # Select best candidate
  244. if len(candidates) == 1:
  245. selected = candidates[0]
  246. reason = f"Highest volume: {selected.volume:.3f}"
  247. else:
  248. # Multiple candidates within hysteresis range
  249. selected = self._break_tie(candidates)
  250. reason = f"Highest volume with tie-break: {selected.volume:.3f} ({len(candidates)} tied)"
  251. processing_time = (time.time() - start_time) * 1000
  252. return SelectionResult(
  253. selected_satellite_id=selected.satellite_id,
  254. selection_criteria=SelectionCriteria.VOLUME,
  255. selection_score=selected.volume,
  256. algorithm_used=ArbitrationAlgorithm.VOLUME_BASED,
  257. processing_time_ms=processing_time,
  258. total_candidates=len(valid_reports),
  259. tied_candidates=[r.satellite_id for r in candidates if r != selected],
  260. selection_reason=reason,
  261. confidence=min(1.0, selected.volume / self.config.volume_threshold),
  262. alternative_selections={
  263. r.satellite_id: r.volume for r in valid_reports if r != selected
  264. }
  265. )
  266. def _apply_volume_smoothing(self, reports: List[WakewordReport]) -> List[WakewordReport]:
  267. """Apply volume smoothing to reduce noise."""
  268. if len(reports) <= 1:
  269. return reports
  270. volumes = [r.volume for r in reports]
  271. median_volume = statistics.median(volumes)
  272. smoothed_reports = []
  273. for report in reports:
  274. smoothing_factor = self.config.volume_smoothing_factor
  275. smoothed_volume = (
  276. report.volume * (1 - smoothing_factor) +
  277. median_volume * smoothing_factor
  278. )
  279. # Create new report with smoothed volume
  280. smoothed_report = WakewordReport(
  281. satellite_id=report.satellite_id,
  282. wakeword_id=report.wakeword_id,
  283. volume=smoothed_volume,
  284. confidence=report.confidence,
  285. timestamp=report.timestamp,
  286. room_id=report.room_id,
  287. satellite_alias=report.satellite_alias,
  288. mac_address=report.mac_address,
  289. speaker_id=report.speaker_id,
  290. speaker_name=report.speaker_name,
  291. speaker_confidence=report.speaker_confidence,
  292. audio_buffer_length=report.audio_buffer_length,
  293. sample_rate=report.sample_rate,
  294. signal_to_noise_ratio=report.signal_to_noise_ratio,
  295. frequency_analysis=report.frequency_analysis,
  296. estimated_distance=report.estimated_distance,
  297. distance_confidence=report.distance_confidence,
  298. processing_time_ms=report.processing_time_ms,
  299. model_version=report.model_version,
  300. )
  301. smoothed_reports.append(smoothed_report)
  302. return smoothed_reports
  303. class DistanceBasedAlgorithm(BaseAlgorithm):
  304. """
  305. Distance-based satellite selection algorithm.
  306. Selects the satellite closest to the detected speaker based on
  307. audio analysis and distance estimation.
  308. """
  309. def select_satellite(self, reports: List[WakewordReport]) -> SelectionResult:
  310. """Select satellite with shortest estimated distance."""
  311. start_time = time.time()
  312. valid_reports = self._filter_valid_reports(reports)
  313. if not valid_reports:
  314. raise ArbitrationConfigError("No valid satellite reports for distance-based selection")
  315. # Estimate distances for reports that don't have them
  316. reports_with_distance = self._ensure_distance_estimates(valid_reports)
  317. # Find shortest distance
  318. min_distance = min(
  319. report.estimated_distance for report in reports_with_distance
  320. if report.estimated_distance is not None
  321. )
  322. # Apply distance threshold and falloff
  323. distance_threshold = min_distance + self.config.distance_threshold
  324. candidates = [
  325. r for r in reports_with_distance
  326. if r.estimated_distance is not None and r.estimated_distance <= distance_threshold
  327. ]
  328. if not candidates:
  329. # Fallback to all reports if none meet threshold
  330. candidates = reports_with_distance
  331. # Select best candidate
  332. if len(candidates) == 1:
  333. selected = candidates[0]
  334. reason = f"Shortest distance: {selected.estimated_distance:.2f}m"
  335. else:
  336. # Multiple candidates within threshold
  337. selected = min(candidates, key=lambda r: r.estimated_distance or float('inf'))
  338. reason = f"Shortest distance with tie-break: {selected.estimated_distance:.2f}m"
  339. processing_time = (time.time() - start_time) * 1000
  340. return SelectionResult(
  341. selected_satellite_id=selected.satellite_id,
  342. selection_criteria=SelectionCriteria.DISTANCE,
  343. selection_score=1.0 / (selected.estimated_distance + 0.1), # Inverted distance
  344. algorithm_used=ArbitrationAlgorithm.DISTANCE_BASED,
  345. processing_time_ms=processing_time,
  346. total_candidates=len(valid_reports),
  347. tied_candidates=[
  348. r.satellite_id for r in candidates
  349. if r != selected and abs((r.estimated_distance or 0) - (selected.estimated_distance or 0)) < 0.5
  350. ],
  351. selection_reason=reason,
  352. confidence=min(1.0, selected.distance_confidence or 0.5),
  353. alternative_selections={
  354. r.satellite_id: r.estimated_distance or float('inf')
  355. for r in valid_reports if r != selected
  356. }
  357. )
  358. def _ensure_distance_estimates(self, reports: List[WakewordReport]) -> List[WakewordReport]:
  359. """Ensure all reports have distance estimates."""
  360. processed_reports = []
  361. for report in reports:
  362. if report.estimated_distance is None and self.config.use_estimated_distance:
  363. # Estimate distance from volume and audio characteristics
  364. estimated_distance = self._estimate_distance_from_audio(report)
  365. # Create new report with estimated distance
  366. new_report = WakewordReport(
  367. satellite_id=report.satellite_id,
  368. wakeword_id=report.wakeword_id,
  369. volume=report.volume,
  370. confidence=report.confidence,
  371. timestamp=report.timestamp,
  372. room_id=report.room_id,
  373. satellite_alias=report.satellite_alias,
  374. mac_address=report.mac_address,
  375. speaker_id=report.speaker_id,
  376. speaker_name=report.speaker_name,
  377. speaker_confidence=report.speaker_confidence,
  378. audio_buffer_length=report.audio_buffer_length,
  379. sample_rate=report.sample_rate,
  380. signal_to_noise_ratio=report.signal_to_noise_ratio,
  381. frequency_analysis=report.frequency_analysis,
  382. estimated_distance=estimated_distance,
  383. distance_confidence=0.6, # Moderate confidence for estimates
  384. processing_time_ms=report.processing_time_ms,
  385. model_version=report.model_version,
  386. )
  387. processed_reports.append(new_report)
  388. else:
  389. processed_reports.append(report)
  390. return processed_reports
  391. def _estimate_distance_from_audio(self, report: WakewordReport) -> float:
  392. """
  393. Estimate distance from audio characteristics.
  394. This is a simplified distance estimation based on volume and
  395. signal characteristics. In a real implementation, this would
  396. use more sophisticated audio analysis.
  397. """
  398. # Base estimation from volume (inverse relationship)
  399. if report.volume > 0:
  400. base_distance = (1.0 / report.volume) * 2.0 # Rough scaling
  401. else:
  402. base_distance = 10.0 # Default for very low volume
  403. # Apply signal-to-noise ratio if available
  404. if report.signal_to_noise_ratio is not None:
  405. # Higher SNR suggests closer distance
  406. snr_factor = max(0.5, min(2.0, 1.0 / (report.signal_to_noise_ratio + 0.1)))
  407. base_distance *= snr_factor
  408. # Apply distance falloff factor
  409. base_distance *= self.config.distance_falloff
  410. # Clamp to reasonable range
  411. return max(0.1, min(20.0, base_distance))
  412. class PriorityBasedAlgorithm(BaseAlgorithm):
  413. """
  414. Priority-based satellite selection algorithm.
  415. Selects satellite based on configured room and satellite priorities.
  416. Higher priority rooms and satellites are preferred.
  417. """
  418. def select_satellite(self, reports: List[WakewordReport]) -> SelectionResult:
  419. """Select satellite with highest priority."""
  420. start_time = time.time()
  421. valid_reports = self._filter_valid_reports(reports)
  422. if not valid_reports:
  423. raise ArbitrationConfigError("No valid satellite reports for priority-based selection")
  424. # Calculate priority scores
  425. priority_scores = []
  426. for report in valid_reports:
  427. room_settings = self.config.get_room_settings(report.room_id)
  428. priority_score = room_settings.priority
  429. # Apply volume weighting if configured
  430. volume_weight = getattr(room_settings, 'volume_weight', 0.1)
  431. priority_score += report.volume * volume_weight
  432. priority_scores.append((report, priority_score))
  433. # Sort by priority score (highest first)
  434. priority_scores.sort(key=lambda x: x[1], reverse=True)
  435. # Find highest priority
  436. max_priority = priority_scores[0][1]
  437. top_candidates = [
  438. report for report, score in priority_scores
  439. if abs(score - max_priority) < 0.01 # Small tolerance for floating point
  440. ]
  441. # Select best candidate
  442. if len(top_candidates) == 1:
  443. selected = top_candidates[0]
  444. reason = f"Highest priority: {max_priority:.3f}"
  445. else:
  446. # Multiple candidates with same priority
  447. selected = self._break_tie(top_candidates)
  448. reason = f"Highest priority with tie-break: {max_priority:.3f} ({len(top_candidates)} tied)"
  449. processing_time = (time.time() - start_time) * 1000
  450. return SelectionResult(
  451. selected_satellite_id=selected.satellite_id,
  452. selection_criteria=SelectionCriteria.PRIORITY,
  453. selection_score=max_priority,
  454. algorithm_used=ArbitrationAlgorithm.PRIORITY_BASED,
  455. processing_time_ms=processing_time,
  456. total_candidates=len(valid_reports),
  457. tied_candidates=[r.satellite_id for r in top_candidates if r != selected],
  458. selection_reason=reason,
  459. confidence=1.0, # Priority is deterministic
  460. alternative_selections={
  461. report.satellite_id: score for report, score in priority_scores if report != selected
  462. }
  463. )
  464. class RoundRobinAlgorithm(BaseAlgorithm):
  465. """
  466. Round-robin satellite selection algorithm.
  467. Provides fair rotation among available satellites, ensuring each
  468. satellite gets an equal opportunity to handle conversations.
  469. """
  470. def __init__(self, config: ArbitrationConfig):
  471. """Initialize round-robin algorithm."""
  472. super().__init__(config)
  473. self.selection_history: List[str] = []
  474. self.selection_count = 0
  475. def select_satellite(self, reports: List[WakewordReport]) -> SelectionResult:
  476. """Select satellite using round-robin strategy."""
  477. start_time = time.time()
  478. valid_reports = self._filter_valid_reports(reports)
  479. if not valid_reports:
  480. raise ArbitrationConfigError("No valid satellite reports for round-robin selection")
  481. # Sort satellites for consistent ordering
  482. valid_reports.sort(key=lambda r: r.satellite_id)
  483. satellite_ids = [r.satellite_id for r in valid_reports]
  484. # Apply fairness mode if enabled
  485. if self.config.round_robin_fairness_mode:
  486. selected_id = self._fair_round_robin_selection(satellite_ids)
  487. else:
  488. selected_id = self._simple_round_robin_selection(satellite_ids)
  489. # Find the selected report
  490. selected = next(r for r in valid_reports if r.satellite_id == selected_id)
  491. # Update history
  492. self.selection_history.append(selected_id)
  493. self.selection_count += 1
  494. # Reset history if needed
  495. if self.selection_count >= self.config.round_robin_reset_interval:
  496. self.selection_history.clear()
  497. self.selection_count = 0
  498. pprint("Round-robin history reset")
  499. processing_time = (time.time() - start_time) * 1000
  500. reason = f"Round-robin selection (position {self.selection_count})"
  501. return SelectionResult(
  502. selected_satellite_id=selected.satellite_id,
  503. selection_criteria=SelectionCriteria.ROUND_ROBIN,
  504. selection_score=1.0, # All satellites have equal score in round-robin
  505. algorithm_used=ArbitrationAlgorithm.ROUND_ROBIN,
  506. processing_time_ms=processing_time,
  507. total_candidates=len(valid_reports),
  508. tied_candidates=[], # No ties in round-robin
  509. selection_reason=reason,
  510. confidence=1.0, # Round-robin is deterministic
  511. alternative_selections={
  512. r.satellite_id: 1.0 for r in valid_reports if r != selected
  513. }
  514. )
  515. def _simple_round_robin_selection(self, satellite_ids: List[str]) -> str:
  516. """Simple round-robin selection based on position."""
  517. if not satellite_ids:
  518. raise ValueError("No satellite IDs provided")
  519. index = self.selection_count % len(satellite_ids)
  520. return satellite_ids[index]
  521. def _fair_round_robin_selection(self, satellite_ids: List[str]) -> str:
  522. """Fair round-robin that considers selection history."""
  523. if not satellite_ids:
  524. raise ValueError("No satellite IDs provided")
  525. # Count recent selections for each satellite
  526. recent_history = self.selection_history[-len(satellite_ids) * 2:] # Look at recent history
  527. selection_counts = {sid: recent_history.count(sid) for sid in satellite_ids}
  528. # Find satellites with minimum selections
  529. min_selections = min(selection_counts.values())
  530. least_selected = [sid for sid, count in selection_counts.items() if count == min_selections]
  531. # If multiple satellites have minimum selections, rotate among them
  532. if len(least_selected) == 1:
  533. return least_selected[0]
  534. else:
  535. # Use simple round-robin among least selected
  536. index = self.selection_count % len(least_selected)
  537. return least_selected[index]
  538. def reset_history(self) -> None:
  539. """Reset round-robin selection history."""
  540. self.selection_history.clear()
  541. self.selection_count = 0
  542. pprint("Round-robin history manually reset")
  543. class HybridAlgorithm(BaseAlgorithm):
  544. """
  545. Hybrid satellite selection algorithm.
  546. Combines multiple algorithms using weighted scoring to make
  547. more sophisticated selection decisions.
  548. """
  549. def __init__(self, config: ArbitrationConfig):
  550. """Initialize hybrid algorithm with sub-algorithms."""
  551. super().__init__(config)
  552. # Initialize sub-algorithms
  553. self.volume_algorithm = VolumeBasedAlgorithm(config)
  554. self.distance_algorithm = DistanceBasedAlgorithm(config)
  555. self.priority_algorithm = PriorityBasedAlgorithm(config)
  556. def select_satellite(self, reports: List[WakewordReport]) -> SelectionResult:
  557. """Select satellite using hybrid weighted algorithm."""
  558. start_time = time.time()
  559. valid_reports = self._filter_valid_reports(reports)
  560. if not valid_reports:
  561. raise ArbitrationConfigError("No valid satellite reports for hybrid selection")
  562. # Calculate composite scores
  563. composite_scores = self._calculate_composite_scores(valid_reports)
  564. # Find highest score
  565. max_score = max(composite_scores.values())
  566. best_candidates = [
  567. satellite_id for satellite_id, score in composite_scores.items()
  568. if abs(score - max_score) < 0.001 # Small tolerance
  569. ]
  570. # Select best candidate
  571. if len(best_candidates) == 1:
  572. selected_id = best_candidates[0]
  573. reason = f"Highest hybrid score: {max_score:.3f}"
  574. else:
  575. # Break tie using volume as primary factor
  576. volume_scores = {r.satellite_id: r.volume for r in valid_reports}
  577. selected_id = max(best_candidates, key=lambda sid: volume_scores.get(sid, 0))
  578. reason = f"Highest hybrid score with volume tie-break: {max_score:.3f}"
  579. selected = next(r for r in valid_reports if r.satellite_id == selected_id)
  580. processing_time = (time.time() - start_time) * 1000
  581. return SelectionResult(
  582. selected_satellite_id=selected.satellite_id,
  583. selection_criteria=SelectionCriteria.HYBRID,
  584. selection_score=max_score,
  585. algorithm_used=ArbitrationAlgorithm.HYBRID,
  586. processing_time_ms=processing_time,
  587. total_candidates=len(valid_reports),
  588. tied_candidates=[sid for sid in best_candidates if sid != selected_id],
  589. selection_reason=reason,
  590. confidence=min(1.0, max_score),
  591. alternative_selections={
  592. satellite_id: score for satellite_id, score in composite_scores.items()
  593. if satellite_id != selected_id
  594. }
  595. )
  596. def _calculate_composite_scores(self, reports: List[WakewordReport]) -> Dict[str, float]:
  597. """Calculate composite scores for all reports."""
  598. scores = {}
  599. weights = self.config.algorithm_weights
  600. for report in reports:
  601. score = 0.0
  602. # Volume component
  603. if weights.volume_weight > 0:
  604. score += report.volume * weights.volume_weight
  605. # Distance component (inverted)
  606. if weights.distance_weight > 0 and report.estimated_distance is not None:
  607. distance_score = max(0.0, 1.0 - (report.estimated_distance / 10.0))
  608. score += distance_score * weights.distance_weight
  609. # Priority component
  610. if weights.priority_weight > 0:
  611. room_priority = self.config.get_room_priority(report.room_id)
  612. score += room_priority * weights.priority_weight
  613. # Confidence component
  614. if weights.confidence_weight > 0:
  615. score += report.confidence * weights.confidence_weight
  616. # History component (would require history tracking)
  617. if weights.history_weight > 0:
  618. # Placeholder for history-based scoring
  619. history_score = 0.5 # Neutral score
  620. score += history_score * weights.history_weight
  621. scores[report.satellite_id] = score
  622. return scores
  623. class ArbitrationAlgorithms:
  624. """
  625. Main class for managing and executing arbitration algorithms.
  626. This class provides a unified interface for all arbitration algorithms
  627. and handles algorithm selection, execution, and result processing.
  628. """
  629. def __init__(self, config: ArbitrationConfig):
  630. """
  631. Initialize arbitration algorithms processor.
  632. Args:
  633. config: Arbitration configuration
  634. """
  635. self.config = config
  636. self.algorithms = self._initialize_algorithms()
  637. self.executor = ThreadPoolExecutor(max_workers=config.thread_pool_size)
  638. # Round-robin state (shared across instances)
  639. self.round_robin_state = {}
  640. pprint(f"Initialized arbitration algorithms with {len(self.algorithms)} algorithms")
  641. def _initialize_algorithms(self) -> Dict[ArbitrationAlgorithm, BaseAlgorithm]:
  642. """Initialize all available algorithms."""
  643. algorithms = {
  644. ArbitrationAlgorithm.VOLUME_BASED: VolumeBasedAlgorithm(self.config),
  645. ArbitrationAlgorithm.DISTANCE_BASED: DistanceBasedAlgorithm(self.config),
  646. ArbitrationAlgorithm.PRIORITY_BASED: PriorityBasedAlgorithm(self.config),
  647. ArbitrationAlgorithm.ROUND_ROBIN: RoundRobinAlgorithm(self.config),
  648. ArbitrationAlgorithm.HYBRID: HybridAlgorithm(self.config),
  649. }
  650. return algorithms
  651. def select_satellite(
  652. self,
  653. reports: List[WakewordReport],
  654. algorithm: Optional[ArbitrationAlgorithm] = None
  655. ) -> SelectionResult:
  656. """
  657. Select the best satellite from wakeword reports.
  658. Args:
  659. reports: List of wakeword detection reports
  660. algorithm: Algorithm to use (defaults to config primary algorithm)
  661. Returns:
  662. SelectionResult: Selection result with chosen satellite
  663. Raises:
  664. ArbitrationConfigError: If no valid reports or algorithm fails
  665. """
  666. if not reports:
  667. raise ArbitrationConfigError("No wakeword reports provided")
  668. # Use configured algorithm if none specified
  669. if algorithm is None:
  670. algorithm = self.config.primary_algorithm
  671. pprint(f"Selecting satellite using {algorithm.value} algorithm from {len(reports)} reports")
  672. try:
  673. # Get the algorithm instance
  674. if algorithm not in self.algorithms:
  675. raise ArbitrationConfigError(f"Unsupported algorithm: {algorithm}")
  676. algorithm_instance = self.algorithms[algorithm]
  677. # Execute selection
  678. result = algorithm_instance.select_satellite(reports)
  679. pprint(f"Selected satellite {result.selected_satellite_id} with score {result.selection_score:.3f}")
  680. return result
  681. except Exception as e:
  682. # Try fallback algorithm if primary fails
  683. if algorithm != self.config.fallback_algorithm:
  684. pprint(f"Primary algorithm failed: {e}, trying fallback")
  685. return self.select_satellite(reports, self.config.fallback_algorithm)
  686. else:
  687. raise ArbitrationConfigError(f"Arbitration failed: {e}")
  688. async def select_satellite_async(
  689. self,
  690. reports: List[WakewordReport],
  691. algorithm: Optional[ArbitrationAlgorithm] = None
  692. ) -> SelectionResult:
  693. """
  694. Asynchronous satellite selection.
  695. Args:
  696. reports: List of wakeword detection reports
  697. algorithm: Algorithm to use (defaults to config primary algorithm)
  698. Returns:
  699. SelectionResult: Selection result with chosen satellite
  700. """
  701. loop = asyncio.get_event_loop()
  702. return await loop.run_in_executor(
  703. self.executor,
  704. self.select_satellite,
  705. reports,
  706. algorithm
  707. )
  708. def validate_reports(self, reports: List[WakewordReport]) -> List[str]:
  709. """
  710. Validate wakeword reports and return list of issues.
  711. Args:
  712. reports: Reports to validate
  713. Returns:
  714. List[str]: List of validation issues (empty if all valid)
  715. """
  716. issues = []
  717. if not reports:
  718. issues.append("No reports provided")
  719. return issues
  720. for i, report in enumerate(reports):
  721. if not report.satellite_id:
  722. issues.append(f"Report {i}: Missing satellite_id")
  723. if not 0 <= report.volume <= 1:
  724. issues.append(f"Report {i}: Invalid volume {report.volume}")
  725. if not 0 <= report.confidence <= 1:
  726. issues.append(f"Report {i}: Invalid confidence {report.confidence}")
  727. if report.timestamp <= 0:
  728. issues.append(f"Report {i}: Invalid timestamp {report.timestamp}")
  729. return issues
  730. def get_algorithm_info(self, algorithm: ArbitrationAlgorithm) -> Dict[str, Any]:
  731. """
  732. Get information about a specific algorithm.
  733. Args:
  734. algorithm: Algorithm to get information about
  735. Returns:
  736. Dict[str, Any]: Algorithm information
  737. """
  738. if algorithm not in self.algorithms:
  739. raise ArbitrationConfigError(f"Unknown algorithm: {algorithm}")
  740. return {
  741. 'name': algorithm.value,
  742. 'description': self.algorithms[algorithm].__class__.__doc__,
  743. 'config': self.config.get_algorithm_config(algorithm),
  744. 'supported': True,
  745. }
  746. def get_supported_algorithms(self) -> List[ArbitrationAlgorithm]:
  747. """Get list of supported algorithms."""
  748. return list(self.algorithms.keys())
  749. def benchmark_algorithms(
  750. self,
  751. test_reports: List[WakewordReport],
  752. iterations: int = 100
  753. ) -> Dict[ArbitrationAlgorithm, Dict[str, float]]:
  754. """
  755. Benchmark all algorithms with test data.
  756. Args:
  757. test_reports: Test wakeword reports
  758. iterations: Number of benchmark iterations
  759. Returns:
  760. Dict[ArbitrationAlgorithm, Dict[str, float]]: Benchmark results
  761. """
  762. results = {}
  763. for algorithm in self.algorithms.keys():
  764. if algorithm == ArbitrationAlgorithm.CUSTOM:
  765. continue # Skip custom algorithm in benchmarks
  766. times = []
  767. for _ in range(iterations):
  768. start_time = time.time()
  769. try:
  770. self.select_satellite(test_reports, algorithm)
  771. elapsed = (time.time() - start_time) * 1000
  772. times.append(elapsed)
  773. except Exception:
  774. pass # Skip failed iterations
  775. if times:
  776. results[algorithm] = {
  777. 'avg_time_ms': statistics.mean(times),
  778. 'min_time_ms': min(times),
  779. 'max_time_ms': max(times),
  780. 'median_time_ms': statistics.median(times),
  781. 'std_dev_ms': statistics.stdev(times) if len(times) > 1 else 0.0,
  782. 'success_rate': len(times) / iterations,
  783. }
  784. return results
  785. def cleanup(self) -> None:
  786. """Clean up resources."""
  787. if self.executor:
  788. self.executor.shutdown(wait=True)
  789. pprint("Arbitration algorithms cleaned up")
  790. def create_test_reports(num_reports: int = 3) -> List[WakewordReport]:
  791. """
  792. Create test wakeword reports for testing and benchmarking.
  793. Args:
  794. num_reports: Number of test reports to create
  795. Returns:
  796. List[WakewordReport]: Test reports
  797. """
  798. reports = []
  799. for i in range(num_reports):
  800. report = WakewordReport(
  801. satellite_id=f"satellite_{i+1}",
  802. wakeword_id="trixy",
  803. volume=random.uniform(0.3, 0.9),
  804. confidence=random.uniform(0.7, 0.95),
  805. timestamp=time.time() - random.uniform(0, 0.5),
  806. room_id=random.choice(["kitchen", "living_room", "bedroom", "office"]),
  807. satellite_alias=f"Satellite {i+1}",
  808. speaker_id="test_speaker",
  809. speaker_name="Test Speaker",
  810. estimated_distance=random.uniform(1.0, 5.0),
  811. distance_confidence=0.8,
  812. )
  813. reports.append(report)
  814. return reports