channel / socket_service.py
achin42's picture
Upload folder using huggingface_hub
bc1e6dc verified
Raw
History Blame Contribute Delete
5.15 kB
from typing import Dict, Optional, Any
from uuid import UUID
from fastapi import WebSocket, BackgroundTasks
from datetime import datetime, timedelta, timezone
import asyncio
from fastapi.encoders import jsonable_encoder
# Constants
HEARTBEAT_INTERVAL = 30 # seconds
CONNECTION_TIMEOUT = 24 * 60 * 60 # 24 hours in seconds
CLEANUP_INTERVAL = 60 * 60 # 1 hour in seconds
class SocketService:
def __init__(self):
self.active_connections: Dict[str, WebSocket] = {}
self.connection_times: Dict[str, datetime] = {}
self.background_tasks: BackgroundTasks = None
def _get_connection_key(self, user_id: UUID, device_id: str) -> str:
device_id = device_id.strip().replace(" ", "_")
return f"socket:{user_id}:{device_id}"
async def connect(self, websocket: WebSocket, user_id: UUID, device_id: str, metadata: Optional[Dict[str, Any]] = None):
if not device_id.strip():
raise ValueError("Device ID cannot be empty")
connection_key = self._get_connection_key(user_id, device_id)
# If there's an existing connection for this device, disconnect it first
if connection_key in self.active_connections:
try:
existing_ws = self.active_connections[connection_key]
await existing_ws.close()
except Exception:
pass
finally:
del self.active_connections[connection_key]
del self.connection_times[connection_key]
# Accept the new connection
await websocket.accept()
self.active_connections[connection_key] = websocket
self.connection_times[connection_key] = datetime.now(timezone.utc)
async def disconnect(self, user_id: UUID, device_id: str):
connection_key = self._get_connection_key(user_id, device_id)
if connection_key in self.active_connections:
websocket = self.active_connections[connection_key]
try:
await websocket.close()
except Exception:
pass
finally:
del self.active_connections[connection_key]
del self.connection_times[connection_key]
async def update_heartbeat(self, user_id: UUID, device_id: str):
connection_key = self._get_connection_key(user_id, device_id)
if connection_key in self.active_connections:
self.connection_times[connection_key] = datetime.now(timezone.utc)
async def cleanup_expired_connections(self, iterations: Optional[int] = None, cleanup_interval: Optional[float] = None):
"""
Cleanup expired connections.
:param iterations: If set, runs only this many iterations (for testing)
:param cleanup_interval: Override default CLEANUP_INTERVAL (useful for testing)
"""
iteration_count = 0
interval = cleanup_interval if cleanup_interval is not None else CLEANUP_INTERVAL
while True:
try:
current_time = datetime.now(timezone.utc)
expired_keys = [
key for key, last_time in self.connection_times.items()
if current_time - last_time > timedelta(seconds=CONNECTION_TIMEOUT)
]
for key in expired_keys:
user_id, device_id = key.split(":")[1:]
await self.disconnect(UUID(user_id), device_id)
if iterations is not None:
iteration_count += 1
if iteration_count >= iterations:
break
await asyncio.sleep(interval)
except Exception:
if iterations is not None:
iteration_count += 1
if iteration_count >= iterations:
break
await asyncio.sleep(interval)
async def get_user_connections(self, user_id: UUID) -> Dict[str, WebSocket]:
pattern = f"socket:{user_id}:"
connections = {}
for key, ws in self.active_connections.items():
if key.startswith(pattern):
# Extract the device_id part from the key
device_id = key.split(":")[-1]
connections[device_id] = ws
return connections
async def broadcast_message(self, user_id: UUID, message: Any, exclude_device_id: Optional[str] = None) -> int:
if exclude_device_id:
exclude_device_id = exclude_device_id.strip().replace(" ", "_")
connections = await self.get_user_connections(user_id)
json_message = jsonable_encoder(message)
broadcast_count = 0
for device_id, websocket in connections.items():
if exclude_device_id and device_id == exclude_device_id:
continue
try:
await websocket.send_json(json_message)
broadcast_count += 1
except Exception:
continue
return broadcast_count