""" Firebase app-user lookup for the MQTT auth/ACL endpoints. App users connect to Mosquitto as "app_". Their allowed devices come from the `device_serials` array on their `users` doc. User docs are resolved by the `uid` FIELD, never by document ID: docs created by the Console (users.service.create_user) get random IDs via .add(), while FlutterFlow uses the Firebase uid as the doc ID. Per-message ACL checks would otherwise hit Firestore on every publish and delivery, so lookups are cached in-process for CACHE_TTL_SECONDS. Anything that changes a user's device_serials or status calls invalidate(uid). """ import logging import threading import time from dataclasses import dataclass from shared.firebase import get_db logger = logging.getLogger("mqtt.app_users") USERS_COLLECTION = "users" CACHE_TTL_SECONDS = 60.0 @dataclass(frozen=True) class AppUser: uid: str blocked: bool device_serials: frozenset[str] _cache: dict[str, tuple[float, AppUser | None]] = {} _lock = threading.Lock() def _load(uid: str) -> AppUser | None: db = get_db() if db is None: raise RuntimeError("Firestore not initialized") docs = list(db.collection(USERS_COLLECTION).where("uid", "==", uid).limit(5).stream()) if not docs: return None if len(docs) > 1: # Duplicate profiles for one Firebase account: be conservative on # blocked, permissive on devices (union), and make it visible. logger.warning("Multiple user docs share uid=%s: %s", uid, [d.id for d in docs]) blocked = False serials: set[str] = set() for doc in docs: data = doc.to_dict() or {} if data.get("status") == "blocked": blocked = True serials.update(s for s in (data.get("device_serials") or []) if isinstance(s, str) and s) return AppUser(uid=uid, blocked=blocked, device_serials=frozenset(serials)) def get_app_user(uid: str, *, use_cache: bool = True) -> AppUser | None: """Return the app user for a Firebase uid, or None if no user doc has that uid. use_cache=False forces a Firestore read (and refreshes the cache). """ now = time.monotonic() if use_cache: with _lock: hit = _cache.get(uid) if hit and hit[0] > now: return hit[1] user = _load(uid) with _lock: _cache[uid] = (now + CACHE_TTL_SECONDS, user) return user def invalidate(uid: str) -> None: with _lock: _cache.pop(uid, None) def clear_cache() -> None: with _lock: _cache.clear()