diff --git a/backend/devices/router.py b/backend/devices/router.py index d5d9c83..293613c 100644 --- a/backend/devices/router.py +++ b/backend/devices/router.py @@ -13,6 +13,7 @@ from devices.models import ( ResetStatsRequest, ResetStatsResult, ) from devices import service +from users import service as users_service import database as mqtt_db from mqtt.models import DeviceAlertEntry, DeviceAlertsResponse from shared.firebase import get_db as get_firestore @@ -461,10 +462,15 @@ async def add_user_to_device( elif isinstance(entry, str): existing_ids.add(entry.split("/")[-1]) + # user_list and the user's device_serials (MQTT ACL lookup) commit together. + user_ref = fs.collection("users").document(body.user_id) + batch = fs.batch() if body.user_id not in existing_ids: - user_ref = fs.collection("users").document(body.user_id) user_list.append(user_ref) - device_ref.update({"user_list": user_list}) + batch.update(device_ref, {"user_list": user_list}) + users_service.stage_device_serial_link(batch, user_ref, users_service.device_serial_of(data), linked=True) + batch.commit() + users_service.invalidate_mqtt_acl_cache(user_doc.to_dict()) await log_action(db, _user.sub, _user.name or _user.email, "UPDATE", "device", device_id, device_id, meta={"action_detail": "user_added", @@ -500,7 +506,15 @@ async def remove_user_from_device( # Remove any entry that resolves to this user_id (handles both DocRef and string paths) new_list = [entry for entry in user_list if not resolves_to(entry, user_id)] - device_ref.update({"user_list": new_list}) + batch = fs.batch() + batch.update(device_ref, {"user_list": new_list}) + user_ref = fs.collection("users").document(user_id) + user_doc = user_ref.get() + if user_doc.exists: + users_service.stage_device_serial_link(batch, user_ref, users_service.device_serial_of(data), linked=False) + batch.commit() + if user_doc.exists: + users_service.invalidate_mqtt_acl_cache(user_doc.to_dict()) await log_action(db, _user.sub, _user.name or _user.email, "UPDATE", "device", device_id, device_id, meta={"action_detail": "user_removed", diff --git a/backend/mqtt/app_users.py b/backend/mqtt/app_users.py new file mode 100644 index 0000000..e79db32 --- /dev/null +++ b/backend/mqtt/app_users.py @@ -0,0 +1,89 @@ +""" +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() diff --git a/backend/users/service.py b/backend/users/service.py index 4695ba3..a1e426c 100644 --- a/backend/users/service.py +++ b/backend/users/service.py @@ -1,6 +1,6 @@ from datetime import datetime -from google.cloud.firestore_v1 import DocumentReference +from google.cloud.firestore_v1 import DocumentReference, ArrayUnion, ArrayRemove from firebase_admin import auth as firebase_auth from shared.firebase import get_db, get_bucket @@ -9,6 +9,37 @@ from users.models import UserCreate, UserUpdate, UserInDB COLLECTION = "users" +# Firestore "in" queries accept at most 30 values. +_IN_QUERY_LIMIT = 30 + + +def device_serial_of(device_data: dict) -> str: + """The board serial as used in MQTT topics (vesper/{serial}/...). + + `serial_number` is what gets flashed into NVS and what the firmware uses + as its MQTT username/topic id. Legacy docs predating the lifecycle system + only carry it in `device_id`. + """ + return (device_data.get("serial_number") or device_data.get("device_id") or "").strip() + + +def stage_device_serial_link(batch, user_ref: DocumentReference, serial: str, linked: bool) -> None: + """Add (linked=True) or remove the serial from the user's `device_serials` + inside the caller's write batch, so it commits together with the device's + `user_list` change.""" + if not serial: + return + op = ArrayUnion([serial]) if linked else ArrayRemove([serial]) + batch.update(user_ref, {"device_serials": op}) + + +def invalidate_mqtt_acl_cache(user_data: dict | None) -> None: + """Drop the MQTT ACL cache entry for this user's Firebase uid.""" + uid = (user_data or {}).get("uid") or "" + if uid: + from mqtt.app_users import invalidate + invalidate(uid) + def _convert_firestore_value(val): """Convert Firestore-specific types (Timestamp, DocumentReference) to strings.""" @@ -121,6 +152,7 @@ def update_user(user_doc_id: str, data: UserUpdate) -> UserInDB: update_data = data.model_dump(exclude_none=True) doc_ref.update(update_data) + invalidate_mqtt_acl_cache(doc.to_dict()) # status may have changed updated_doc = doc_ref.get() return _doc_to_user(updated_doc) @@ -142,6 +174,7 @@ def delete_user(user_doc_id: str) -> None: pass doc_ref.delete() + invalidate_mqtt_acl_cache(doc.to_dict()) def block_user(user_doc_id: str) -> UserInDB: @@ -153,6 +186,7 @@ def block_user(user_doc_id: str) -> UserInDB: raise NotFoundError("User") doc_ref.update({"status": "blocked"}) + invalidate_mqtt_acl_cache(doc.to_dict()) updated_doc = doc_ref.get() return _doc_to_user(updated_doc) @@ -166,6 +200,7 @@ def unblock_user(user_doc_id: str) -> UserInDB: raise NotFoundError("User") doc_ref.update({"status": "active"}) + invalidate_mqtt_acl_cache(doc.to_dict()) updated_doc = doc_ref.get() return _doc_to_user(updated_doc) @@ -202,9 +237,15 @@ def assign_device(user_doc_id: str, device_doc_id: str) -> UserInDB: already_assigned = True break + # Device user_list and the user's device_serials are written in one atomic batch. + # device_serials is re-synced even when already assigned, to heal drift. + batch = db.batch() if not already_assigned: user_list.append(user_path) - device_ref.update({"user_list": user_list}) + batch.update(device_ref, {"user_list": user_list}) + stage_device_serial_link(batch, user_ref, device_serial_of(device_data), linked=True) + batch.commit() + invalidate_mqtt_acl_cache(user_doc.to_dict()) return _doc_to_user(user_ref.get()) @@ -238,13 +279,17 @@ def unassign_device(user_doc_id: str, device_doc_id: str) -> UserInDB: elif entry != user_path: new_list.append(entry) - device_ref.update({"user_list": new_list}) + batch = db.batch() + batch.update(device_ref, {"user_list": new_list}) + stage_device_serial_link(batch, user_ref, device_serial_of(device_data), linked=False) + batch.commit() + invalidate_mqtt_acl_cache(user_doc.to_dict()) return _doc_to_user(user_ref.get()) def get_user_devices(user_doc_id: str) -> list[dict]: - """Get all devices assigned to a user.""" + """Get all devices assigned to a user, via the user's `device_serials`.""" db = get_db() # Verify user exists @@ -253,25 +298,32 @@ def get_user_devices(user_doc_id: str) -> list[dict]: if not user_doc.exists: raise NotFoundError("User") - user_path = f"users/{user_doc_id}" + serials = [s for s in dict.fromkeys(user_doc.to_dict().get("device_serials") or []) if s] + + # Match on serial_number first; fall back to the legacy device_id field for the rest. + found: dict[str, object] = {} + for field in ("serial_number", "device_id"): + remaining = [s for s in serials if s not in found] + for i in range(0, len(remaining), _IN_QUERY_LIMIT): + chunk = remaining[i:i + _IN_QUERY_LIMIT] + for doc in db.collection("devices").where(field, "in", chunk).stream(): + sn = doc.to_dict().get(field) + if sn not in found: + found[sn] = doc - # Search all devices for this user in their user_list devices = [] - for doc in db.collection("devices").stream(): + for sn in serials: + doc = found.get(sn) + if doc is None: + continue data = doc.to_dict() - user_list = data.get("user_list", []) - - for entry in user_list: - entry_path = entry.path if isinstance(entry, DocumentReference) else entry - if entry_path == user_path: - devices.append({ - "id": doc.id, - "device_name": data.get("device_name", ""), - "device_id": data.get("device_id", ""), - "device_location": data.get("device_location", ""), - "is_Online": data.get("is_Online", False), - }) - break + devices.append({ + "id": doc.id, + "device_name": data.get("device_name", ""), + "device_id": data.get("device_id", ""), + "device_location": data.get("device_location", ""), + "is_Online": data.get("is_Online", False), + }) return devices @@ -279,7 +331,7 @@ def get_user_devices(user_doc_id: str) -> list[dict]: def set_password(user_doc_id: str, new_password: str) -> None: """Set a Firebase Auth password for a user via their Firestore document ID. - Requires the user document to have a non-empty `uid` field — populated + Requires the user document to have a non-empty `uid` field — populated automatically for users who registered via the Flutter app. """ if not new_password or len(new_password) < 6: @@ -293,7 +345,7 @@ def set_password(user_doc_id: str, new_password: str) -> None: uid = doc.to_dict().get("uid", "") if not uid: - raise ValidationError("This user has no Firebase Auth UID — they may not have signed up via the app yet.") + raise ValidationError("This user has no Firebase Auth UID — they may not have signed up via the app yet.") try: firebase_auth.update_user(uid, password=new_password)