diff --git a/backend/tests/conftest.py b/backend/tests/conftest.py new file mode 100644 index 0000000..37048cc --- /dev/null +++ b/backend/tests/conftest.py @@ -0,0 +1,5 @@ +import sys +from pathlib import Path + +# Tests import backend modules the same way the app does (cwd = backend/). +sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) diff --git a/backend/tests/test_mqtt_auth.py b/backend/tests/test_mqtt_auth.py new file mode 100644 index 0000000..175d963 --- /dev/null +++ b/backend/tests/test_mqtt_auth.py @@ -0,0 +1,357 @@ +""" +Tests for the mosquitto-go-auth HTTP backend endpoints (mqtt/auth.py): +POST /mqtt/auth/user and POST /mqtt/auth/acl. + +firebase_admin token verification and Firestore are mocked; nothing here +touches the network. +""" + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient +from firebase_admin import auth as firebase_auth + +from config import settings +from mqtt import app_users +from mqtt import auth as mqtt_auth + +UID = "uidAlice123" +OTHER_UID = "uidBob456" +SERIAL = "BSVSPR-26C13X-STD01R-X7KQA" +FOREIGN_SERIAL = "PV25L22BP01R01" +APP_USER = f"app_{UID}" +APP_CLIENT = f"app_{UID}_phone1" +GOOD_TOKEN = "good-token" + + +# --------------------------------------------------------------------------- +# Fakes +# --------------------------------------------------------------------------- + +class _FakeDoc: + def __init__(self, doc_id, data): + self.id = doc_id + self._data = data + + def to_dict(self): + return dict(self._data) + + +class _FakeQuery: + def __init__(self, docs, field=None, value=None): + self._docs = docs + self._field = field + self._value = value + + def where(self, field, op, value): + assert op == "==" + return _FakeQuery(self._docs, field, value) + + def limit(self, _n): + return self + + def stream(self): + return [d for d in self._docs if d.to_dict().get(self._field) == self._value] + + +class FakeFirestore: + """Just enough of the Firestore client for app_users._load().""" + + def __init__(self): + self.users: list[_FakeDoc] = [] + self.reads = 0 + + def add_user(self, doc_id, **data): + self.users.append(_FakeDoc(doc_id, data)) + + def collection(self, name): + assert name == "users" + self.reads += 1 + return _FakeQuery(self.users) + + +@pytest.fixture +def fs(monkeypatch): + db = FakeFirestore() + # Doc ID deliberately != uid, like Console-created users. + db.add_user("randomDocId1", uid=UID, status="active", device_serials=[SERIAL]) + db.add_user(OTHER_UID, uid=OTHER_UID, status="blocked", device_serials=[FOREIGN_SERIAL]) + monkeypatch.setattr(app_users, "get_db", lambda: db) + app_users.clear_cache() + yield db + app_users.clear_cache() + + +@pytest.fixture +def verify(monkeypatch): + """Mock verify_id_token. Set `verify.result` to a dict or an exception.""" + + class _Verify: + result = {"uid": UID} + calls = [] + + def __call__(self, token, check_revoked=False): + self.calls.append((token, check_revoked)) + if isinstance(self.result, Exception): + raise self.result + return self.result + + v = _Verify() + v.calls = [] + monkeypatch.setattr(mqtt_auth.firebase_auth, "verify_id_token", v) + return v + + +@pytest.fixture +def client(): + app = FastAPI() + app.include_router(mqtt_auth.router) + return TestClient(app) + + +@pytest.fixture(autouse=True) +def _defaults(monkeypatch): + monkeypatch.setattr(settings, "mqtt_secret", "test-secret") + monkeypatch.setattr(settings, "mqtt_allow_legacy_password", True) + mqtt_auth._legacy_log_last.clear() + + +def auth_user(client, username, password, clientid="c1"): + return client.post( + "/mqtt/auth/user", + data={"username": username, "password": password, "clientid": clientid}, + ).status_code + + +def acl(client, username, topic, acc, clientid=APP_CLIENT): + return client.post( + "/mqtt/auth/acl", + data={"username": username, "topic": topic, "clientid": clientid, "acc": acc}, + ).status_code + + +# --------------------------------------------------------------------------- +# /mqtt/auth/user — devices +# --------------------------------------------------------------------------- + +@pytest.mark.parametrize("username", [SERIAL, FOREIGN_SERIAL, f"{FOREIGN_SERIAL}-kiosk"]) +def test_device_hmac_ok(client, username): + assert auth_user(client, username, mqtt_auth._derive_password(username)) == 200 + + +def test_device_hmac_wrong_password(client): + assert auth_user(client, SERIAL, "0" * 32) == 403 + + +def test_device_hmac_of_other_serial_rejected(client): + assert auth_user(client, SERIAL, mqtt_auth._derive_password(FOREIGN_SERIAL)) == 403 + + +def test_legacy_password_accepted_when_flag_on(client): + assert auth_user(client, FOREIGN_SERIAL, "vesper") == 200 + assert auth_user(client, f"{FOREIGN_SERIAL}-kiosk", "vesper") == 200 + + +def test_legacy_password_rejected_when_flag_off(client, monkeypatch): + monkeypatch.setattr(settings, "mqtt_allow_legacy_password", False) + assert auth_user(client, FOREIGN_SERIAL, "vesper") == 403 + # HMAC still works with the flag off + assert auth_user(client, FOREIGN_SERIAL, mqtt_auth._derive_password(FOREIGN_SERIAL)) == 200 + + +@pytest.mark.parametrize("username", ["somebody", "NodeRED", "pv25l22bp01r01", "PV-", "PV 25"]) +def test_legacy_password_rejected_for_non_device_usernames(client, username): + assert auth_user(client, username, "vesper") == 403 + + +def test_legacy_password_never_accepted_for_app_users(client, fs, verify): + verify.result = firebase_auth.InvalidIdTokenError("not a jwt") + assert auth_user(client, APP_USER, "vesper") == 403 + + +def test_hmac_never_accepted_for_app_users(client, fs, verify): + verify.result = firebase_auth.InvalidIdTokenError("not a jwt") + assert auth_user(client, APP_USER, mqtt_auth._derive_password(APP_USER)) == 403 + + +def test_legacy_login_logged_once_per_hour(client, caplog): + with caplog.at_level("WARNING", logger="mqtt.auth"): + for _ in range(3): + assert auth_user(client, FOREIGN_SERIAL, "vesper") == 200 + assert auth_user(client, SERIAL, "vesper") == 200 + legacy = [r for r in caplog.records if "legacy password" in r.getMessage()] + assert [r.getMessage().split(" for ")[1].split(" ")[0] for r in legacy] == [FOREIGN_SERIAL, SERIAL] + + +# --------------------------------------------------------------------------- +# /mqtt/auth/user — app users +# --------------------------------------------------------------------------- + +def test_app_valid_token(client, fs, verify): + assert auth_user(client, APP_USER, GOOD_TOKEN) == 200 + assert verify.calls == [(GOOD_TOKEN, True)] # check_revoked=True + + +def test_app_token_for_other_uid(client, fs, verify): + verify.result = {"uid": OTHER_UID} + assert auth_user(client, APP_USER, GOOD_TOKEN) == 403 + + +def test_app_revoked_token(client, fs, verify): + verify.result = firebase_auth.RevokedIdTokenError("revoked") + assert auth_user(client, APP_USER, GOOD_TOKEN) == 403 + + +def test_app_expired_token(client, fs, verify): + verify.result = firebase_auth.ExpiredIdTokenError("expired", cause=None) + assert auth_user(client, APP_USER, GOOD_TOKEN) == 403 + + +def test_app_blocked_user(client, fs, verify): + verify.result = {"uid": OTHER_UID} + assert auth_user(client, f"app_{OTHER_UID}", GOOD_TOKEN) == 403 + + +def test_app_unknown_user(client, fs, verify): + verify.result = {"uid": "ghost"} + assert auth_user(client, "app_ghost", GOOD_TOKEN) == 403 + + +def test_app_empty_uid(client, fs, verify): + verify.result = {"uid": ""} + assert auth_user(client, "app_", GOOD_TOKEN) == 403 + + +def test_app_deny_never_logs_token(client, fs, verify, caplog): + secret_token = "eyJhbGciOi.SECRET-TOKEN-VALUE.sig" + verify.result = {"uid": OTHER_UID} + with caplog.at_level("DEBUG"): + assert auth_user(client, APP_USER, secret_token) == 403 + assert "token uid does not match" in caplog.text + assert secret_token not in caplog.text + + +# --------------------------------------------------------------------------- +# /mqtt/auth/acl — app users +# --------------------------------------------------------------------------- + +@pytest.mark.parametrize("leaf", ["control/ack", "status/heartbeat", "status/playback"]) +@pytest.mark.parametrize("acc", [mqtt_auth.ACC_READ, mqtt_auth.ACC_SUBSCRIBE]) +def test_app_acl_read_subscribe_allowed(client, fs, leaf, acc): + assert acl(client, APP_USER, f"vesper/{SERIAL}/{leaf}", acc) == 200 + + +def test_app_acl_publish_command_allowed(client, fs): + assert acl(client, APP_USER, f"vesper/{SERIAL}/control/command", mqtt_auth.ACC_WRITE) == 200 + + +@pytest.mark.parametrize("leaf", ["control/ack", "status/heartbeat", "status/playback", "system/info"]) +def test_app_acl_publish_other_topics_denied(client, fs, leaf): + assert acl(client, APP_USER, f"vesper/{SERIAL}/{leaf}", mqtt_auth.ACC_WRITE) == 403 + + +@pytest.mark.parametrize("acc", [mqtt_auth.ACC_READ, mqtt_auth.ACC_SUBSCRIBE]) +@pytest.mark.parametrize("leaf", ["control/command", "system/info", "system/alerts", "control/reports"]) +def test_app_acl_read_subscribe_other_topics_denied(client, fs, acc, leaf): + assert acl(client, APP_USER, f"vesper/{SERIAL}/{leaf}", acc) == 403 + + +@pytest.mark.parametrize("acc", [0, 3, 5, 8]) +def test_app_acl_unknown_acc_denied(client, fs, acc): + assert acl(client, APP_USER, f"vesper/{SERIAL}/control/ack", acc) == 403 + + +@pytest.mark.parametrize("topic", [ + f"vesper/{SERIAL}/#", + f"vesper/{SERIAL}/status/+", + "vesper/+/control/ack", + "#", +]) +@pytest.mark.parametrize("acc", [mqtt_auth.ACC_READ, mqtt_auth.ACC_WRITE, mqtt_auth.ACC_SUBSCRIBE]) +def test_app_acl_wildcards_denied(client, fs, topic, acc): + assert acl(client, APP_USER, topic, acc) == 403 + + +@pytest.mark.parametrize("acc", [mqtt_auth.ACC_READ, mqtt_auth.ACC_WRITE, mqtt_auth.ACC_SUBSCRIBE]) +def test_app_acl_foreign_serial_denied(client, fs, acc): + leaf = "control/command" if acc == mqtt_auth.ACC_WRITE else "control/ack" + assert acl(client, APP_USER, f"vesper/{FOREIGN_SERIAL}/{leaf}", acc) == 403 + + +@pytest.mark.parametrize("clientid", [ + "", + "phone1", + f"app_{UID}", # missing trailing "_" + f"app_{OTHER_UID}_x", # another user's prefix + f"app_{UID}x_phone", # uid prefix collision +]) +def test_app_acl_wrong_clientid_denied(client, fs, clientid): + assert acl(client, APP_USER, f"vesper/{SERIAL}/control/ack", mqtt_auth.ACC_SUBSCRIBE, clientid=clientid) == 403 + + +@pytest.mark.parametrize("topic", [ + f"vesper/{SERIAL}/control/command/extra", + f"vesper/{SERIAL}/control", + f"other/{SERIAL}/control/command", + "vesper//control/command", +]) +def test_app_acl_malformed_topics_denied(client, fs, topic): + assert acl(client, APP_USER, topic, mqtt_auth.ACC_WRITE) == 403 + + +def test_app_acl_blocked_user_denied(client, fs): + user = f"app_{OTHER_UID}" + assert acl(client, user, f"vesper/{FOREIGN_SERIAL}/control/ack", mqtt_auth.ACC_SUBSCRIBE, + clientid=f"{user}_phone") == 403 + + +def test_app_acl_unknown_user_denied(client, fs): + assert acl(client, "app_ghost", f"vesper/{SERIAL}/control/ack", mqtt_auth.ACC_SUBSCRIBE, + clientid="app_ghost_phone") == 403 + + +def test_app_acl_uses_cache_and_invalidate_refreshes(client, fs): + topic = f"vesper/{SERIAL}/control/ack" + assert acl(client, APP_USER, topic, mqtt_auth.ACC_SUBSCRIBE) == 200 + reads = fs.reads + for _ in range(5): + assert acl(client, APP_USER, topic, mqtt_auth.ACC_READ) == 200 + assert fs.reads == reads # served from cache + + # Unassign the device; the cached entry still allows until invalidated. + fs.users[0]._data["device_serials"] = [] + assert acl(client, APP_USER, topic, mqtt_auth.ACC_READ) == 200 + app_users.invalidate(UID) + assert acl(client, APP_USER, topic, mqtt_auth.ACC_READ) == 403 + + +def test_app_acl_cache_expires(client, fs, monkeypatch): + topic = f"vesper/{SERIAL}/control/ack" + assert acl(client, APP_USER, topic, mqtt_auth.ACC_READ) == 200 + fs.users[0]._data["device_serials"] = [] + clock = app_users.time.monotonic() + app_users.CACHE_TTL_SECONDS + 1 + monkeypatch.setattr(app_users.time, "monotonic", lambda: clock) + assert acl(client, APP_USER, topic, mqtt_auth.ACC_READ) == 403 + + +# --------------------------------------------------------------------------- +# /mqtt/auth/acl — devices (unchanged behaviour) +# --------------------------------------------------------------------------- + +@pytest.mark.parametrize("acc", [1, 2, 4]) +def test_device_acl_own_topics(client, acc): + assert acl(client, SERIAL, f"vesper/{SERIAL}/status/heartbeat", acc, clientid="") == 200 + + +@pytest.mark.parametrize("acc", [1, 2, 4]) +def test_device_acl_foreign_topics_denied(client, acc): + assert acl(client, SERIAL, f"vesper/{FOREIGN_SERIAL}/status/heartbeat", acc, clientid="") == 403 + + +def test_kiosk_acl_base_device_topics(client): + assert acl(client, f"{FOREIGN_SERIAL}-kiosk", f"vesper/{FOREIGN_SERIAL}/control/command", 2, clientid="") == 200 + assert acl(client, f"{FOREIGN_SERIAL}-kiosk", f"vesper/{SERIAL}/control/command", 2, clientid="") == 403 + + +def test_superuser_acl(client): + assert acl(client, "NodeRED", "vesper/#", 4, clientid="") == 200