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, _principal_idm_context, _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.core.idm import OrganizationFunctionAssignmentRef 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() def test_acting_assignment_requires_exact_session_selection(self) -> None: class Directory: def organization_function_assignments_for_account( self, account_id: str, *, tenant_id: str | None = None, effective_at=None, ): del account_id, effective_at return ( OrganizationFunctionAssignmentRef( id="direct-1", tenant_id=str(tenant_id), identity_id="identity-1", account_id="account-1", function_id="function-direct", organization_unit_id="unit-1", ), OrganizationFunctionAssignmentRef( id="acting-1", tenant_id=str(tenant_id), identity_id="identity-1", account_id="account-1", function_id="function-acting", organization_unit_id="unit-1", source="acting_for", acting_for_account_id="represented-1", ), ) engine = create_engine("sqlite:///:memory:") create_scope_tables(engine) AccessBase.metadata.create_all(bind=engine) 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="actor@example.test", normalized_email="actor@example.test", ) user = User( id="user-1", tenant_id=tenant.id, account_id=account.id, email=account.email, ) auth_session = AuthSession( id="session-1", tenant_id=tenant.id, user_id=user.id, account_id=account.id, token_hash="token-hash", expires_at=utc_now() + timedelta(hours=1), ) session.add_all((tenant, account, user, auth_session)) session.flush() ordinary, _ = _principal_idm_context( session, user=user, account=account, tenant_id=tenant.id, idm_directory=Directory(), # type: ignore[arg-type] organization_directory=None, ) self.assertEqual([item.id for item in ordinary], ["direct-1"]) auth_session.acting_assignment_id = "acting-1" auth_session.acting_for_account_id = "represented-1" selected, _ = _principal_idm_context( session, user=user, account=account, tenant_id=tenant.id, idm_directory=Directory(), # type: ignore[arg-type] organization_directory=None, auth_session=auth_session, ) self.assertEqual( [item.id for item in selected], ["direct-1", "acting-1"], ) auth_session.acting_for_account_id = "wrong-account" mismatched, _ = _principal_idm_context( session, user=user, account=account, tenant_id=tenant.id, idm_directory=Directory(), # type: ignore[arg-type] organization_directory=None, auth_session=auth_session, ) self.assertEqual([item.id for item in mismatched], ["direct-1"]) finally: AccessBase.metadata.drop_all(bind=engine) scope_registry.metadata.drop_all(bind=engine) engine.dispose() if __name__ == "__main__": unittest.main()