diff --git a/main.py b/main.py index c89d72f..51e72ae 100644 --- a/main.py +++ b/main.py @@ -393,9 +393,7 @@ async def general_exception_handler(request: Request, exc: Exception): @app.get("/") -async def read_home( - request: Request, _: None = Depends(require_unauthenticated_client) -): +def read_home(request: Request, _: None = Depends(require_unauthenticated_client)): return templates.TemplateResponse(request, "index.html", {"user": None}) diff --git a/routers/core/account.py b/routers/core/account.py index d9aaf65..4a2d2c3 100644 --- a/routers/core/account.py +++ b/routers/core/account.py @@ -188,7 +188,7 @@ def logout( @router.get("/login") -async def read_login( +def read_login( request: Request, _: None = Depends(require_unauthenticated_unless_invitation_warning), invitation_token: Optional[str] = Query(None), @@ -215,7 +215,7 @@ async def read_login( @router.get("/register") -async def read_register( +def read_register( request: Request, _: None = Depends(require_unauthenticated_unless_invitation_warning), email: Optional[EmailStr] = Query(None), @@ -246,7 +246,7 @@ async def read_register( @router.get("/forgot_password") -async def read_forgot_password( +def read_forgot_password( request: Request, _: None = Depends(require_unauthenticated_client), show_form: Optional[str] = "true", @@ -262,7 +262,7 @@ async def read_forgot_password( @router.get("/reset_password") -async def read_reset_password( +def read_reset_password( request: Request, email: str, token: str, @@ -291,7 +291,7 @@ async def read_reset_password( @router.post("/delete", response_class=RedirectResponse) -async def delete_account( +def delete_account( account: Account = Depends(get_verified_account), session: Session = Depends(get_session), ): @@ -314,7 +314,7 @@ async def delete_account( @router.post("/register", response_class=RedirectResponse) -async def register( +def register( request: Request, _ip_check: None = Depends(check_register_ip_rate_limit), name: str = Form( @@ -473,7 +473,7 @@ async def register( @router.post("/login", response_class=RedirectResponse) -async def login( +def login( request: Request, _ip_check: None = Depends(check_login_ip_rate_limit), _email_check: EmailStr = Depends(check_login_email_rate_limit), @@ -588,7 +588,7 @@ async def login( # Updated refresh_token endpoint @router.post("/refresh", response_class=RedirectResponse) -async def refresh_token( +def refresh_token( tokens: tuple[Optional[str], Optional[str]] = Depends(oauth2_scheme_cookie), session: Session = Depends(get_session), ) -> RedirectResponse: @@ -666,7 +666,7 @@ async def refresh_token( @router.post("/forgot_password") -async def forgot_password( +def forgot_password( background_tasks: BackgroundTasks, request: Request, _ip_check: None = Depends(check_forgot_password_ip_rate_limit), @@ -703,7 +703,7 @@ async def forgot_password( @router.post("/reset_password") -async def reset_password( +def reset_password( request: Request, email: EmailStr = Form(..., title="Email", description="Account email address"), token: str = Form( @@ -765,7 +765,7 @@ async def reset_password( @router.get("/recover") -async def recover_account_confirm( +def recover_account_confirm( request: Request, token: str = Query(...), session: Session = Depends(get_session), @@ -784,7 +784,7 @@ async def recover_account_confirm( @router.post("/recover") -async def recover_account( +def recover_account( token: str = Form(...), session: Session = Depends(get_session), ): @@ -852,7 +852,7 @@ async def recover_account( @router.post("/emails/add") -async def add_email( +def add_email( request: Request, new_email: EmailStr = Form( ..., title="New email", description="New email address to add" @@ -905,7 +905,7 @@ async def add_email( @router.get("/emails/verify") -async def verify_email( +def verify_email( token: str, session: Session = Depends(get_session), ): @@ -958,7 +958,7 @@ async def verify_email( @router.post("/emails/promote") -async def promote_email( +def promote_email( request: Request, email_id: int = Form( ..., title="Email ID", description="ID of the email to promote" @@ -1047,7 +1047,7 @@ async def promote_email( @router.post("/emails/remove") -async def remove_email( +def remove_email( request: Request, email_id: int = Form( ..., title="Email ID", description="ID of the email to remove" diff --git a/routers/core/dashboard.py b/routers/core/dashboard.py index 82e8a95..0dc282f 100644 --- a/routers/core/dashboard.py +++ b/routers/core/dashboard.py @@ -15,7 +15,7 @@ @router.get("/") -async def read_dashboard( +def read_dashboard( request: Request, user: User = Depends(get_user_with_relations), session: Session = Depends(get_session), @@ -77,7 +77,7 @@ async def read_dashboard( @router.post("/select-organization/{org_id}") -async def select_organization( +def select_organization( request: Request, org_id: int, user: User = Depends(get_user_with_relations), diff --git a/routers/core/invitation.py b/routers/core/invitation.py index 473ec42..8bdb745 100644 --- a/routers/core/invitation.py +++ b/routers/core/invitation.py @@ -108,7 +108,7 @@ def _members_table_response( @router.post("/", name="create_invitation") -async def create_invitation( +def create_invitation( request: Request, current_user: User = Depends(get_authenticated_user), session: Session = Depends(get_session), @@ -205,7 +205,7 @@ async def create_invitation( @router.post("/resend", name="resend_invitation", response_class=RedirectResponse) -async def resend_invitation( +def resend_invitation( request: Request, current_user: User = Depends(get_authenticated_user), session: Session = Depends(get_session), @@ -277,7 +277,7 @@ async def resend_invitation( @router.post("/delete", name="delete_invitation", response_class=RedirectResponse) -async def delete_invitation( +def delete_invitation( request: Request, current_user: User = Depends(get_authenticated_user), session: Session = Depends(get_session), @@ -318,7 +318,7 @@ async def delete_invitation( @router.get("/accept", name="accept_invitation") -async def accept_invitation( +def accept_invitation( token: str = Query(...), current_user: Optional[User] = Depends(get_optional_user), session: Session = Depends(get_session), diff --git a/routers/core/organization.py b/routers/core/organization.py index 4c79312..89e7e90 100644 --- a/routers/core/organization.py +++ b/routers/core/organization.py @@ -36,7 +36,7 @@ @router.get("/{org_id}") -async def read_organization( +def read_organization( org_id: int, request: Request, user: User = Depends(get_user_with_relations), diff --git a/routers/core/static_pages.py b/routers/core/static_pages.py index ca859fb..56e81d2 100644 --- a/routers/core/static_pages.py +++ b/routers/core/static_pages.py @@ -16,7 +16,7 @@ @router.get("/{page_name}", name="read_static_page") -async def read_static_page( +def read_static_page( page_name: str, request: Request, user: Optional[User] = Depends(get_optional_user) ): """ diff --git a/routers/core/user.py b/routers/core/user.py index 983e562..7dff819 100644 --- a/routers/core/user.py +++ b/routers/core/user.py @@ -1,5 +1,6 @@ from fastapi import APIRouter, Depends, Form, UploadFile, File, Request, HTTPException from fastapi.responses import RedirectResponse, Response +from starlette.concurrency import run_in_threadpool from sqlmodel import Session, select, col from typing import Optional, List from fastapi.templating import Jinja2Templates @@ -51,7 +52,7 @@ @router.get("/profile") -async def read_profile( +def read_profile( request: Request, user: User = Depends(get_user_with_relations), session: Session = Depends(get_session), @@ -82,7 +83,7 @@ async def read_profile( @router.get("/edit-form") -async def edit_profile_form( +def edit_profile_form( request: Request, user: User = Depends(get_authenticated_user), ): @@ -104,7 +105,7 @@ async def edit_profile_form( @router.get("/profile-display") -async def profile_display( +def profile_display( request: Request, user: User = Depends(get_authenticated_user), ): @@ -131,17 +132,17 @@ async def update_profile( ): avatar_changed = bool(avatar_file and avatar_file.filename) - # Handle avatar update + # Async upload read stays on the event loop. CPU-bound image work is + # offloaded; Session/ORM mutations stay here so the Session is not passed + # into the thread pool. if avatar_changed: assert avatar_file is not None reject_oversized_content_length( request.headers.get("content-length"), MAX_AVATAR_UPLOAD_BYTES ) avatar_data = await read_upload_with_size_limit(avatar_file, MAX_FILE_SIZE) - avatar_content_type = avatar_file.content_type - - processed_image, content_type = validate_and_process_image( - avatar_data, avatar_content_type + processed_image, content_type = await run_in_threadpool( + validate_and_process_image, avatar_data, avatar_file.content_type ) if user.avatar: user.avatar.avatar_data = processed_image @@ -154,9 +155,7 @@ async def update_profile( avatar_content_type=content_type, ) - # Update user details user.name = name - session.commit() session.refresh(user) @@ -185,7 +184,7 @@ async def update_profile( @router.post("/communication-preferences", response_class=RedirectResponse) -async def update_communication_preferences( +def update_communication_preferences( request: Request, comm_opt_in: Optional[str] = Form(None), comm_updates: Optional[str] = Form(None), @@ -210,7 +209,7 @@ async def update_communication_preferences( @router.get("/avatar") -async def get_avatar(user: User = Depends(get_authenticated_user)): +def get_avatar(user: User = Depends(get_authenticated_user)): """Serve avatar image from database""" if not user.avatar: raise DataIntegrityError(resource="User avatar") diff --git a/tests/utils/test_dependencies.py b/tests/utils/test_dependencies.py index 3ed865a..07eb153 100644 --- a/tests/utils/test_dependencies.py +++ b/tests/utils/test_dependencies.py @@ -1,5 +1,10 @@ +import asyncio +from concurrent.futures import ThreadPoolExecutor from unittest.mock import MagicMock, patch from datetime import datetime, timedelta, UTC +from typing import Any, Coroutine, TypeVar +from starlette.concurrency import run_in_threadpool +from starlette.requests import Request from utils.core.models import ( Account, AccountRecoveryToken, @@ -14,6 +19,8 @@ get_authenticated_account, validate_token_and_get_user, get_user_from_tokens, + get_user_from_request, + _get_user_from_request_sync, get_authenticated_user, get_optional_user, get_account_from_reset_token, @@ -684,3 +691,92 @@ def test_get_account_from_recovery_token_invalid() -> None: account, token = get_account_from_recovery_token("nonexistent", session) assert account is None assert token is None + + +def _request_with_auth_cookies( + access_token: str | None = "access_token", + refresh_token: str | None = "refresh_token", +) -> Request: + cookie_parts: list[str] = [] + if access_token is not None: + cookie_parts.append(f"access_token={access_token}") + if refresh_token is not None: + cookie_parts.append(f"refresh_token={refresh_token}") + headers: list[tuple[bytes, bytes]] = [] + if cookie_parts: + headers.append((b"cookie", "; ".join(cookie_parts).encode())) + return Request( + { + "type": "http", + "headers": headers, + "method": "GET", + "path": "/error", + "query_string": b"", + } + ) + + +T = TypeVar("T") + + +def _run_async(coro: Coroutine[Any, Any, T]) -> T: + """Drive an async helper from sync tests. + + The full suite may already have an event loop (e.g. after Playwright), so + asyncio.run() on the main thread can fail. Always run in a fresh thread. + """ + + def _runner() -> T: + return asyncio.run(coro) + + with ThreadPoolExecutor(max_workers=1) as pool: + return pool.submit(_runner).result() + + +def test_get_user_from_request_resolves_user_via_threadpool() -> None: + """ + Exception handlers await get_user_from_request directly (not via Depends). + Cookies are read on the event loop; sync DB work is offloaded through + run_in_threadpool and must return the user resolved from those tokens. + """ + mock_user = User(id=1, name="Test User") + mock_user.avatar = None + mock_session = MagicMock() + request = _request_with_auth_cookies() + + with ( + patch( + "utils.core.dependencies.run_in_threadpool", + wraps=run_in_threadpool, + ) as mock_threadpool, + patch("utils.core.dependencies.get_engine", return_value=MagicMock()), + patch("utils.core.dependencies.Session") as mock_session_cls, + patch("utils.core.dependencies.get_user_from_tokens") as mock_get_user, + ): + mock_session_cls.return_value.__enter__.return_value = mock_session + mock_session_cls.return_value.__exit__.return_value = None + mock_get_user.return_value = (mock_user, None, None) + + user = _run_async(get_user_from_request(request)) + + assert user is mock_user + mock_threadpool.assert_called_once_with( + _get_user_from_request_sync, "access_token", "refresh_token" + ) + mock_get_user.assert_called_once_with( + ("access_token", "refresh_token"), mock_session + ) + + # No auth cookies → no user (still via the same async/threadpool path) + bare_request = _request_with_auth_cookies(access_token=None, refresh_token=None) + with ( + patch("utils.core.dependencies.get_engine", return_value=MagicMock()), + patch("utils.core.dependencies.Session") as mock_session_cls, + patch("utils.core.dependencies.get_user_from_tokens") as mock_get_user, + ): + mock_session_cls.return_value.__enter__.return_value = mock_session + mock_session_cls.return_value.__exit__.return_value = None + mock_get_user.return_value = (None, None, None) + + assert _run_async(get_user_from_request(bare_request)) is None + mock_get_user.assert_called_once_with((None, None), mock_session) diff --git a/utils/core/dependencies.py b/utils/core/dependencies.py index 7d86ce2..607a384 100644 --- a/utils/core/dependencies.py +++ b/utils/core/dependencies.py @@ -1,5 +1,6 @@ import logging from fastapi import Depends, Form, Query, Request +from starlette.concurrency import run_in_threadpool from pydantic import EmailStr from sqlmodel import Session, select from sqlalchemy.orm import selectinload @@ -427,9 +428,22 @@ async def get_user_from_request(request: Request) -> Optional[User]: """ Helper function to get user from request cookies in exception handlers. Exception handlers can't use Depends(), so we manually extract tokens and get the user. + + Cookie reads stay on the event loop; sync DB/session work runs in the thread + pool. This is called directly (not via Depends()) from async exception + handlers, so without offloading it would block the loop while querying. """ access_token = request.cookies.get("access_token") refresh_token = request.cookies.get("refresh_token") + return await run_in_threadpool( + _get_user_from_request_sync, access_token, refresh_token + ) + + +def _get_user_from_request_sync( + access_token: Optional[str], + refresh_token: Optional[str], +) -> Optional[User]: tokens = (access_token, refresh_token) # Get a database session