diff --git a/tests/conftest.py b/tests/conftest.py index 6b78bc7..b2719d2 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -26,6 +26,14 @@ ) from main import app from datetime import datetime, UTC, timedelta +from utils.core.rate_limit import clear_all_rate_limiters + + +@pytest.fixture(autouse=True) +def reset_rate_limiters() -> Generator[None, None, None]: + clear_all_rate_limiters() + yield + clear_all_rate_limiters() # Define a custom exception for test setup errors diff --git a/tests/routers/core/test_account.py b/tests/routers/core/test_account.py index 382c40a..63d263e 100644 --- a/tests/routers/core/test_account.py +++ b/tests/routers/core/test_account.py @@ -1,4 +1,3 @@ -import pytest from fastapi.testclient import TestClient from starlette.datastructures import URLPath from sqlmodel import Session, select @@ -28,28 +27,13 @@ get_password_hash, ) from utils.core.rate_limit import ( - login_ip_limiter, + forgot_password_email_limiter, + forgot_password_ip_limiter, login_email_limiter, + login_ip_limiter, register_ip_limiter, - forgot_password_ip_limiter, - forgot_password_email_limiter, ) - -@pytest.fixture(autouse=True) -def _reset_rate_limiters(): - """Reset all rate limiter state between tests to avoid cross-test pollution.""" - yield - for limiter in ( - login_ip_limiter, - login_email_limiter, - register_ip_limiter, - forgot_password_ip_limiter, - forgot_password_email_limiter, - ): - limiter._attempts.clear() - - # --- API Endpoint Tests --- diff --git a/tests/test_htmx.py b/tests/test_htmx.py index 050558c..1fe9332 100644 --- a/tests/test_htmx.py +++ b/tests/test_htmx.py @@ -8,34 +8,15 @@ - Non-HTMX paths remain unchanged (303 RedirectResponse or full-page error). """ -import pytest from starlette.requests import Request from fastapi.templating import Jinja2Templates from tests.conftest import htmx_headers from utils.core.htmx import is_htmx_request, toast_response, append_toast from utils.core.rate_limit import ( - login_ip_limiter, - login_email_limiter, - register_ip_limiter, forgot_password_ip_limiter, - forgot_password_email_limiter, + login_ip_limiter, ) - -@pytest.fixture(autouse=True) -def _reset_rate_limiters(): - """Reset all rate limiter state between tests.""" - yield - for limiter in ( - login_ip_limiter, - login_email_limiter, - register_ip_limiter, - forgot_password_ip_limiter, - forgot_password_email_limiter, - ): - limiter._attempts.clear() - - # --------------------------------------------------------------------------- # 1.3 — is_htmx_request helper # --------------------------------------------------------------------------- diff --git a/utils/core/organizations.py b/utils/core/organizations.py index 27e2ed6..27788a4 100644 --- a/utils/core/organizations.py +++ b/utils/core/organizations.py @@ -1,5 +1,7 @@ +from typing import Any, cast + from sqlmodel import Session, select -from sqlalchemy.orm import selectinload +from sqlalchemy.orm import InstrumentedAttribute, selectinload from utils.core.models import Organization, Role, User, Invitation @@ -21,13 +23,15 @@ def load_org_for_members_partial( select(Organization) .where(Organization.id == organization_id) .options( - selectinload(Organization.roles) - .selectinload(Role.users) - .selectinload(User.account), - selectinload(Organization.roles) - .selectinload(Role.users) - .selectinload(User.roles), - selectinload(Organization.roles).selectinload(Role.permissions), + selectinload(cast(InstrumentedAttribute[Any], Organization.roles)) + .selectinload(cast(InstrumentedAttribute[Any], Role.users)) + .selectinload(cast(InstrumentedAttribute[Any], User.account)), + selectinload(cast(InstrumentedAttribute[Any], Organization.roles)) + .selectinload(cast(InstrumentedAttribute[Any], Role.users)) + .selectinload(cast(InstrumentedAttribute[Any], User.roles)), + selectinload( + cast(InstrumentedAttribute[Any], Organization.roles) + ).selectinload(cast(InstrumentedAttribute[Any], Role.permissions)), ) ).first() user_permissions = _user_permissions_for_org(user, organization_id) @@ -43,8 +47,12 @@ def load_org_for_roles_partial( select(Organization) .where(Organization.id == organization_id) .options( - selectinload(Organization.roles).selectinload(Role.users), - selectinload(Organization.roles).selectinload(Role.permissions), + selectinload( + cast(InstrumentedAttribute[Any], Organization.roles) + ).selectinload(cast(InstrumentedAttribute[Any], Role.users)), + selectinload( + cast(InstrumentedAttribute[Any], Organization.roles) + ).selectinload(cast(InstrumentedAttribute[Any], Role.permissions)), ) ).first() user_permissions = _user_permissions_for_org(user, organization_id) diff --git a/utils/core/rate_limit.py b/utils/core/rate_limit.py index bf2cdda..18a3d35 100644 --- a/utils/core/rate_limit.py +++ b/utils/core/rate_limit.py @@ -154,6 +154,20 @@ def _int_env(name: str, default: int) -> int: window_seconds=_int_env("FORGOT_PASSWORD_EMAIL_WINDOW_SECONDS", 60), ) +_ALL_LIMITERS = ( + login_ip_limiter, + login_email_limiter, + register_ip_limiter, + forgot_password_ip_limiter, + forgot_password_email_limiter, +) + + +def clear_all_rate_limiters() -> None: + """Clear all in-memory rate limiter state.""" + for limiter in _ALL_LIMITERS: + limiter._attempts.clear() + # --- Dependency helpers ---