"""Device credential repository backed by Redis.""" from __future__ import annotations from typing import Dict, List, Optional from pydantic import BaseModel, ValidationError from .core.logging import get_logger from .core.services.redis_client import redis_manager logger = get_logger() class DeviceCredentials(BaseModel): """Validated credential tuple for one device.""" device_id: str api_key: str secret: str is_active: bool class DeviceRepository: """Loads and stores device credentials in memory for low-latency lookup.""" def __init__(self) -> None: self.redis_client = redis_manager.client self._credentials_by_id: Dict[str, DeviceCredentials] = {} self._credentials_by_key: Dict[str, DeviceCredentials] = {} async def _load_credentials(self) -> None: """Read `device:*` hashes from Redis and cache them in memory.""" if not self.redis_client: raise ConnectionError("Redis client is not configured.") logger.info("Loading device credentials from Redis") loaded = 0 async for key in self.redis_client.scan_iter("device:*"): data = await self.redis_client.hgetall(key) if not data: continue data["device_id"] = key.split(":", 1)[1] data["is_active"] = data.get("is_active") == "1" try: creds = DeviceCredentials(**data) except ValidationError as exc: logger.error("Skipping invalid device credentials", redis_key=key, error=str(exc)) continue self._credentials_by_id[creds.device_id] = creds self._credentials_by_key[creds.api_key] = creds loaded += 1 logger.info("Device credentials loaded", count=loaded) def get_by_api_key(self, api_key: str) -> Optional[DeviceCredentials]: """Get credentials by api key.""" return self._credentials_by_key.get(api_key) def get_by_device_id(self, device_id: str) -> Optional[DeviceCredentials]: """Get credentials by device id.""" return self._credentials_by_id.get(device_id) def get_all_devices(self) -> List[DeviceCredentials]: """Return all loaded devices.""" return list(self._credentials_by_id.values()) def count(self) -> int: """Return count of loaded credentials.""" return len(self._credentials_by_id)