from __future__ import annotations import unittest from datetime import timedelta from typing import Iterable from fastapi import HTTPException from sqlalchemy import create_engine from sqlalchemy.orm import sessionmaker from starlette.requests import Request from govoplan_access.backend.auth.dependencies import ( _extract_token, _requires_csrf, _resolve_legacy_principal_context, _resolve_legacy_principal_ref, ) from govoplan_access.backend.auth.principal_cache import principal_summary_cache from govoplan_access.backend.auth.tokens import hash_secret from govoplan_access.backend.db.base import AccessBase from govoplan_access.backend.db.models import Account, AuthSession, Role, User, UserRoleAssignment from govoplan_core.core.change_sequence import ChangeSequenceEntry, ChangeSequenceRetentionFloor from govoplan_core.core.principal_cache import invalidate_auth_principals from govoplan_core.db.base import Base from govoplan_core.security.time import utc_now from govoplan_core.settings import settings from govoplan_core.tenancy.scope import Tenant, create_scope_tables, scope_registry def request_for(*, method: str = "GET", headers: Iterable[tuple[str, str]] = ()) -> Request: return Request( { "type": "http", "method": method, "path": "/", "headers": [(key.lower().encode("latin-1"), value.encode("latin-1")) for key, value in headers], } ) class AuthDependencyTests(unittest.TestCase): def tearDown(self) -> None: principal_summary_cache.clear() def test_extract_token_prefers_explicit_api_key(self) -> None: request = request_for(headers=[("authorization", "Bearer session-token")]) self.assertEqual(_extract_token(request, "Bearer session-token", "api-key-token"), ("api-key-token", "api_key")) def test_extract_token_supports_bearer_and_cookie_sources(self) -> None: self.assertEqual(_extract_token(request_for(), "Bearer session-token", None), ("session-token", "bearer")) cookie_request = request_for(headers=[("cookie", f"{settings.auth_session_cookie_name}=cookie-token")]) self.assertEqual(_extract_token(cookie_request, None, None), ("cookie-token", "cookie")) def test_requires_csrf_only_for_mutating_methods(self) -> None: self.assertFalse(_requires_csrf(request_for(method="GET"))) self.assertTrue(_requires_csrf(request_for(method="POST"))) def test_legacy_principal_resolver_rejects_missing_token_before_db_lookup(self) -> None: with self.assertRaises(HTTPException) as raised: _resolve_legacy_principal_ref(request_for(), None, authorization=None, x_api_key=None) # type: ignore[arg-type] self.assertEqual(raised.exception.status_code, 401) self.assertEqual(raised.exception.detail, "Missing API key or session token") def test_permission_revision_invalidates_cached_principal(self) -> None: engine = create_engine("sqlite:///:memory:") create_scope_tables(engine) AccessBase.metadata.create_all(bind=engine) Base.metadata.create_all( bind=engine, tables=[ ChangeSequenceEntry.__table__, ChangeSequenceRetentionFloor.__table__, ], ) SessionLocal = sessionmaker(bind=engine) try: with SessionLocal() as session: tenant = Tenant(id="tenant-1", slug="tenant-1", name="Tenant 1") account = Account( id="account-1", email="owner@example.test", normalized_email="owner@example.test", ) user = User( id="user-1", tenant_id=tenant.id, account_id=account.id, email=account.email, ) role = Role( id="role-1", tenant_id=tenant.id, slug="reader", name="Reader", permissions=["files:file:read"], ) assignment = UserRoleAssignment( tenant_id=tenant.id, user_id=user.id, role_id=role.id, ) token = "ms_test-session-token" auth_session = AuthSession( id="session-1", tenant_id=tenant.id, user_id=user.id, account_id=account.id, token_hash=hash_secret(token), expires_at=utc_now() + timedelta(hours=1), ) session.add_all( [tenant, account, user, role, assignment, auth_session] ) session.commit() request = request_for() first = _resolve_legacy_principal_context( request, session, authorization=f"Bearer {token}", x_api_key=None, ) self.assertIn("files:file:read", first.principal.scopes) role.permissions = ["files:file:write"] session.add(role) invalidate_auth_principals( session, tenant_id=tenant.id, source_module="access", resource_type="role", resource_id=role.id, ) session.commit() second = _resolve_legacy_principal_context( request, session, authorization=f"Bearer {token}", x_api_key=None, ) self.assertNotIn("files:file:read", second.principal.scopes) self.assertIn("files:file:write", second.principal.scopes) finally: AccessBase.metadata.drop_all(bind=engine) scope_registry.metadata.drop_all(bind=engine) Base.metadata.drop_all( bind=engine, tables=[ ChangeSequenceEntry.__table__, ChangeSequenceRetentionFloor.__table__, ], ) engine.dispose() if __name__ == "__main__": unittest.main()