| 1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495969798991001011021031041051061071081091101111121131141151161171181191201211221231241251261271281291301311321331341351361371381391401411421431441451461471481491501511521531541551561571581591601611621631641651661671681691701711721731741751761771781791801811821831841851861871881891901911921931941951961971981992002012022032042052062072082092102112122132142152162172182192202212222232242252262272282292302312322332342352362372382392402412422432442452462472482492502512522532542552562572582592602612622632642652662672682692702712722732742752762772782792802812822832842852862872882892902912922932942952962972982993003013023033043053063073083093103113123133143153163173183193203213223233243253263273283293303313323333343353363373383393403413423433443453463473483493503513523533543553563573583593603613623633643653663673683693703713723733743753763773783793803813823833843853863873883893903913923933943953963973983994004014024034044054064074084094104114124134144154164174184194204214224234244254264274284294304314324334344354364374384394404414424434444454464474484494504514524534544554564574584594604614624634644654664674684694704714724734744754764774784794804814824834844854864874884894904914924934944954964974984995005015025035045055065075085095105115125135145155165175185195205215225235245255265275285295305315325335345355365375385395405415425435445455465475485495505515525535545555565575585595605615625635645655665675685695705715725735745755765775785795805815825835845855865875885895905915925935945955965975985996006016026036046056066076086096106116126136146156166176186196206216226236246256266276286296306316326336346356366376386396406416426436446456466476486496506516526536546556566576586596606616626636646656666676686696706716726736746756766776786796806816826836846856866876886896906916926936946956966976986997007017027037047057067077087097107117127137147157167177187197207217227237247257267277287297307317327337347357367377387397407417427437447457467477487497507517527537547557567577587597607617627637647657667677687697707717727737747757767777787797807817827837847857867877887897907917927937947957967977987998008018028038048058068078088098108118128138148158168178188198208218228238248258268278288298308318328338348358368378388398408418428438448458468478488498508518528538548558568578588598608618628638648658668678688698708718728738748758768778788798808818828838848858868878888898908918928938948958968978988999009019029039049059069079089099109119129139149159169179189199209219229239249259269279289299309319329339349359369379389399409419429439449459469479489499509519529539549559569579589599609619629639649659669679689699709719729739749759769779789799809819829839849859869879889899909919929939949959969979989991000100110021003100410051006100710081009101010111012101310141015101610171018101910201021 |
- """
- Arbitration Algorithms for Trixy Satellite Selection
- This module implements various algorithms for selecting the optimal satellite
- when multiple satellites detect the wakeword simultaneously. Each algorithm
- uses different criteria and strategies for making the selection decision.
- Supported Algorithms:
- - Volume-based: Select satellite with highest wakeword volume
- - Distance-based: Select satellite closest to speaker (estimated from audio)
- - Priority-based: Select based on room/satellite priority settings
- - Round-robin: Fair rotation among available satellites
- - Hybrid: Weighted combination of multiple algorithms
- - Custom: Plugin-based custom algorithms
- The algorithms implement the core arbitration logic as specified in CLAUDE.md,
- particularly the volume-based selection which is the primary algorithm.
- Usage:
- from trixy_core.arbitration import ArbitrationAlgorithms, ArbitrationAlgorithm
-
- # Create algorithms processor
- algorithms = ArbitrationAlgorithms(config)
-
- # Get satellite reports
- reports = [...] # List of WakewordReport objects
-
- # Select best satellite using volume algorithm
- selected = algorithms.select_satellite(
- reports, ArbitrationAlgorithm.VOLUME_BASED
- )
- """
- import asyncio
- import time
- import math
- import random
- from abc import ABC, abstractmethod
- from dataclasses import dataclass, field
- from typing import Dict, List, Optional, Any, Tuple, Union, Callable
- from enum import Enum
- import statistics
- from concurrent.futures import ThreadPoolExecutor
- from .arbitration_config import (
- ArbitrationConfig, ArbitrationAlgorithm, AlgorithmWeights,
- RoomSettings, ArbitrationConfigError
- )
- def pprint(message: str) -> None:
- """Arbitration algorithms logging function."""
- print(f"[ARBITRATION_ALGORITHMS] {message}")
- class SelectionCriteria(Enum):
- """Criteria used for satellite selection."""
- VOLUME = "volume"
- DISTANCE = "distance"
- PRIORITY = "priority"
- CONFIDENCE = "confidence"
- HISTORY = "history"
- ROUND_ROBIN = "round_robin"
- HYBRID = "hybrid"
- @dataclass
- class WakewordReport:
- """
- Report of wakeword detection from a satellite.
-
- This class encapsulates all information about a wakeword detection
- event from a specific satellite, including audio characteristics,
- satellite information, and timing data.
- """
- satellite_id: str
- wakeword_id: str
- volume: float
- confidence: float
- timestamp: float
-
- # Satellite information
- room_id: str = ""
- satellite_alias: str = ""
- mac_address: str = ""
-
- # Speaker information
- speaker_id: str = ""
- speaker_name: str = ""
- speaker_confidence: float = 0.0
-
- # Audio characteristics
- audio_buffer_length: float = 0.0
- sample_rate: int = 16000
- signal_to_noise_ratio: Optional[float] = None
- frequency_analysis: Dict[str, float] = field(default_factory=dict)
-
- # Distance estimation (if available)
- estimated_distance: Optional[float] = None
- distance_confidence: Optional[float] = None
-
- # Processing metadata
- processing_time_ms: float = 0.0
- model_version: str = ""
-
- def __post_init__(self):
- """Post-initialization processing."""
- # Ensure volume and confidence are within valid ranges
- self.volume = max(0.0, min(1.0, self.volume))
- self.confidence = max(0.0, min(1.0, self.confidence))
- self.speaker_confidence = max(0.0, min(1.0, self.speaker_confidence))
-
- def get_weighted_score(self, weights: AlgorithmWeights) -> float:
- """
- Calculate weighted score for hybrid algorithm.
-
- Args:
- weights: Algorithm weights configuration
-
- Returns:
- float: Weighted score for this report
- """
- score = 0.0
-
- # Volume component
- score += self.volume * weights.volume_weight
-
- # Distance component (inverted - closer is better)
- if self.estimated_distance is not None:
- distance_score = max(0.0, 1.0 - (self.estimated_distance / 10.0))
- score += distance_score * weights.distance_weight
-
- # Confidence component
- score += self.confidence * weights.confidence_weight
-
- return score
-
- def to_dict(self) -> Dict[str, Any]:
- """Convert report to dictionary representation."""
- return {
- 'satellite_id': self.satellite_id,
- 'wakeword_id': self.wakeword_id,
- 'volume': self.volume,
- 'confidence': self.confidence,
- 'timestamp': self.timestamp,
- 'room_id': self.room_id,
- 'satellite_alias': self.satellite_alias,
- 'mac_address': self.mac_address,
- 'speaker_id': self.speaker_id,
- 'speaker_name': self.speaker_name,
- 'speaker_confidence': self.speaker_confidence,
- 'audio_buffer_length': self.audio_buffer_length,
- 'sample_rate': self.sample_rate,
- 'signal_to_noise_ratio': self.signal_to_noise_ratio,
- 'frequency_analysis': self.frequency_analysis,
- 'estimated_distance': self.estimated_distance,
- 'distance_confidence': self.distance_confidence,
- 'processing_time_ms': self.processing_time_ms,
- 'model_version': self.model_version,
- }
- @dataclass
- class SelectionResult:
- """
- Result of satellite selection algorithm.
-
- Contains the selected satellite and metadata about the selection process.
- """
- selected_satellite_id: str
- selection_criteria: SelectionCriteria
- selection_score: float
- algorithm_used: ArbitrationAlgorithm
- processing_time_ms: float
-
- # Selection metadata
- total_candidates: int = 0
- tied_candidates: List[str] = field(default_factory=list)
- selection_reason: str = ""
- confidence: float = 1.0
-
- # Alternative selections (for debugging/analysis)
- alternative_selections: Dict[str, float] = field(default_factory=dict)
-
- def to_dict(self) -> Dict[str, Any]:
- """Convert result to dictionary representation."""
- return {
- 'selected_satellite_id': self.selected_satellite_id,
- 'selection_criteria': self.selection_criteria.value,
- 'selection_score': self.selection_score,
- 'algorithm_used': self.algorithm_used.value,
- 'processing_time_ms': self.processing_time_ms,
- 'total_candidates': self.total_candidates,
- 'tied_candidates': self.tied_candidates,
- 'selection_reason': self.selection_reason,
- 'confidence': self.confidence,
- 'alternative_selections': self.alternative_selections,
- }
- class BaseAlgorithm(ABC):
- """Base class for arbitration algorithms."""
-
- def __init__(self, config: ArbitrationConfig):
- """
- Initialize algorithm with configuration.
-
- Args:
- config: Arbitration configuration
- """
- self.config = config
-
- @abstractmethod
- def select_satellite(self, reports: List[WakewordReport]) -> SelectionResult:
- """
- Select the best satellite from available reports.
-
- Args:
- reports: List of wakeword reports from satellites
-
- Returns:
- SelectionResult: Selection result with chosen satellite
- """
- pass
-
- def _filter_valid_reports(self, reports: List[WakewordReport]) -> List[WakewordReport]:
- """Filter out invalid or disabled satellite reports."""
- valid_reports = []
-
- for report in reports:
- # Check if room is enabled
- if not self.config.is_room_enabled(report.room_id):
- continue
-
- # Check volume threshold
- if report.volume < self.config.volume_threshold:
- continue
-
- # Check confidence threshold (if applicable)
- if hasattr(self.config, 'confidence_threshold'):
- if report.confidence < getattr(self.config, 'confidence_threshold', 0.0):
- continue
-
- valid_reports.append(report)
-
- return valid_reports
-
- def _break_tie(self, tied_reports: List[WakewordReport]) -> WakewordReport:
- """
- Break tie between reports with equal scores.
-
- Args:
- tied_reports: Reports with tied scores
-
- Returns:
- WakewordReport: Selected report from tied candidates
- """
- if not tied_reports:
- raise ValueError("No tied reports provided")
-
- if len(tied_reports) == 1:
- return tied_reports[0]
-
- # Use secondary criteria for tie-breaking
- # 1. Highest confidence
- max_confidence = max(report.confidence for report in tied_reports)
- confidence_winners = [r for r in tied_reports if r.confidence == max_confidence]
-
- if len(confidence_winners) == 1:
- return confidence_winners[0]
-
- # 2. Most recent timestamp
- latest_timestamp = max(report.timestamp for report in confidence_winners)
- timestamp_winners = [r for r in confidence_winners if r.timestamp == latest_timestamp]
-
- if len(timestamp_winners) == 1:
- return timestamp_winners[0]
-
- # 3. Random selection as final fallback
- return random.choice(timestamp_winners)
- class VolumeBasedAlgorithm(BaseAlgorithm):
- """
- Volume-based satellite selection algorithm.
-
- Selects the satellite with the highest wakeword detection volume.
- This is the primary algorithm specified in CLAUDE.md.
- """
-
- def select_satellite(self, reports: List[WakewordReport]) -> SelectionResult:
- """Select satellite with highest volume."""
- start_time = time.time()
-
- valid_reports = self._filter_valid_reports(reports)
-
- if not valid_reports:
- raise ArbitrationConfigError("No valid satellite reports for volume-based selection")
-
- # Apply volume smoothing if configured
- if self.config.volume_smoothing_factor > 0:
- valid_reports = self._apply_volume_smoothing(valid_reports)
-
- # Find highest volume
- max_volume = max(report.volume for report in valid_reports)
-
- # Apply hysteresis to prevent selection bouncing
- hysteresis_threshold = max_volume - self.config.volume_hysteresis
- candidates = [r for r in valid_reports if r.volume >= hysteresis_threshold]
-
- # Select best candidate
- if len(candidates) == 1:
- selected = candidates[0]
- reason = f"Highest volume: {selected.volume:.3f}"
- else:
- # Multiple candidates within hysteresis range
- selected = self._break_tie(candidates)
- reason = f"Highest volume with tie-break: {selected.volume:.3f} ({len(candidates)} tied)"
-
- processing_time = (time.time() - start_time) * 1000
-
- return SelectionResult(
- selected_satellite_id=selected.satellite_id,
- selection_criteria=SelectionCriteria.VOLUME,
- selection_score=selected.volume,
- algorithm_used=ArbitrationAlgorithm.VOLUME_BASED,
- processing_time_ms=processing_time,
- total_candidates=len(valid_reports),
- tied_candidates=[r.satellite_id for r in candidates if r != selected],
- selection_reason=reason,
- confidence=min(1.0, selected.volume / self.config.volume_threshold),
- alternative_selections={
- r.satellite_id: r.volume for r in valid_reports if r != selected
- }
- )
-
- def _apply_volume_smoothing(self, reports: List[WakewordReport]) -> List[WakewordReport]:
- """Apply volume smoothing to reduce noise."""
- if len(reports) <= 1:
- return reports
-
- volumes = [r.volume for r in reports]
- median_volume = statistics.median(volumes)
-
- smoothed_reports = []
- for report in reports:
- smoothing_factor = self.config.volume_smoothing_factor
- smoothed_volume = (
- report.volume * (1 - smoothing_factor) +
- median_volume * smoothing_factor
- )
-
- # Create new report with smoothed volume
- smoothed_report = WakewordReport(
- satellite_id=report.satellite_id,
- wakeword_id=report.wakeword_id,
- volume=smoothed_volume,
- confidence=report.confidence,
- timestamp=report.timestamp,
- room_id=report.room_id,
- satellite_alias=report.satellite_alias,
- mac_address=report.mac_address,
- speaker_id=report.speaker_id,
- speaker_name=report.speaker_name,
- speaker_confidence=report.speaker_confidence,
- audio_buffer_length=report.audio_buffer_length,
- sample_rate=report.sample_rate,
- signal_to_noise_ratio=report.signal_to_noise_ratio,
- frequency_analysis=report.frequency_analysis,
- estimated_distance=report.estimated_distance,
- distance_confidence=report.distance_confidence,
- processing_time_ms=report.processing_time_ms,
- model_version=report.model_version,
- )
- smoothed_reports.append(smoothed_report)
-
- return smoothed_reports
- class DistanceBasedAlgorithm(BaseAlgorithm):
- """
- Distance-based satellite selection algorithm.
-
- Selects the satellite closest to the detected speaker based on
- audio analysis and distance estimation.
- """
-
- def select_satellite(self, reports: List[WakewordReport]) -> SelectionResult:
- """Select satellite with shortest estimated distance."""
- start_time = time.time()
-
- valid_reports = self._filter_valid_reports(reports)
-
- if not valid_reports:
- raise ArbitrationConfigError("No valid satellite reports for distance-based selection")
-
- # Estimate distances for reports that don't have them
- reports_with_distance = self._ensure_distance_estimates(valid_reports)
-
- # Find shortest distance
- min_distance = min(
- report.estimated_distance for report in reports_with_distance
- if report.estimated_distance is not None
- )
-
- # Apply distance threshold and falloff
- distance_threshold = min_distance + self.config.distance_threshold
- candidates = [
- r for r in reports_with_distance
- if r.estimated_distance is not None and r.estimated_distance <= distance_threshold
- ]
-
- if not candidates:
- # Fallback to all reports if none meet threshold
- candidates = reports_with_distance
-
- # Select best candidate
- if len(candidates) == 1:
- selected = candidates[0]
- reason = f"Shortest distance: {selected.estimated_distance:.2f}m"
- else:
- # Multiple candidates within threshold
- selected = min(candidates, key=lambda r: r.estimated_distance or float('inf'))
- reason = f"Shortest distance with tie-break: {selected.estimated_distance:.2f}m"
-
- processing_time = (time.time() - start_time) * 1000
-
- return SelectionResult(
- selected_satellite_id=selected.satellite_id,
- selection_criteria=SelectionCriteria.DISTANCE,
- selection_score=1.0 / (selected.estimated_distance + 0.1), # Inverted distance
- algorithm_used=ArbitrationAlgorithm.DISTANCE_BASED,
- processing_time_ms=processing_time,
- total_candidates=len(valid_reports),
- tied_candidates=[
- r.satellite_id for r in candidates
- if r != selected and abs((r.estimated_distance or 0) - (selected.estimated_distance or 0)) < 0.5
- ],
- selection_reason=reason,
- confidence=min(1.0, selected.distance_confidence or 0.5),
- alternative_selections={
- r.satellite_id: r.estimated_distance or float('inf')
- for r in valid_reports if r != selected
- }
- )
-
- def _ensure_distance_estimates(self, reports: List[WakewordReport]) -> List[WakewordReport]:
- """Ensure all reports have distance estimates."""
- processed_reports = []
-
- for report in reports:
- if report.estimated_distance is None and self.config.use_estimated_distance:
- # Estimate distance from volume and audio characteristics
- estimated_distance = self._estimate_distance_from_audio(report)
-
- # Create new report with estimated distance
- new_report = WakewordReport(
- satellite_id=report.satellite_id,
- wakeword_id=report.wakeword_id,
- volume=report.volume,
- confidence=report.confidence,
- timestamp=report.timestamp,
- room_id=report.room_id,
- satellite_alias=report.satellite_alias,
- mac_address=report.mac_address,
- speaker_id=report.speaker_id,
- speaker_name=report.speaker_name,
- speaker_confidence=report.speaker_confidence,
- audio_buffer_length=report.audio_buffer_length,
- sample_rate=report.sample_rate,
- signal_to_noise_ratio=report.signal_to_noise_ratio,
- frequency_analysis=report.frequency_analysis,
- estimated_distance=estimated_distance,
- distance_confidence=0.6, # Moderate confidence for estimates
- processing_time_ms=report.processing_time_ms,
- model_version=report.model_version,
- )
- processed_reports.append(new_report)
- else:
- processed_reports.append(report)
-
- return processed_reports
-
- def _estimate_distance_from_audio(self, report: WakewordReport) -> float:
- """
- Estimate distance from audio characteristics.
-
- This is a simplified distance estimation based on volume and
- signal characteristics. In a real implementation, this would
- use more sophisticated audio analysis.
- """
- # Base estimation from volume (inverse relationship)
- if report.volume > 0:
- base_distance = (1.0 / report.volume) * 2.0 # Rough scaling
- else:
- base_distance = 10.0 # Default for very low volume
-
- # Apply signal-to-noise ratio if available
- if report.signal_to_noise_ratio is not None:
- # Higher SNR suggests closer distance
- snr_factor = max(0.5, min(2.0, 1.0 / (report.signal_to_noise_ratio + 0.1)))
- base_distance *= snr_factor
-
- # Apply distance falloff factor
- base_distance *= self.config.distance_falloff
-
- # Clamp to reasonable range
- return max(0.1, min(20.0, base_distance))
- class PriorityBasedAlgorithm(BaseAlgorithm):
- """
- Priority-based satellite selection algorithm.
-
- Selects satellite based on configured room and satellite priorities.
- Higher priority rooms and satellites are preferred.
- """
-
- def select_satellite(self, reports: List[WakewordReport]) -> SelectionResult:
- """Select satellite with highest priority."""
- start_time = time.time()
-
- valid_reports = self._filter_valid_reports(reports)
-
- if not valid_reports:
- raise ArbitrationConfigError("No valid satellite reports for priority-based selection")
-
- # Calculate priority scores
- priority_scores = []
- for report in valid_reports:
- room_settings = self.config.get_room_settings(report.room_id)
- priority_score = room_settings.priority
-
- # Apply volume weighting if configured
- volume_weight = getattr(room_settings, 'volume_weight', 0.1)
- priority_score += report.volume * volume_weight
-
- priority_scores.append((report, priority_score))
-
- # Sort by priority score (highest first)
- priority_scores.sort(key=lambda x: x[1], reverse=True)
-
- # Find highest priority
- max_priority = priority_scores[0][1]
- top_candidates = [
- report for report, score in priority_scores
- if abs(score - max_priority) < 0.01 # Small tolerance for floating point
- ]
-
- # Select best candidate
- if len(top_candidates) == 1:
- selected = top_candidates[0]
- reason = f"Highest priority: {max_priority:.3f}"
- else:
- # Multiple candidates with same priority
- selected = self._break_tie(top_candidates)
- reason = f"Highest priority with tie-break: {max_priority:.3f} ({len(top_candidates)} tied)"
-
- processing_time = (time.time() - start_time) * 1000
-
- return SelectionResult(
- selected_satellite_id=selected.satellite_id,
- selection_criteria=SelectionCriteria.PRIORITY,
- selection_score=max_priority,
- algorithm_used=ArbitrationAlgorithm.PRIORITY_BASED,
- processing_time_ms=processing_time,
- total_candidates=len(valid_reports),
- tied_candidates=[r.satellite_id for r in top_candidates if r != selected],
- selection_reason=reason,
- confidence=1.0, # Priority is deterministic
- alternative_selections={
- report.satellite_id: score for report, score in priority_scores if report != selected
- }
- )
- class RoundRobinAlgorithm(BaseAlgorithm):
- """
- Round-robin satellite selection algorithm.
-
- Provides fair rotation among available satellites, ensuring each
- satellite gets an equal opportunity to handle conversations.
- """
-
- def __init__(self, config: ArbitrationConfig):
- """Initialize round-robin algorithm."""
- super().__init__(config)
- self.selection_history: List[str] = []
- self.selection_count = 0
-
- def select_satellite(self, reports: List[WakewordReport]) -> SelectionResult:
- """Select satellite using round-robin strategy."""
- start_time = time.time()
-
- valid_reports = self._filter_valid_reports(reports)
-
- if not valid_reports:
- raise ArbitrationConfigError("No valid satellite reports for round-robin selection")
-
- # Sort satellites for consistent ordering
- valid_reports.sort(key=lambda r: r.satellite_id)
- satellite_ids = [r.satellite_id for r in valid_reports]
-
- # Apply fairness mode if enabled
- if self.config.round_robin_fairness_mode:
- selected_id = self._fair_round_robin_selection(satellite_ids)
- else:
- selected_id = self._simple_round_robin_selection(satellite_ids)
-
- # Find the selected report
- selected = next(r for r in valid_reports if r.satellite_id == selected_id)
-
- # Update history
- self.selection_history.append(selected_id)
- self.selection_count += 1
-
- # Reset history if needed
- if self.selection_count >= self.config.round_robin_reset_interval:
- self.selection_history.clear()
- self.selection_count = 0
- pprint("Round-robin history reset")
-
- processing_time = (time.time() - start_time) * 1000
-
- reason = f"Round-robin selection (position {self.selection_count})"
-
- return SelectionResult(
- selected_satellite_id=selected.satellite_id,
- selection_criteria=SelectionCriteria.ROUND_ROBIN,
- selection_score=1.0, # All satellites have equal score in round-robin
- algorithm_used=ArbitrationAlgorithm.ROUND_ROBIN,
- processing_time_ms=processing_time,
- total_candidates=len(valid_reports),
- tied_candidates=[], # No ties in round-robin
- selection_reason=reason,
- confidence=1.0, # Round-robin is deterministic
- alternative_selections={
- r.satellite_id: 1.0 for r in valid_reports if r != selected
- }
- )
-
- def _simple_round_robin_selection(self, satellite_ids: List[str]) -> str:
- """Simple round-robin selection based on position."""
- if not satellite_ids:
- raise ValueError("No satellite IDs provided")
-
- index = self.selection_count % len(satellite_ids)
- return satellite_ids[index]
-
- def _fair_round_robin_selection(self, satellite_ids: List[str]) -> str:
- """Fair round-robin that considers selection history."""
- if not satellite_ids:
- raise ValueError("No satellite IDs provided")
-
- # Count recent selections for each satellite
- recent_history = self.selection_history[-len(satellite_ids) * 2:] # Look at recent history
- selection_counts = {sid: recent_history.count(sid) for sid in satellite_ids}
-
- # Find satellites with minimum selections
- min_selections = min(selection_counts.values())
- least_selected = [sid for sid, count in selection_counts.items() if count == min_selections]
-
- # If multiple satellites have minimum selections, rotate among them
- if len(least_selected) == 1:
- return least_selected[0]
- else:
- # Use simple round-robin among least selected
- index = self.selection_count % len(least_selected)
- return least_selected[index]
-
- def reset_history(self) -> None:
- """Reset round-robin selection history."""
- self.selection_history.clear()
- self.selection_count = 0
- pprint("Round-robin history manually reset")
- class HybridAlgorithm(BaseAlgorithm):
- """
- Hybrid satellite selection algorithm.
-
- Combines multiple algorithms using weighted scoring to make
- more sophisticated selection decisions.
- """
-
- def __init__(self, config: ArbitrationConfig):
- """Initialize hybrid algorithm with sub-algorithms."""
- super().__init__(config)
-
- # Initialize sub-algorithms
- self.volume_algorithm = VolumeBasedAlgorithm(config)
- self.distance_algorithm = DistanceBasedAlgorithm(config)
- self.priority_algorithm = PriorityBasedAlgorithm(config)
-
- def select_satellite(self, reports: List[WakewordReport]) -> SelectionResult:
- """Select satellite using hybrid weighted algorithm."""
- start_time = time.time()
-
- valid_reports = self._filter_valid_reports(reports)
-
- if not valid_reports:
- raise ArbitrationConfigError("No valid satellite reports for hybrid selection")
-
- # Calculate composite scores
- composite_scores = self._calculate_composite_scores(valid_reports)
-
- # Find highest score
- max_score = max(composite_scores.values())
- best_candidates = [
- satellite_id for satellite_id, score in composite_scores.items()
- if abs(score - max_score) < 0.001 # Small tolerance
- ]
-
- # Select best candidate
- if len(best_candidates) == 1:
- selected_id = best_candidates[0]
- reason = f"Highest hybrid score: {max_score:.3f}"
- else:
- # Break tie using volume as primary factor
- volume_scores = {r.satellite_id: r.volume for r in valid_reports}
- selected_id = max(best_candidates, key=lambda sid: volume_scores.get(sid, 0))
- reason = f"Highest hybrid score with volume tie-break: {max_score:.3f}"
-
- selected = next(r for r in valid_reports if r.satellite_id == selected_id)
-
- processing_time = (time.time() - start_time) * 1000
-
- return SelectionResult(
- selected_satellite_id=selected.satellite_id,
- selection_criteria=SelectionCriteria.HYBRID,
- selection_score=max_score,
- algorithm_used=ArbitrationAlgorithm.HYBRID,
- processing_time_ms=processing_time,
- total_candidates=len(valid_reports),
- tied_candidates=[sid for sid in best_candidates if sid != selected_id],
- selection_reason=reason,
- confidence=min(1.0, max_score),
- alternative_selections={
- satellite_id: score for satellite_id, score in composite_scores.items()
- if satellite_id != selected_id
- }
- )
-
- def _calculate_composite_scores(self, reports: List[WakewordReport]) -> Dict[str, float]:
- """Calculate composite scores for all reports."""
- scores = {}
- weights = self.config.algorithm_weights
-
- for report in reports:
- score = 0.0
-
- # Volume component
- if weights.volume_weight > 0:
- score += report.volume * weights.volume_weight
-
- # Distance component (inverted)
- if weights.distance_weight > 0 and report.estimated_distance is not None:
- distance_score = max(0.0, 1.0 - (report.estimated_distance / 10.0))
- score += distance_score * weights.distance_weight
-
- # Priority component
- if weights.priority_weight > 0:
- room_priority = self.config.get_room_priority(report.room_id)
- score += room_priority * weights.priority_weight
-
- # Confidence component
- if weights.confidence_weight > 0:
- score += report.confidence * weights.confidence_weight
-
- # History component (would require history tracking)
- if weights.history_weight > 0:
- # Placeholder for history-based scoring
- history_score = 0.5 # Neutral score
- score += history_score * weights.history_weight
-
- scores[report.satellite_id] = score
-
- return scores
- class ArbitrationAlgorithms:
- """
- Main class for managing and executing arbitration algorithms.
-
- This class provides a unified interface for all arbitration algorithms
- and handles algorithm selection, execution, and result processing.
- """
-
- def __init__(self, config: ArbitrationConfig):
- """
- Initialize arbitration algorithms processor.
-
- Args:
- config: Arbitration configuration
- """
- self.config = config
- self.algorithms = self._initialize_algorithms()
- self.executor = ThreadPoolExecutor(max_workers=config.thread_pool_size)
-
- # Round-robin state (shared across instances)
- self.round_robin_state = {}
-
- pprint(f"Initialized arbitration algorithms with {len(self.algorithms)} algorithms")
-
- def _initialize_algorithms(self) -> Dict[ArbitrationAlgorithm, BaseAlgorithm]:
- """Initialize all available algorithms."""
- algorithms = {
- ArbitrationAlgorithm.VOLUME_BASED: VolumeBasedAlgorithm(self.config),
- ArbitrationAlgorithm.DISTANCE_BASED: DistanceBasedAlgorithm(self.config),
- ArbitrationAlgorithm.PRIORITY_BASED: PriorityBasedAlgorithm(self.config),
- ArbitrationAlgorithm.ROUND_ROBIN: RoundRobinAlgorithm(self.config),
- ArbitrationAlgorithm.HYBRID: HybridAlgorithm(self.config),
- }
-
- return algorithms
-
- def select_satellite(
- self,
- reports: List[WakewordReport],
- algorithm: Optional[ArbitrationAlgorithm] = None
- ) -> SelectionResult:
- """
- Select the best satellite from wakeword reports.
-
- Args:
- reports: List of wakeword detection reports
- algorithm: Algorithm to use (defaults to config primary algorithm)
-
- Returns:
- SelectionResult: Selection result with chosen satellite
-
- Raises:
- ArbitrationConfigError: If no valid reports or algorithm fails
- """
- if not reports:
- raise ArbitrationConfigError("No wakeword reports provided")
-
- # Use configured algorithm if none specified
- if algorithm is None:
- algorithm = self.config.primary_algorithm
-
- pprint(f"Selecting satellite using {algorithm.value} algorithm from {len(reports)} reports")
-
- try:
- # Get the algorithm instance
- if algorithm not in self.algorithms:
- raise ArbitrationConfigError(f"Unsupported algorithm: {algorithm}")
-
- algorithm_instance = self.algorithms[algorithm]
-
- # Execute selection
- result = algorithm_instance.select_satellite(reports)
-
- pprint(f"Selected satellite {result.selected_satellite_id} with score {result.selection_score:.3f}")
- return result
-
- except Exception as e:
- # Try fallback algorithm if primary fails
- if algorithm != self.config.fallback_algorithm:
- pprint(f"Primary algorithm failed: {e}, trying fallback")
- return self.select_satellite(reports, self.config.fallback_algorithm)
- else:
- raise ArbitrationConfigError(f"Arbitration failed: {e}")
-
- async def select_satellite_async(
- self,
- reports: List[WakewordReport],
- algorithm: Optional[ArbitrationAlgorithm] = None
- ) -> SelectionResult:
- """
- Asynchronous satellite selection.
-
- Args:
- reports: List of wakeword detection reports
- algorithm: Algorithm to use (defaults to config primary algorithm)
-
- Returns:
- SelectionResult: Selection result with chosen satellite
- """
- loop = asyncio.get_event_loop()
- return await loop.run_in_executor(
- self.executor,
- self.select_satellite,
- reports,
- algorithm
- )
-
- def validate_reports(self, reports: List[WakewordReport]) -> List[str]:
- """
- Validate wakeword reports and return list of issues.
-
- Args:
- reports: Reports to validate
-
- Returns:
- List[str]: List of validation issues (empty if all valid)
- """
- issues = []
-
- if not reports:
- issues.append("No reports provided")
- return issues
-
- for i, report in enumerate(reports):
- if not report.satellite_id:
- issues.append(f"Report {i}: Missing satellite_id")
-
- if not 0 <= report.volume <= 1:
- issues.append(f"Report {i}: Invalid volume {report.volume}")
-
- if not 0 <= report.confidence <= 1:
- issues.append(f"Report {i}: Invalid confidence {report.confidence}")
-
- if report.timestamp <= 0:
- issues.append(f"Report {i}: Invalid timestamp {report.timestamp}")
-
- return issues
-
- def get_algorithm_info(self, algorithm: ArbitrationAlgorithm) -> Dict[str, Any]:
- """
- Get information about a specific algorithm.
-
- Args:
- algorithm: Algorithm to get information about
-
- Returns:
- Dict[str, Any]: Algorithm information
- """
- if algorithm not in self.algorithms:
- raise ArbitrationConfigError(f"Unknown algorithm: {algorithm}")
-
- return {
- 'name': algorithm.value,
- 'description': self.algorithms[algorithm].__class__.__doc__,
- 'config': self.config.get_algorithm_config(algorithm),
- 'supported': True,
- }
-
- def get_supported_algorithms(self) -> List[ArbitrationAlgorithm]:
- """Get list of supported algorithms."""
- return list(self.algorithms.keys())
-
- def benchmark_algorithms(
- self,
- test_reports: List[WakewordReport],
- iterations: int = 100
- ) -> Dict[ArbitrationAlgorithm, Dict[str, float]]:
- """
- Benchmark all algorithms with test data.
-
- Args:
- test_reports: Test wakeword reports
- iterations: Number of benchmark iterations
-
- Returns:
- Dict[ArbitrationAlgorithm, Dict[str, float]]: Benchmark results
- """
- results = {}
-
- for algorithm in self.algorithms.keys():
- if algorithm == ArbitrationAlgorithm.CUSTOM:
- continue # Skip custom algorithm in benchmarks
-
- times = []
- for _ in range(iterations):
- start_time = time.time()
- try:
- self.select_satellite(test_reports, algorithm)
- elapsed = (time.time() - start_time) * 1000
- times.append(elapsed)
- except Exception:
- pass # Skip failed iterations
-
- if times:
- results[algorithm] = {
- 'avg_time_ms': statistics.mean(times),
- 'min_time_ms': min(times),
- 'max_time_ms': max(times),
- 'median_time_ms': statistics.median(times),
- 'std_dev_ms': statistics.stdev(times) if len(times) > 1 else 0.0,
- 'success_rate': len(times) / iterations,
- }
-
- return results
-
- def cleanup(self) -> None:
- """Clean up resources."""
- if self.executor:
- self.executor.shutdown(wait=True)
- pprint("Arbitration algorithms cleaned up")
- def create_test_reports(num_reports: int = 3) -> List[WakewordReport]:
- """
- Create test wakeword reports for testing and benchmarking.
-
- Args:
- num_reports: Number of test reports to create
-
- Returns:
- List[WakewordReport]: Test reports
- """
- reports = []
-
- for i in range(num_reports):
- report = WakewordReport(
- satellite_id=f"satellite_{i+1}",
- wakeword_id="trixy",
- volume=random.uniform(0.3, 0.9),
- confidence=random.uniform(0.7, 0.95),
- timestamp=time.time() - random.uniform(0, 0.5),
- room_id=random.choice(["kitchen", "living_room", "bedroom", "office"]),
- satellite_alias=f"Satellite {i+1}",
- speaker_id="test_speaker",
- speaker_name="Test Speaker",
- estimated_distance=random.uniform(1.0, 5.0),
- distance_confidence=0.8,
- )
- reports.append(report)
-
- return reports
|