feat(users): store device_serials on user docs for direct user->device lookup
Adds a `device_serials: [string]` array to Firestore `users` docs so the
MQTT ACL (and get_user_devices) can answer "which boards may this user
reach?" without streaming the entire devices collection.
- The serial is the value used in MQTT topics vesper/{serial}/...: the
device doc's `serial_number` (flashed into NVS, used by the firmware as
its MQTT id), falling back to the legacy `device_id` for old docs.
Centralised in users.service.device_serial_of().
- assign_device / unassign_device now write the device's user_list and the
user's device_serials (ArrayUnion/ArrayRemove) in one atomic batch.
- The device Manage tab endpoints (POST/DELETE /api/devices/{id}/user-list)
also edit user_list, so they get the same batched sync - otherwise the
most common assignment path would silently leave device_serials stale.
- get_user_devices resolves devices via device_serials with chunked
Firestore "in" queries instead of a full collection scan. Requires the
backfill script (next commit) to be run for existing assignments.
- New mqtt/app_users.py: resolves users by the `uid` FIELD (not doc id -
create_user uses .add(), FlutterFlow uses uid as doc id) with a 60s
in-process TTL cache. Assign/unassign, update, block/unblock and delete
invalidate that uid's entry.
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
This commit is contained in:
@@ -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])
|
||||
|
||||
if body.user_id not in existing_ids:
|
||||
# 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_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",
|
||||
|
||||
@@ -0,0 +1,89 @@
|
||||
"""
|
||||
Firebase app-user lookup for the MQTT auth/ACL endpoints.
|
||||
|
||||
App users connect to Mosquitto as "app_<firebase_uid>". 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()
|
||||
+67
-15
@@ -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,17 +298,25 @@ 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", ""),
|
||||
@@ -271,7 +324,6 @@ def get_user_devices(user_doc_id: str) -> list[dict]:
|
||||
"device_location": data.get("device_location", ""),
|
||||
"is_Online": data.get("is_Online", False),
|
||||
})
|
||||
break
|
||||
|
||||
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)
|
||||
|
||||
Reference in New Issue
Block a user