From 21062cc8c32105a7066a23e0e970c56dc4702550 Mon Sep 17 00:00:00 2001 From: Adamskiee Date: Sun, 16 Aug 2026 11:08:07 +0800 Subject: [PATCH 1/2] fix(server-api): harden error handling logs --- server/app/api/admin.py | 144 ++++++++++++------ server/app/api/peer_connection.py | 74 ++++++--- server/app/api/update_info.py | 9 +- server/app/db_operations/GPS_manager.py | 17 ++- server/app/db_operations/auth.py | 21 ++- .../app/db_operations/connection_manager.py | 23 ++- server/app/db_operations/token.py | 13 +- server/app/structured_logging.py | 16 ++ server/app/tests/test_admin_me.py | 51 ++++++- server/app/tests/test_gps.py | 40 ++++- server/app/tests/test_websocket_pool.py | 25 +++ 11 files changed, 354 insertions(+), 79 deletions(-) create mode 100644 server/app/structured_logging.py diff --git a/server/app/api/admin.py b/server/app/api/admin.py index 9c532a89..a0e96c1f 100644 --- a/server/app/api/admin.py +++ b/server/app/api/admin.py @@ -1,6 +1,7 @@ from app.models.announcement import Announcement, PriorityType, AnnouncementStatusType, AudienceType from app.models.users import User from datetime import datetime, timedelta, timezone +import logging import os from uuid import UUID, uuid4 from math import ceil @@ -46,6 +47,8 @@ refresh_token, ) from app.models.rescuer import Rescuer +from app.structured_logging import log_context +from app.limiter import limiter from app.models.token import Token from app.models.users import ( User, @@ -59,6 +62,7 @@ router = APIRouter( prefix="/api/admin", tags=["admin"], responses={404: {"description": "Not Found"}} ) +logger = logging.getLogger("app") class AdminLoginResponse(BaseModel): @@ -125,6 +129,7 @@ async def login_for_access_token( @router.post("/refresh") +@limiter.limit("10/minute") async def refresh_access_token( request: Request, response: Response, session: SessionDep ): @@ -156,8 +161,12 @@ async def refresh_access_token( # max_age=604800, # 7 days # ) return {"status": "refreshed", "refresh_token": new_access_token.refresh_token, "access_token": new_access_token.access_token} - except: - raise HTTPException(status_code=401) + except Exception as exc: + logger.info( + "Admin refresh token validation failed", + extra=log_context(None, "admin_refresh_failed"), + ) + raise HTTPException(status_code=401) from exc @router.post("/logout") @@ -244,6 +253,11 @@ def perform_ping_probe(): # Success if return code is 0 ping_history.append(result.returncode == 0) except Exception: + logger.warning( + "Network ping probe failed", + exc_info=True, + extra=log_context(None, "network_ping_probe_failed"), + ) ping_history.append(False) @@ -343,7 +357,7 @@ def get_network_details(): struct.pack("256s", iface[:15].encode("utf-8")), )[20:24] ) - except Exception: + except OSError: # Likely no IPv4 assigned to this interface ip_addr = "N/A" @@ -461,24 +475,26 @@ def get_admin_users( } -def makeAdmin(user: User, session: SessionDep): +def makeAdmin(user: User, session: SessionDep, commit: bool = True): if not user.id: raise HTTPException(500) admin = Admin(user_id=user.id) session.add(admin) - session.commit() - session.refresh(admin) + if commit: + session.commit() + session.refresh(admin) -def makeRescuer(user: User, session: SessionDep): +def makeRescuer(user: User, session: SessionDep, commit: bool = True): if not user.id: raise HTTPException(500) rescuer = Rescuer( user_id=user.id, ) session.add(rescuer) - session.commit() - session.refresh(rescuer) + if commit: + session.commit() + session.refresh(rescuer) @router.post("/create/user/rescuer") @@ -499,8 +515,14 @@ def create_rescuer( session.commit() return {"status": "ok"} except IntegrityError as _: + session.rollback() raise HTTPException(403, "user is already a rescuer") - except Exception as _: + except Exception: + session.rollback() + logger.exception( + "Failed to grant rescuer role", + extra=log_context(current_user.id, "admin_rescuer_grant_failed", user_id), + ) raise HTTPException(500) @@ -517,13 +539,18 @@ def create_admin( session.commit() return {"status": "ok"} except IntegrityError as _: + session.rollback() raise HTTPException(403, "user is already an admin") - except Exception as _: + except Exception: session.rollback() + logger.exception( + "Failed to grant admin role", + extra=log_context(current_user.id, "admin_role_grant_failed", user_id), + ) raise HTTPException(500) -def removeAdmin(user: User, session: SessionDep): +def removeAdmin(user: User, session: SessionDep, commit: bool = True): if not user.id: raise HTTPException(500) if not user.id: @@ -534,14 +561,15 @@ def removeAdmin(user: User, session: SessionDep): # 2. If it exists, delete it if admin: session.delete(admin) - session.commit() + if commit: + session.commit() return {"message": "Admin deleted successfully"} # 3. Handle the case where it doesn't exist raise HTTPException(404, "Admin not found") -def removeRescuer(user: User, session: SessionDep): +def removeRescuer(user: User, session: SessionDep, commit: bool = True): if not user.id: raise HTTPException(500) statement = select(Rescuer).where(Rescuer.user_id == user.id) @@ -550,7 +578,8 @@ def removeRescuer(user: User, session: SessionDep): # 2. If it exists, delete it if rescuer: session.delete(rescuer) - session.commit() + if commit: + session.commit() return {"message": "Rescuer deleted successfully"} # 3. Handle the case where it doesn't exist @@ -567,8 +596,15 @@ def remove_admin( try: removeAdmin(user, session) return {"status": "ok"} - except Exception as _: + except HTTPException: session.rollback() + raise + except Exception: + session.rollback() + logger.exception( + "Failed to remove admin role", + extra=log_context(current_user.id, "admin_role_removal_failed", user_id), + ) raise HTTPException(500) @@ -582,8 +618,15 @@ def remove_rescuer( try: removeRescuer(user, session) return {"status": "ok"} - except Exception as _: + except HTTPException: session.rollback() + raise + except Exception: + session.rollback() + logger.exception( + "Failed to remove rescuer role", + extra=log_context(current_user.id, "admin_rescuer_removal_failed", user_id), + ) raise HTTPException(500) @@ -609,8 +652,12 @@ def create_user( except HTTPException as e: session.rollback() raise e - except Exception as e: + except Exception: session.rollback() + logger.exception( + "Failed to create user through admin console", + extra=log_context(current_user.id, "admin_user_creation_failed"), + ) raise HTTPException(500) @@ -625,41 +672,38 @@ def edit_user( if not user or not user.id: raise HTTPException(404, "user not found") - print("here") update_user_info( - user, UserUpdate(**userData.model_dump(exclude_unset=True)), session + user, UserUpdate(**userData.model_dump(exclude_unset=True)), session, commit=False ) - print("here after") updated_user = get_user_by_ID(session, userData.id) if not updated_user: - print("raise 500") raise HTTPException(500) - print("here after 500") - try: - if userData.is_rescuer: - makeRescuer(updated_user, session) - else: - removeRescuer(updated_user, session) - except: - pass + if userData.is_rescuer is not None: + if userData.is_rescuer and not updated_user.rescuer: + makeRescuer(updated_user, session, commit=False) + elif not userData.is_rescuer and updated_user.rescuer: + removeRescuer(updated_user, session, commit=False) - try: - if userData.is_admin: - makeAdmin(updated_user, session) - else: - removeAdmin(updated_user, session) - except: - pass + if userData.is_admin is not None: + if userData.is_admin and not updated_user.admin: + makeAdmin(updated_user, session, commit=False) + elif not userData.is_admin and updated_user.admin: + removeAdmin(updated_user, session, commit=False) + + session.commit() return {"status": "ok"} except HTTPException as e: session.rollback() raise e - except Exception as e: + except Exception: session.rollback() - print("E", e) + logger.exception( + "Failed to update user", + extra=log_context(current_user.id, "admin_user_update_failed", userData.id), + ) raise HTTPException(500) @@ -677,7 +721,12 @@ def delete_user( session.commit() except HTTPException as e: raise e - except Exception as e: + except Exception: + session.rollback() + logger.exception( + "Failed to delete user", + extra=log_context(_.id, "admin_user_deletion_failed", user_id), + ) raise HTTPException(500) return {"status": "ok"} @@ -715,8 +764,12 @@ def ban_user( session.refresh(ban) except HTTPException as e: raise e - except Exception as e: - print(e) + except Exception: + session.rollback() + logger.exception( + "Failed to ban user", + extra=log_context(_.id, "admin_user_ban_failed", user_id), + ) raise HTTPException(500) return {"status": "ok"} @@ -741,7 +794,12 @@ def unban_user( session.commit() except HTTPException as e: raise e - except Exception as e: + except Exception: + session.rollback() + logger.exception( + "Failed to unban user", + extra=log_context(_.id, "admin_user_unban_failed", user_id), + ) raise HTTPException(500) return {"status": "ok"} diff --git a/server/app/api/peer_connection.py b/server/app/api/peer_connection.py index b2036d9e..466e169a 100644 --- a/server/app/api/peer_connection.py +++ b/server/app/api/peer_connection.py @@ -19,20 +19,25 @@ from app.models.queued import Queue from app.models.signalling import SignalMessage from fastapi import Query, WebSocketDisconnect +from pydantic import ValidationError from app.db_operations.websockets import authenticate_websocket, relay_message, relay_public_message, validate_message_sender, validate_sender, relay_signal, receive_signal_message, WebSocketAuthError from app.db_operations.connection_manager import manager from app.db_operations.activity import set_user_status from app.models.websocketComms import MessageData, PublicMessageData +from app.structured_logging import log_context -logger = logging.getLogger(__name__) +logger = logging.getLogger("app") def _set_status_bg(user_id: UUID, status: str) -> None: try: with Session(engine) as session: set_user_status(session, user_id, status) - except Exception as e: - print(f"[activity] status update failed for {user_id}: {e}") + except Exception: + logger.exception( + "Activity status update failed", + extra=log_context(user_id, "websocket_activity_status_update_failed", metadata={"status": status}), + ) router = APIRouter( prefix='/ws', @@ -129,8 +134,11 @@ def get_queued_messages(user_id: UUID, session: SessionDep, limit: int = 100): try: statement = select(Queue).where(Queue.to == user_id).limit(limit) return session.exec(statement).all() - except Exception as e: - print("EX", e) + except Exception: + logger.exception( + "Failed to fetch queued messages", + extra=log_context(user_id, "websocket_queue_fetch_failed"), + ) return None @@ -149,10 +157,10 @@ def deep_parse_dict(data): if (data.startswith('{') and data.endswith('}')) or (data.startswith('[') and data.endswith(']')): try: data = json.loads(data) - except Exception: + except json.JSONDecodeError: try: data = ast.literal_eval(data) - except Exception: + except (ValueError, SyntaxError): return data else: return data @@ -179,15 +187,23 @@ async def main_web_socket(token: str, websocket: WebSocket, target_id: UUID|None try: user_id = await authenticate_websocket(websocket, token) except WebSocketAuthError: - logger.warning("WebSocket auth rejected: invalid or expired token client=%s", websocket.client) + logger.warning( + "WebSocket auth rejected: invalid or expired token client=%s", + websocket.client, + extra=log_context(None, "websocket_auth_rejected"), + ) return await manager.connect(UUID(user_id), websocket) asyncio.get_event_loop().run_in_executor(None, _set_status_bg, UUID(user_id), "Active") try: await manager.broadcast({"type": "status-update", "user_id": user_id, 'status': "online"}) - except: - pass + except Exception: + logger.warning( + "Failed to broadcast WebSocket online status", + exc_info=True, + extra=log_context(user_id, "websocket_online_status_broadcast_failed"), + ) try: with Session(engine) as session: @@ -211,10 +227,16 @@ async def main_web_socket(token: str, websocket: WebSocket, target_id: UUID|None if message.payload_type == 'seen': session.delete(message) session.commit() - except Exception as e: - print(f"[drain] failed to deliver queued message {message.id}: {e}") - except Exception as e: - print(f"[drain] failed to fetch queued messages for {user_id}: {e}") + except Exception: + logger.exception( + "Failed to deliver queued WebSocket message", + extra=log_context(user_id, "websocket_queue_delivery_failed", message.id), + ) + except Exception: + logger.exception( + "Failed to drain queued WebSocket messages", + extra=log_context(user_id, "websocket_queue_drain_failed"), + ) try: while True: @@ -227,15 +249,23 @@ async def main_web_socket(token: str, websocket: WebSocket, target_id: UUID|None if raw_type == "public-chat": try: payload = PublicMessageData.model_validate(raw_payload) - except Exception: + except ValidationError: + logger.debug( + "Invalid public WebSocket payload", + extra=log_context(user_id, "websocket_invalid_payload", metadata={"type": raw_type}), + ) payload = raw_payload else: try: payload = MessageData.model_validate(raw_payload) - except Exception: + except ValidationError: try: - payload = SignalMessage(**raw_payload) - except Exception: + payload = SignalMessage.model_validate(raw_payload) + except ValidationError: + logger.debug( + "Invalid WebSocket payload", + extra=log_context(user_id, "websocket_invalid_payload", metadata={"type": raw_type}), + ) payload = raw_payload if isinstance(payload, dict) and payload.get("type") == "ping": await manager.send_personal_message(UUID(user_id), {"type": "pong"}) @@ -261,7 +291,11 @@ async def main_web_socket(token: str, websocket: WebSocket, target_id: UUID|None except WebSocketDisconnect: try: await manager.broadcast({"type": "status-update","user_id": user_id, 'status': "offline"}) - except: - pass + except Exception: + logger.warning( + "Failed to broadcast WebSocket offline status", + exc_info=True, + extra=log_context(user_id, "websocket_offline_status_broadcast_failed"), + ) asyncio.get_event_loop().run_in_executor(None, _set_status_bg, UUID(user_id), "Inactive") await manager.disconnect(UUID(user_id)) diff --git a/server/app/api/update_info.py b/server/app/api/update_info.py index eba8d4a9..5b319823 100644 --- a/server/app/api/update_info.py +++ b/server/app/api/update_info.py @@ -1,4 +1,5 @@ from typing import Annotated +import logging from sqlalchemy import or_ from fastapi import Depends, HTTPException from fastapi.routing import APIRouter @@ -9,6 +10,7 @@ from app.models.users import UserUpdate, User from app.db_operations.auth import update_user_info from app.db_operations.auth import SessionDep +from app.structured_logging import log_context from sqlalchemy.exc import IntegrityError from fastapi import status @@ -20,6 +22,7 @@ 404: {'description': 'Not Found'} } ) +logger = logging.getLogger("app") @router.post("/", status_code=status.HTTP_200_OK) @@ -73,8 +76,12 @@ def update_user( raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail=detail) - except Exception as e: + except Exception: session.rollback() + logger.exception( + "Failed to update profile", + extra=log_context(current_user.id, "profile_update_failed", current_user.id), + ) raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="An unexpected error occurred" diff --git a/server/app/db_operations/GPS_manager.py b/server/app/db_operations/GPS_manager.py index 59ef327e..4a1065bd 100644 --- a/server/app/db_operations/GPS_manager.py +++ b/server/app/db_operations/GPS_manager.py @@ -1,6 +1,10 @@ import json +import logging from fastapi import WebSocket from typing import Dict +from app.structured_logging import log_context + +logger = logging.getLogger("app") class GPSManager: def __init__(self): @@ -17,11 +21,18 @@ def disconnect_monitor(self, user_id: str): async def broadcast_to_rescuers(self, message: dict): """Send the GPS packet to every connected Rescuer.""" - for user_id, connection in self.active_monitors.items(): + stale_monitors: list[tuple[str, WebSocket]] = [] + for user_id, connection in list(self.active_monitors.items()): try: await connection.send_json(message) except Exception: - # If a connection is dead, we'll clean it up later or on next fail - pass + logger.info( + "GPS broadcast delivery failed for disconnected monitor", + extra=log_context(user_id, "gps_monitor_disconnected"), + ) + stale_monitors.append((user_id, connection)) + for user_id, connection in stale_monitors: + if self.active_monitors.get(user_id) is connection: + self.disconnect_monitor(user_id) gps_manager = GPSManager() diff --git a/server/app/db_operations/auth.py b/server/app/db_operations/auth.py index e627ff10..5dd5a349 100644 --- a/server/app/db_operations/auth.py +++ b/server/app/db_operations/auth.py @@ -1,16 +1,19 @@ import os +import logging from datetime import datetime, timezone from typing import Annotated, Dict from uuid import UUID from fastapi import Depends, HTTPException, Request from pwdlib import PasswordHash -from sqlalchemy.exc import IntegrityError +from sqlalchemy.exc import IntegrityError, SQLAlchemyError from pwdlib.hashers.argon2 import Argon2Hasher from sqlmodel import SQLModel, Session, create_engine, select, or_ from app.models.users import User, UserCreate from app.models.users import UserUpdate, UserPasswordUpdate +from app.structured_logging import log_context +logger = logging.getLogger("app") # Reduced from recommended() defaults (time_cost=2, memory_cost=65536) @@ -60,7 +63,12 @@ def db_create_user(user: UserCreate, session: SessionDep): user_in_db = get_user_by_ID(session, user.id) if user.id else None except HTTPException: user_in_db = None - except: + except SQLAlchemyError: + session.rollback() + logger.exception( + "Failed to look up existing user during user creation", + extra=log_context(None, "user_creation_lookup_failed"), + ) raise HTTPException(500, "Internal server error.") errors: Dict[str, str] = {} @@ -192,15 +200,18 @@ def authenticate_user( return user -def update_user_info(user: User, new_user_data : UserUpdate, session : SessionDep): +def update_user_info( + user: User, new_user_data: UserUpdate, session: SessionDep, commit: bool = True +): new_user_dump = new_user_data.model_dump(exclude_unset=True) for field, value in new_user_dump.items(): setattr(user, field, value) session.add(user) - session.commit() - session.refresh(user) + if commit: + session.commit() + session.refresh(user) PASSWORD_MIN_LENGTH = 8 diff --git a/server/app/db_operations/connection_manager.py b/server/app/db_operations/connection_manager.py index 55802969..fd2478c6 100644 --- a/server/app/db_operations/connection_manager.py +++ b/server/app/db_operations/connection_manager.py @@ -5,10 +5,11 @@ from uuid import UUID, uuid4 from fastapi import WebSocket from typing import Dict, Optional +from app.structured_logging import log_context import redis.asyncio as aioredis -logger = logging.getLogger(__name__) +logger = logging.getLogger("app") _PRESENCE_KEY = "ws:online_users" # Sorted set; score = expiry epoch (float) _BROADCAST_CHANNEL = "ws:broadcast" @@ -93,6 +94,10 @@ async def send_personal_message(self, target_id: UUID, message: dict) -> None: await ws.send_json(message) return except Exception: + logger.info( + "WebSocket delivery failed for disconnected user", + extra=log_context(target_id, "websocket_user_disconnected"), + ) await self.disconnect(target_id) # Cross-worker: publish so the holding worker delivers it if self._redis: @@ -150,8 +155,11 @@ async def _user_sub_loop( try: data = json.loads(raw["data"]) await websocket.send_json(data) - except Exception as exc: - logger.debug("[ws-sub] delivery failed user=%s: %s", user_id, exc) + except Exception: + logger.info( + "WebSocket subscription delivery failed for disconnected user", + extra=log_context(user_id, "websocket_subscription_disconnected"), + ) break except asyncio.CancelledError: pass @@ -170,12 +178,21 @@ async def _broadcast_loop(self, pubsub: aioredis.client.PubSub) -> None: continue message = envelope["msg"] except Exception: + logger.warning( + "Discarded malformed WebSocket broadcast", + exc_info=True, + extra=log_context(None, "websocket_broadcast_malformed"), + ) continue stale: list[UUID] = [] for uid, ws in list(self._local.items()): try: await ws.send_json(message) except Exception: + logger.info( + "WebSocket broadcast delivery failed for disconnected user", + extra=log_context(uid, "websocket_broadcast_disconnected"), + ) stale.append(uid) for uid in stale: await self.disconnect(uid) diff --git a/server/app/db_operations/token.py b/server/app/db_operations/token.py index b2aab2be..594b40a6 100644 --- a/server/app/db_operations/token.py +++ b/server/app/db_operations/token.py @@ -1,6 +1,7 @@ #!/usr/bin/env python3 import os import threading +import logging import redis as _redis_module from uuid import UUID, uuid4 from pydantic import BaseModel @@ -20,6 +21,9 @@ from app.models.users import UserCreate from app.models.jti import BlacklistedToken from app.db_operations.auth import get_user, get_user_by_ID +from app.structured_logging import log_context + +logger = logging.getLogger("app") _REDIS_URL = os.getenv("REDIS_URL", "redis://localhost:6379") @@ -27,7 +31,12 @@ _redis: _redis_module.Redis = _redis_module.from_url(_REDIS_URL, decode_responses=True) _redis.ping() _REDIS_AVAILABLE = True -except Exception: +except _redis_module.RedisError: + logger.warning( + "Redis blacklist cache is unavailable; using database checks", + exc_info=True, + extra=log_context(None, "token_blacklist_cache_unavailable"), + ) _redis = None # type: ignore[assignment] _REDIS_AVAILABLE = False @@ -90,7 +99,7 @@ def get_user_id_from_header(request: Request): token = auth_header.split(" ")[1] payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM]) return payload.get("sub") # or payload.get("user_id") - except Exception: + except (IndexError, PyJWTError): return None def verify_token(token: str): diff --git a/server/app/structured_logging.py b/server/app/structured_logging.py new file mode 100644 index 00000000..ac470f8f --- /dev/null +++ b/server/app/structured_logging.py @@ -0,0 +1,16 @@ +from typing import Any +from uuid import UUID + + +def log_context( + user_id: UUID | str | None, + action: str, + entity_id: UUID | str | None = None, + metadata: dict[str, Any] | None = None, +) -> dict[str, Any]: + return { + "user_id": str(user_id) if user_id else "ANONYMOUS", + "action": action, + "entity_id": str(entity_id) if entity_id else None, + "metadata_json": metadata or {}, + } diff --git a/server/app/tests/test_admin_me.py b/server/app/tests/test_admin_me.py index 27e4b162..566cf538 100644 --- a/server/app/tests/test_admin_me.py +++ b/server/app/tests/test_admin_me.py @@ -1,8 +1,12 @@ from fastapi.testclient import TestClient +from fastapi import HTTPException +import pytest from sqlmodel import Session, select +from app.api import admin from app.models.admin import Admin -from app.models.users import User +from app.models.rescuer import Rescuer +from app.models.users import User, UserUpdateThroughAdmin def _login_as_admin(client: TestClient, session: Session, username: str, password: str) -> str: @@ -74,3 +78,48 @@ def test_logout_without_refresh_token_cookie_returns_401_not_500(client: TestCli ) assert response.status_code == 401 + + +def test_admin_edit_rolls_back_profile_when_role_change_fails(session: Session, monkeypatch): + user = session.exec(select(User).where(User.username == "test")).one() + user_id = user.id + original_username = user.username + update = UserUpdateThroughAdmin( + id=user_id, + username="updated-admin-user", + is_admin=True, + is_rescuer=False, + ) + + def fail_role_grant(*args, **kwargs): + raise RuntimeError("role write failed") + + monkeypatch.setattr(admin, "makeAdmin", fail_role_grant) + + with pytest.raises(HTTPException) as exc_info: + admin.edit_user(user, update, session) + + assert exc_info.value.status_code == 500 + session.expire_all() + persisted_user = session.exec(select(User).where(User.id == user_id)).one() + assert persisted_user.username == original_username + + +def test_admin_edit_preserves_roles_when_role_fields_are_omitted(session: Session): + user = session.exec(select(User).where(User.username == "test")).one() + session.add_all([Admin(user_id=user.id), Rescuer(user_id=user.id)]) + session.commit() + session.expire_all() + user = session.exec(select(User).where(User.username == "test")).one() + + result = admin.edit_user( + user, + UserUpdateThroughAdmin(id=user.id, username="renamed-admin-user"), + session, + ) + + assert result == {"status": "ok"} + session.expire_all() + persisted_user = session.exec(select(User).where(User.id == user.id)).one() + assert persisted_user.admin is not None + assert persisted_user.rescuer is not None diff --git a/server/app/tests/test_gps.py b/server/app/tests/test_gps.py index 3288afd8..6e7bd37c 100644 --- a/server/app/tests/test_gps.py +++ b/server/app/tests/test_gps.py @@ -1,4 +1,6 @@ import pytest +import asyncio +import logging from datetime import datetime, timedelta, timezone from fastapi.testclient import TestClient from fastapi import APIRouter, WebSocket, WebSocketDisconnect, Depends @@ -10,6 +12,43 @@ import json from app.tests.test_db_utils import get_auth_headers +from app.db_operations.GPS_manager import GPSManager + + +class FailingWebSocket: + async def send_json(self, message: dict) -> None: + raise RuntimeError("connection closed") + + +def test_gps_broadcast_logs_and_removes_failed_monitor(caplog): + manager = GPSManager() + manager.active_monitors["rescuer-1"] = FailingWebSocket() + caplog.set_level(logging.INFO, logger="app") + + asyncio.run(manager.broadcast_to_rescuers({"lat": 1, "lng": 2})) + + assert manager.active_monitors == {} + record = next(record for record in caplog.records if "GPS broadcast delivery failed" in record.message) + assert record.user_id == "rescuer-1" + assert record.action == "gps_monitor_disconnected" + assert record.entity_id is None + assert record.metadata_json == {} + + +def test_gps_broadcast_keeps_monitor_that_reconnects_during_delivery(): + manager = GPSManager() + replacement = object() + + class ReconnectingWebSocket: + async def send_json(self, message: dict) -> None: + manager.active_monitors["rescuer-1"] = replacement + raise RuntimeError("connection closed") + + manager.active_monitors["rescuer-1"] = ReconnectingWebSocket() + + asyncio.run(manager.broadcast_to_rescuers({"lat": 1, "lng": 2})) + + assert manager.active_monitors["rescuer-1"] is replacement def test_stream_gps_location_success(client: TestClient): @@ -184,4 +223,3 @@ def test_gps_broadcast_to_rescuer(client, session, test_user_instance, test_resc # After exiting the block, the rescuer is disconnected # and the 'raise' inside your endpoint is handled by the TestClient - diff --git a/server/app/tests/test_websocket_pool.py b/server/app/tests/test_websocket_pool.py index 39489a14..d7f3eef3 100644 --- a/server/app/tests/test_websocket_pool.py +++ b/server/app/tests/test_websocket_pool.py @@ -10,7 +10,10 @@ The fix: open short-lived sessions per DB operation so an idle peer holds zero connections. """ +import asyncio +import logging import time +from uuid import uuid4 import pytest from sqlmodel import Session, SQLModel, create_engine @@ -21,6 +24,28 @@ from app.main import app from app.db_operations.auth import get_session from app.db_operations.token import create_access_token +from app.db_operations.connection_manager import ConnectionManager + + +class FailingWebSocket: + async def send_json(self, message: dict) -> None: + raise RuntimeError("connection closed") + + +def test_personal_delivery_failure_is_logged_and_disconnects(caplog): + manager = ConnectionManager() + user_id = uuid4() + manager._local[user_id] = FailingWebSocket() + caplog.set_level(logging.INFO, logger="app") + + asyncio.run(manager.send_personal_message(user_id, {"type": "ping"})) + + assert user_id not in manager._local + record = next(record for record in caplog.records if "WebSocket delivery failed" in record.message) + assert record.user_id == str(user_id) + assert record.action == "websocket_user_disconnected" + assert record.entity_id is None + assert record.metadata_json == {} @pytest.fixture(name="pool_engine") From 1b0e93801a14748dd011fb3908dec31ec636788e Mon Sep 17 00:00:00 2001 From: Adamskiee Date: Sun, 16 Aug 2026 11:16:57 +0800 Subject: [PATCH 2/2] fix(server-api): make role writes atomic --- docs/api/conventions.md | 1 + server/app/api/admin.py | 13 +++-- server/app/api/peer_connection.py | 2 +- server/app/db_operations/auth.py | 17 +++++-- .../app/db_operations/connection_manager.py | 14 +++--- server/app/tests/test_admin_me.py | 47 ++++++++++++++++++- server/app/tests/test_websocket_pool.py | 44 +++++++++++++++++ 7 files changed, 121 insertions(+), 17 deletions(-) diff --git a/docs/api/conventions.md b/docs/api/conventions.md index 62bdedc8..e7c65a38 100644 --- a/docs/api/conventions.md +++ b/docs/api/conventions.md @@ -95,6 +95,7 @@ set (every decorated route in `server/app/api/`): | `POST /auth/token` | 5/minute | | `POST /auth/` | 3/minute | | `POST /auth/refresh` | 10/minute | +| `POST /api/admin/refresh` | 10/minute | | `POST /auth/reauthenticate` | 5/minute | | `POST /auth/change-password` | 3/minute | | `POST /auth/forgot-password/otp/send` | 3/minute | diff --git a/server/app/api/admin.py b/server/app/api/admin.py index a0e96c1f..ea07b266 100644 --- a/server/app/api/admin.py +++ b/server/app/api/admin.py @@ -637,16 +637,19 @@ def create_user( session: SessionDep, ): try: - user = db_create_user(userData, session) + user = db_create_user(userData, session, commit=False) if not user.id: raise HTTPException(500) if userData.is_rescuer: - makeRescuer(user, session) + makeRescuer(user, session, commit=False) if userData.is_admin: - makeAdmin(user, session) + makeAdmin(user, session, commit=False) + + session.commit() + session.refresh(user) return user except HTTPException as e: @@ -787,7 +790,9 @@ def unban_user( now = datetime.now(timezone.utc).replace(tzinfo=None) nowreal = datetime.now(timezone.utc) statement = ( - update(BannedUser).where(BannedUser.until > now).values(until=nowreal) + update(BannedUser) + .where(BannedUser.user_id == user_id, BannedUser.until > now) + .values(until=nowreal) ) session.exec(statement) diff --git a/server/app/api/peer_connection.py b/server/app/api/peer_connection.py index 466e169a..5c90fdd8 100644 --- a/server/app/api/peer_connection.py +++ b/server/app/api/peer_connection.py @@ -298,4 +298,4 @@ async def main_web_socket(token: str, websocket: WebSocket, target_id: UUID|None extra=log_context(user_id, "websocket_offline_status_broadcast_failed"), ) asyncio.get_event_loop().run_in_executor(None, _set_status_bg, UUID(user_id), "Inactive") - await manager.disconnect(UUID(user_id)) + await manager.disconnect(UUID(user_id), websocket=websocket) diff --git a/server/app/db_operations/auth.py b/server/app/db_operations/auth.py index 5dd5a349..ecddd9f3 100644 --- a/server/app/db_operations/auth.py +++ b/server/app/db_operations/auth.py @@ -58,7 +58,7 @@ def verify_password(plain_password : str, hashed__password : str): return password_hash.verify(plain_password, hashed__password) -def db_create_user(user: UserCreate, session: SessionDep): +def db_create_user(user: UserCreate, session: SessionDep, commit: bool = True): try: user_in_db = get_user_by_ID(session, user.id) if user.id else None except HTTPException: @@ -105,11 +105,15 @@ def db_create_user(user: UserCreate, session: SessionDep): session.add(db_user) try: - session.commit() + if commit: + session.commit() + else: + session.flush() except IntegrityError: session.rollback() raise HTTPException(status_code=400, detail={"username": "Username or contact already registered"}) - session.refresh(db_user) + if commit: + session.refresh(db_user) return db_user elif user_in_db and user_in_db.guest: # modify existing user @@ -129,8 +133,11 @@ def db_create_user(user: UserCreate, session: SessionDep): session.add(user_in_db) # delete guest record session.delete(user_in_db.guest) - session.commit() - session.refresh(user_in_db) + if commit: + session.commit() + session.refresh(user_in_db) + else: + session.flush() return user_in_db # TODO: all guest accounts are disabled from getting a token in any way shape or form diff --git a/server/app/db_operations/connection_manager.py b/server/app/db_operations/connection_manager.py index fd2478c6..5e43e7b3 100644 --- a/server/app/db_operations/connection_manager.py +++ b/server/app/db_operations/connection_manager.py @@ -75,7 +75,9 @@ async def connect(self, user_id: UUID, websocket: WebSocket) -> None: self._heartbeat_loop(user_id), name=f"ws-hb-{user_id}" ) - async def disconnect(self, user_id: UUID) -> None: + async def disconnect(self, user_id: UUID, websocket: WebSocket | None = None) -> None: + if websocket is not None and self._local.get(user_id) is not websocket: + return self._local.pop(user_id, None) for tasks in (self._sub_tasks, self._hb_tasks): task = tasks.pop(user_id, None) @@ -98,7 +100,7 @@ async def send_personal_message(self, target_id: UUID, message: dict) -> None: "WebSocket delivery failed for disconnected user", extra=log_context(target_id, "websocket_user_disconnected"), ) - await self.disconnect(target_id) + await self.disconnect(target_id, websocket=ws) # Cross-worker: publish so the holding worker delivers it if self._redis: await self._redis.publish( @@ -184,7 +186,7 @@ async def _broadcast_loop(self, pubsub: aioredis.client.PubSub) -> None: extra=log_context(None, "websocket_broadcast_malformed"), ) continue - stale: list[UUID] = [] + stale: list[tuple[UUID, WebSocket]] = [] for uid, ws in list(self._local.items()): try: await ws.send_json(message) @@ -193,9 +195,9 @@ async def _broadcast_loop(self, pubsub: aioredis.client.PubSub) -> None: "WebSocket broadcast delivery failed for disconnected user", extra=log_context(uid, "websocket_broadcast_disconnected"), ) - stale.append(uid) - for uid in stale: - await self.disconnect(uid) + stale.append((uid, ws)) + for uid, websocket in stale: + await self.disconnect(uid, websocket=websocket) except asyncio.CancelledError: pass finally: diff --git a/server/app/tests/test_admin_me.py b/server/app/tests/test_admin_me.py index 566cf538..432f7295 100644 --- a/server/app/tests/test_admin_me.py +++ b/server/app/tests/test_admin_me.py @@ -1,12 +1,14 @@ from fastapi.testclient import TestClient from fastapi import HTTPException +from datetime import datetime, timedelta, timezone import pytest from sqlmodel import Session, select from app.api import admin from app.models.admin import Admin +from app.models.banned_user import BannedUser from app.models.rescuer import Rescuer -from app.models.users import User, UserUpdateThroughAdmin +from app.models.users import User, UserCreateThroughAdmin, UserUpdateThroughAdmin def _login_as_admin(client: TestClient, session: Session, username: str, password: str) -> str: @@ -123,3 +125,46 @@ def test_admin_edit_preserves_roles_when_role_fields_are_omitted(session: Sessio persisted_user = session.exec(select(User).where(User.id == user.id)).one() assert persisted_user.admin is not None assert persisted_user.rescuer is not None + + +def test_admin_create_rolls_back_user_and_role_when_role_grant_fails(session: Session, monkeypatch): + actor = session.exec(select(User).where(User.username == "test")).one() + user_data = UserCreateThroughAdmin( + username="atomic-admin-user", + first_name="Atomic", + last_name="Admin", + password="StrongPass123", + is_rescuer=True, + is_admin=True, + ) + + def fail_admin_grant(*args, **kwargs): + raise RuntimeError("role write failed") + + monkeypatch.setattr(admin, "makeAdmin", fail_admin_grant) + + with pytest.raises(HTTPException) as exc_info: + admin.create_user(actor, user_data, session) + + assert exc_info.value.status_code == 500 + assert session.exec(select(User).where(User.username == "atomic-admin-user")).first() is None + + +def test_unban_only_updates_requested_user(session: Session): + actor = session.exec(select(User).where(User.username == "test")).one() + target = session.exec(select(User).where(User.username == "Tony Stark")).one() + other = session.exec(select(User).where(User.username == "Steve Rogers")).one() + expires_at = datetime.now(timezone.utc).replace(tzinfo=None) + timedelta(days=1) + session.add_all([ + BannedUser(user_id=target.id, until=expires_at), + BannedUser(user_id=other.id, until=expires_at), + ]) + session.commit() + + admin.unban_user(actor, target.id, session) + + session.expire_all() + target_ban = session.exec(select(BannedUser).where(BannedUser.user_id == target.id)).one() + other_ban = session.exec(select(BannedUser).where(BannedUser.user_id == other.id)).one() + assert target_ban.until < expires_at + assert other_ban.until == expires_at diff --git a/server/app/tests/test_websocket_pool.py b/server/app/tests/test_websocket_pool.py index d7f3eef3..ff2f04f4 100644 --- a/server/app/tests/test_websocket_pool.py +++ b/server/app/tests/test_websocket_pool.py @@ -48,6 +48,50 @@ def test_personal_delivery_failure_is_logged_and_disconnects(caplog): assert record.metadata_json == {} +def test_personal_delivery_failure_keeps_reconnected_socket(): + manager = ConnectionManager() + user_id = uuid4() + replacement = object() + + class ReconnectingWebSocket: + async def send_json(self, message: dict) -> None: + manager._local[user_id] = replacement + raise RuntimeError("connection closed") + + manager._local[user_id] = ReconnectingWebSocket() + + asyncio.run(manager.send_personal_message(user_id, {"type": "ping"})) + + assert manager._local[user_id] is replacement + + +def test_broadcast_delivery_failure_keeps_reconnected_socket(): + manager = ConnectionManager() + user_id = uuid4() + replacement = object() + + class ReconnectingWebSocket: + async def send_json(self, message: dict) -> None: + manager._local[user_id] = replacement + raise RuntimeError("connection closed") + + class OneMessagePubSub: + async def listen(self): + yield {"type": "message", "data": '{"_wid": "other", "msg": {"type": "ping"}}'} + + async def unsubscribe(self, channel: str) -> None: + return None + + async def aclose(self) -> None: + return None + + manager._local[user_id] = ReconnectingWebSocket() + + asyncio.run(manager._broadcast_loop(OneMessagePubSub())) + + assert manager._local[user_id] is replacement + + @pytest.fixture(name="pool_engine") def pool_engine_fixture(tmp_path, monkeypatch): """A real QueuePool-backed SQLite engine so we can inspect checked-out