feat: cache immutable auth principal summaries

This commit is contained in:
2026-07-29 19:24:00 +02:00
parent 6b0dd8beab
commit b6c2c89adf
4 changed files with 446 additions and 3 deletions

View File

@@ -1,13 +1,30 @@
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_ref
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:
@@ -22,6 +39,9 @@ def request_for(*, method: str = "GET", headers: Iterable[tuple[str, str]] = ())
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")])
@@ -44,6 +64,98 @@ class AuthDependencyTests(unittest.TestCase):
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()