Download validation/app/websocket/handler.py from indrajit4533/voiceguard: direct link, hf CLI and curl.
- Browser
- Download file 21.2 kB
-
https://huggingface.co/indrajit4533/voiceguard/resolve/main/validation/app/websocket/handler.py
- Command line
-
hf download hf://indrajit4533/voiceguard/validation/app/websocket/handler.py
-
curl -L -o handler.py https://huggingface.co/indrajit4533/voiceguard/resolve/main/validation/app/websocket/handler.py
21.2 kB
| import asyncio | |
| import logging | |
| import json | |
| import time | |
| import base64 | |
| import numpy as np | |
| import torch | |
| from typing import Dict, Optional | |
| from fastapi import WebSocket, WebSocketDisconnect | |
| from ..config import settings | |
| from .protocol import ( | |
| MessageType, | |
| ClientMessage, | |
| ClientAudioChunk, | |
| ClientConfig, | |
| ClientPing, | |
| ClientReset, | |
| AuditLogEntry, | |
| ) | |
| from .manager import ConnectionManager | |
| from ..model.inference import InferenceEngine | |
| from ..utils.timing import timestamp_ms | |
| class WebSocketHandler: | |
| """ | |
| Handles WebSocket message routing with per-connection inference engines. | |
| """ | |
| def __init__( | |
| self, | |
| manager: ConnectionManager, | |
| engine_factory, | |
| max_chunk_size_bytes: int = 64 * 1024, | |
| storage=None, | |
| alerts=None, | |
| verification=None, | |
| ): | |
| self.manager = manager | |
| self._engine_factory = engine_factory | |
| self._engines: Dict[str, InferenceEngine] = {} | |
| self._max_chunk_size_bytes = max_chunk_size_bytes | |
| self.storage = storage | |
| self.alerts = alerts | |
| self.verification = verification | |
| self._alert_cb_set = False | |
| self._loop = None | |
| self._logger = logging.getLogger(__name__) | |
| async def handle_connection(self, websocket: WebSocket, connection_id: str) -> None: | |
| """Main entry point for a WebSocket connection.""" | |
| try: | |
| await self.manager.connect(websocket, connection_id) | |
| engine = self._engine_factory(connection_id) | |
| self._engines[connection_id] = engine | |
| self._loop = asyncio.get_running_loop() | |
| if self.alerts is not None and not self._alert_cb_set: | |
| self.alerts.set_push_callback(self._push_alert) | |
| self._alert_cb_set = True | |
| # Send ready message with model info | |
| await self._send_ready(connection_id, engine) | |
| # Main message loop | |
| async for raw_message in websocket.iter_text(): | |
| connection_id = websocket.headers.get("X-Connection-ID") or connection_id | |
| await self._handle_text(connection_id, engine, raw_message) | |
| except WebSocketDisconnect: | |
| self._logger.info(f"Client disconnected: {connection_id}") | |
| except Exception as e: | |
| self._logger.error(f"Unhandled error in connection {connection_id}: {e}", exc_info=True) | |
| try: | |
| await self.manager.send_json(connection_id, { | |
| "type": MessageType.ERROR.value, | |
| "code": "INTERNAL_ERROR", | |
| "message": str(e), | |
| }) | |
| except Exception: | |
| pass | |
| finally: | |
| await self._cleanup(connection_id) | |
| async def _handle_text(self, connection_id: str, engine: InferenceEngine, raw: str) -> None: | |
| """Route a single text message based on its type.""" | |
| try: | |
| data = json.loads(raw) | |
| except json.JSONDecodeError: | |
| await self._send_error(connection_id, "INVALID_JSON", "Malformed JSON payload") | |
| return | |
| msg_type = data.get("type", "") | |
| try: | |
| if msg_type == MessageType.AUDIO_CHUNK.value: | |
| await self._handle_audio_chunk(connection_id, engine, data) | |
| elif msg_type == MessageType.CONFIG.value: | |
| await self._handle_config(connection_id, engine, data) | |
| elif msg_type == MessageType.PING.value: | |
| await self._handle_ping(connection_id, data) | |
| elif msg_type == MessageType.RESET.value: | |
| await self._handle_reset(connection_id, engine) | |
| else: | |
| await self._send_error(connection_id, "UNKNOWN_MESSAGE_TYPE", f"Unknown type: {msg_type}") | |
| except Exception as e: | |
| self._logger.error(f"Error handling {msg_type} from {connection_id}: {e}") | |
| await self._send_error(connection_id, "HANDLER_ERROR", str(e)) | |
| async def _handle_audio_chunk( | |
| self, connection_id: str, engine: InferenceEngine, data: dict | |
| ) -> None: | |
| """Process an audio chunk through the inference engine.""" | |
| t_decode0 = time.perf_counter() | |
| samples = data.get("data") | |
| if samples is None or (isinstance(samples, (list, dict, str)) and len(samples) == 0): | |
| await self._send_error(connection_id, "INVALID_AUDIO_DATA", "Missing or empty audio data") | |
| return | |
| fmt = data.get("format", "int16") | |
| try: | |
| if isinstance(samples, str): | |
| raw_bytes = base64.b64decode(samples) | |
| dtype = np.float32 if fmt == "float32" else np.int16 | |
| audio = np.frombuffer(raw_bytes, dtype=dtype) | |
| elif isinstance(samples, list): | |
| dtype = np.float32 if fmt == "float32" else np.int16 | |
| audio = np.array(samples, dtype=dtype) | |
| elif isinstance(samples, dict): | |
| dtype = np.float32 if fmt == "float32" else np.int16 | |
| audio = np.array(list(samples.values()), dtype=dtype) | |
| else: | |
| await self._send_error(connection_id, "INVALID_AUDIO_DATA", "Unsupported audio data format") | |
| return | |
| except Exception as e: | |
| await self._send_error(connection_id, "INVALID_AUDIO_DATA", f"Error parsing audio: {e}") | |
| return | |
| if len(audio) * 2 > self._max_chunk_size_bytes: | |
| await self._send_error(connection_id, "CHUNK_TOO_LARGE", f"Chunk exceeds {self._max_chunk_size_bytes} bytes") | |
| return | |
| decode_ms = (time.perf_counter() - t_decode0) * 1000.0 | |
| sample_rate = data.get("sample_rate", 16000) | |
| timestamp = data.get("timestamp_ms") | |
| chunk_index = data.get("chunk_index") | |
| fmt = data.get("format", "int16") | |
| self.manager.touch(connection_id, n_chunks=1, n_bytes=len(audio.tobytes())) | |
| # Update chunk index tracking if provided by client | |
| if chunk_index is not None: | |
| setattr(engine, "_last_client_chunk_index", chunk_index) | |
| # Process through inference engine (streaming pipeline) | |
| result = engine.process_stream_chunk(audio, sample_rate) | |
| if result is not None: | |
| # Build detection result message (includes multi-detector fields) | |
| if result.probabilities is not None: | |
| probs_dict = { | |
| settings.class_labels[i]: round(float(p), 4) | |
| for i, p in enumerate(result.probabilities) | |
| } | |
| else: | |
| probs_dict = None | |
| latency_stages = dict(result.latency_stages or {}) | |
| latency_stages["decode_ms"] = round(decode_ms, 2) | |
| detection_msg = { | |
| "type": MessageType.DETECTION_RESULT.value, | |
| "chunk_index": result.chunk_index, | |
| "timestamp_ms": result.timestamp_ms, | |
| "window_id": result.window_id, | |
| "status": result.status, | |
| "probabilities": probs_dict, | |
| "predicted_class": result.predicted_class, | |
| "predicted_label": result.predicted_label, | |
| "confidence": round(float(result.confidence), 4), | |
| "is_bonafide": bool(result.is_bonafide), | |
| "threat_score": (None if result.threat_score is None | |
| else round(float(result.threat_score), 2)), | |
| "authenticity_score": (None if result.authenticity_score is None | |
| else round(float(result.authenticity_score), 2)), | |
| "inference_ms": round(float(result.inference_ms), 2), | |
| "ema_synthetic_probability": (None if result.ema_synthetic_probability is None | |
| else round(float(result.ema_synthetic_probability), 4)), | |
| "tensor_stats": result.tensor_stats, | |
| "latency_stages": latency_stages, | |
| "vad": { | |
| "is_speech": bool(result.vad_result["is_speech"]), | |
| "energy": round(float(result.vad_result["energy"]), 6), | |
| "snr": round(float(result.vad_result["snr"]), 2), | |
| "confidence": round(float(result.vad_result["confidence"]), 4), | |
| }, | |
| "audio_hash": result.audio_hash, | |
| "replay_score": round(float(result.replay_score), 4), | |
| "replay_high_confidence": bool(result.replay_high_confidence), | |
| "speaker_match": None if result.speaker_match is None else round(float(result.speaker_match), 4), | |
| "speaker_similarity": None if result.speaker_similarity is None else round(float(result.speaker_similarity), 4), | |
| "matched_speaker": result.matched_speaker, | |
| "fused_synthetic_probability": round(float(result.fused_synthetic_probability), 4), | |
| "calibrated_synthetic_probability": round(float(result.calibrated_synthetic_probability), 4), | |
| "calibrated_speaker_match": None if result.calibrated_speaker_match is None else round(float(result.calibrated_speaker_match), 4), | |
| "detector_consistency": round(float(result.detector_consistency), 4), | |
| "risk_level": result.risk_level, | |
| "risk_score": round(float(result.risk_score), 2), | |
| "risk_action": result.risk_action, | |
| "decision": result.decision, | |
| "behavioral_risk": (None if result.behavioral_risk is None | |
| else round(float(result.behavioral_risk), 2)), | |
| "context_risk": (None if result.context_risk is None | |
| else round(float(result.context_risk), 2)), | |
| "risk_escalated": bool(result.risk_escalated), | |
| "detections": result.detections, | |
| } | |
| await self.manager.send_json(connection_id, detection_msg) | |
| # Non-speech / low-quality statuses are informational only (no alerts) | |
| if result.status != "detected": | |
| return | |
| # Send threat alert if threat score is elevated | |
| if result.threat_score is not None and result.threat_score > 50: | |
| alert_msg = { | |
| "type": MessageType.THREAT_ALERT.value, | |
| "threat_score": round(result.threat_score, 2), | |
| "authenticity_score": round(result.authenticity_score, 2), | |
| "predicted_label": result.predicted_label, | |
| "confidence": round(result.confidence, 4), | |
| "timestamp_ms": result.timestamp_ms, | |
| "chunk_index": result.chunk_index, | |
| "audio_hash": result.audio_hash, | |
| } | |
| await self.manager.send_json(connection_id, alert_msg) | |
| # Send risk alert on HIGH / CRITICAL levels | |
| if result.risk_level in ("HIGH", "CRITICAL"): | |
| risk_msg = { | |
| "type": MessageType.RISK_ALERT.value, | |
| "risk_level": result.risk_level, | |
| "risk_score": round(float(result.risk_score), 2), | |
| "action": result.risk_action, | |
| "decision": result.decision, | |
| "reasons": list(result.risk_reasons), | |
| "calibrated_synthetic_probability": round(float(result.calibrated_synthetic_probability), 4), | |
| "speaker_match": None if result.speaker_match is None else round(float(result.speaker_match), 4), | |
| "replay_probability": round(float(result.replay_score), 4), | |
| "consecutive_level": 2 if result.risk_escalated else 1, | |
| "escalated": bool(result.risk_escalated), | |
| "timestamp_ms": result.timestamp_ms, | |
| "chunk_index": result.chunk_index, | |
| "audio_hash": result.audio_hash, | |
| } | |
| await self.manager.send_json(connection_id, risk_msg) | |
| self._logger.warning( | |
| f"RISK_ALERT [{result.risk_level}] for {connection_id}: " | |
| f"action={result.risk_action}, risk={result.risk_score:.1f}" | |
| ) | |
| # Create real in-app alert (+ simulated SMS/email), pushed on WS. | |
| if self.alerts is not None: | |
| alert_type = "fraud_team" if result.risk_level == "CRITICAL" else "supervisor" | |
| severity = "critical" if result.risk_level == "CRITICAL" else "warning" | |
| reasons = ", ".join(result.risk_reasons) or "no reasons" | |
| self.alerts.send( | |
| alert_type=alert_type, | |
| severity=severity, | |
| message=(f"risk={result.risk_score:.0f} level={result.risk_level} " | |
| f"decision={result.decision} reasons={reasons}"), | |
| call_id=connection_id, | |
| connection_id=connection_id, | |
| ) | |
| # Proactive verification workflow from risk decision. | |
| if result.decision in ("VERIFY", "BLOCK") and self.verification is not None: | |
| try: | |
| flow = self.verification.begin( | |
| connection_id=connection_id, | |
| call_id=connection_id, | |
| risk_level=result.risk_level, | |
| reason=" | ".join(result.risk_reasons), | |
| ) | |
| await self.manager.send_json(connection_id, { | |
| "type": MessageType.VERIFY_REQUIRED.value, | |
| "flow_id": flow.flow_id, | |
| "call_id": flow.call_id, | |
| "state": flow.state, | |
| "risk_level": flow.risk_level, | |
| "reason": flow.reason, | |
| "attempts": flow.attempts, | |
| "prototype": True, | |
| "timestamp_ms": result.timestamp_ms, | |
| "message": ("Verification requested (SIMULATED OTP workflow for the demo); " | |
| "submit a code via POST /api/verification/submit."), | |
| }) | |
| except Exception as e: # verification must never break the audio loop | |
| self._logger.warning(f"verification begin failed for {connection_id}: {e}") | |
| # Send kill-switch message if triggered | |
| if result.kill_switch_triggered: | |
| kill_msg = { | |
| "type": MessageType.KILL_SWITCH.value, | |
| "action": "TERMINATE_CALL", | |
| "reason": result.kill_switch_reason, | |
| "threat_score": round(result.threat_score, 2), | |
| "consecutive_detections": 2, | |
| "timestamp_ms": result.timestamp_ms, | |
| "audio_hash": result.audio_hash, | |
| } | |
| await self.manager.send_json(connection_id, kill_msg) | |
| self._logger.warning(f"KILL_SWITCH triggered for {connection_id}: {result.kill_switch_reason}") | |
| async def _handle_config(self, connection_id: str, engine: InferenceEngine, data: dict) -> None: | |
| """Update configuration for this connection's engine.""" | |
| threat_threshold = data.get("threat_threshold") | |
| consecutive_windows = data.get("consecutive_windows") | |
| enable_vad = data.get("enable_vad") | |
| enable_kill_switch = data.get("enable_kill_switch") | |
| if threat_threshold is not None: | |
| engine.kill_switch.threshold = float(threat_threshold) | |
| if consecutive_windows is not None: | |
| engine.kill_switch.consecutive_windows = int(consecutive_windows) | |
| if enable_vad is not None: | |
| engine.processor.vad.enabled = not enable_vad # placeholder, adjust as needed | |
| if enable_kill_switch is not None: | |
| engine.kill_switch.enabled = bool(enable_kill_switch) | |
| self._logger.info(f"Config updated for {connection_id}: threshold={threat_threshold}") | |
| # Send updated stats | |
| await self._send_stats(connection_id, engine) | |
| async def _handle_ping(self, connection_id: str, data: dict) -> None: | |
| """Respond to client ping with server timestamp.""" | |
| client_ts = data.get("timestamp_ms", 0) | |
| await self.manager.send_json(connection_id, { | |
| "type": MessageType.PONG.value, | |
| "timestamp_ms": client_ts, | |
| "server_time_ms": timestamp_ms(), | |
| }) | |
| async def _handle_reset(self, connection_id: str, engine: InferenceEngine) -> None: | |
| """Reset engine state for a fresh stream.""" | |
| engine.reset() | |
| await self.manager.send_json(connection_id, { | |
| "type": MessageType.STATS.value, | |
| "total_chunks": 0, | |
| "avg_inference_ms": 0.0, | |
| "kill_switch_triggers": 0, | |
| "errors": 0, | |
| "buffer_fill_ratio": 0.0, | |
| "uptime_ms": 0, | |
| }) | |
| async def _send_ready(self, connection_id: str, engine: InferenceEngine) -> None: | |
| """Send ready message with model info.""" | |
| model_obj = getattr(engine, "model", None) | |
| model_info = None | |
| if model_obj is not None and hasattr(model_obj, "get_model_info"): | |
| try: | |
| model_info = model_obj.get_model_info() | |
| except Exception: | |
| model_info = None | |
| if model_info is None: | |
| model_info = { | |
| "architecture": "VoiceGuardRawNet", | |
| "quantized": settings.quantized, | |
| "num_classes": settings.num_classes, | |
| "class_labels": settings.class_labels, | |
| "input_shape": [1, 1, 24000], | |
| "sample_rate": settings.sample_rate, | |
| } | |
| config = { | |
| "sample_rate": settings.sample_rate, | |
| "chunk_duration_ms": settings.chunk_duration_ms, | |
| "buffer_duration_s": settings.buffer_duration_s, | |
| "window_duration_s": settings.window_duration_s, | |
| "window_step_duration_s": settings.window_step_duration_s, | |
| "threat_threshold": engine.kill_switch.threshold, | |
| "consecutive_windows": engine.kill_switch.consecutive_windows, | |
| "quantized": settings.quantized, | |
| "speaker_enrolled": engine.speaker_verifier.enrolled, | |
| "storage_root": str(self.storage.root) if self.storage is not None else None, | |
| } | |
| await self.manager.send_json(connection_id, { | |
| "type": MessageType.READY.value, | |
| "model_info": model_info, | |
| "config": config, | |
| }) | |
| async def _send_stats(self, connection_id: str, engine: InferenceEngine) -> None: | |
| """Send current engine stats.""" | |
| stats = engine.get_stats() | |
| await self.manager.send_json(connection_id, { | |
| "type": MessageType.STATS.value, | |
| "total_chunks": stats.total_chunks, | |
| "avg_inference_ms": round(stats.avg_inference_ms, 2), | |
| "kill_switch_triggers": stats.kill_switch_triggers, | |
| "errors": stats.errors, | |
| "buffer_fill_ratio": round(engine.buffer.fill_ratio, 4), | |
| "uptime_ms": 0, | |
| }) | |
| async def _send_error(self, connection_id: str, code: str, message: str) -> None: | |
| """Send error message to connection.""" | |
| await self.manager.send_json(connection_id, { | |
| "type": MessageType.ERROR.value, | |
| "code": code, | |
| "message": message, | |
| }) | |
| def _push_alert(self, connection_id: str, alert_data: dict) -> None: | |
| """Sync callback from AlertService -> async WS push (thread-safe).""" | |
| loop = self._loop | |
| if loop is None or loop.is_closed(): | |
| return | |
| msg = dict(alert_data) | |
| msg["type"] = MessageType.ALERT.value | |
| try: | |
| asyncio.run_coroutine_threadsafe( | |
| self.manager.send_json(connection_id, msg), loop | |
| ) | |
| except Exception: # pragma: no cover - push must never break alerts | |
| pass | |
| def _get_quantized_layers(self, engine: InferenceEngine) -> list: | |
| """Safely get list of quantized layer names.""" | |
| try: | |
| model = engine.model.model if hasattr(engine.model, "model") else None | |
| if model is None: | |
| return [] | |
| return [ | |
| name for name, module in model.named_modules() | |
| if hasattr(module, 'weight') and hasattr(module.weight, 'dtype') and module.weight.dtype == torch.qint8 | |
| ] | |
| except Exception: | |
| return [] | |
| async def _cleanup(self, connection_id: str) -> None: | |
| """Cleanup connection resources.""" | |
| if connection_id in self._engines: | |
| self._engines[connection_id].reset() | |
| del self._engines[connection_id] | |
| if self.verification is not None: | |
| try: | |
| self.verification.close(connection_id) | |
| except Exception: | |
| pass | |
| await self.manager.disconnect(connection_id) |