162 lines
6.2 KiB
Python
162 lines
6.2 KiB
Python
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()
|