feat(cloud-auth): verify scoped bearer tokens
This commit is contained in:
@@ -0,0 +1,88 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from hashlib import sha256
|
||||
from hmac import compare_digest
|
||||
from typing import Protocol, runtime_checkable
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Principal:
|
||||
id: str = "anonymous"
|
||||
scopes: frozenset[str] = field(default_factory=frozenset)
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class AuthProvider(Protocol):
|
||||
def authenticate(self, request: object) -> Principal | None: ...
|
||||
|
||||
|
||||
class NullAuthProvider:
|
||||
"""Explicit insecure-development provider with unrestricted scope."""
|
||||
|
||||
def authenticate(self, request: object) -> Principal | None:
|
||||
return Principal(scopes=frozenset({"*"}))
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class BearerCredential:
|
||||
principal_id: str
|
||||
token: str = field(repr=False)
|
||||
scopes: frozenset[str] = field(default_factory=frozenset)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if not self.principal_id.strip():
|
||||
raise ValueError("bearer credential principal_id must not be empty")
|
||||
if not self.token:
|
||||
raise ValueError("bearer credential token must not be empty")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _StoredBearerCredential:
|
||||
principal: Principal
|
||||
token_digest: bytes = field(repr=False)
|
||||
|
||||
|
||||
class ConfiguredBearerAuthProvider:
|
||||
"""Authenticate configured bearer tokens without exposing credential text."""
|
||||
|
||||
def __init__(self, credentials: list[BearerCredential]) -> None:
|
||||
self._credentials = tuple(
|
||||
_StoredBearerCredential(
|
||||
principal=Principal(
|
||||
id=credential.principal_id,
|
||||
scopes=frozenset(credential.scopes),
|
||||
),
|
||||
token_digest=_token_digest(credential.token),
|
||||
)
|
||||
for credential in credentials
|
||||
)
|
||||
|
||||
def authenticate(self, request: object) -> Principal | None:
|
||||
token = _extract_bearer_token(request)
|
||||
if token is None:
|
||||
return None
|
||||
|
||||
candidate_digest = _token_digest(token)
|
||||
matched_principal: Principal | None = None
|
||||
for credential in self._credentials:
|
||||
if compare_digest(candidate_digest, credential.token_digest):
|
||||
matched_principal = credential.principal
|
||||
return matched_principal
|
||||
|
||||
|
||||
def _extract_bearer_token(request: object) -> str | None:
|
||||
headers = getattr(request, "headers", None)
|
||||
if headers is None:
|
||||
return None
|
||||
authorization = headers.get("authorization")
|
||||
if not isinstance(authorization, str):
|
||||
return None
|
||||
scheme, separator, token = authorization.partition(" ")
|
||||
if not separator or scheme.lower() != "bearer" or not token:
|
||||
return None
|
||||
return token
|
||||
|
||||
|
||||
def _token_digest(token: str) -> bytes:
|
||||
return sha256(token.encode("utf-8")).digest()
|
||||
@@ -10,9 +10,9 @@ authentication can be added later without changing route signatures.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Protocol, runtime_checkable
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from cloud.auth import AuthProvider, NullAuthProvider, Principal
|
||||
from cloud.sdk.models import (
|
||||
DeviceResponse,
|
||||
ErrorResponse,
|
||||
@@ -23,34 +23,12 @@ from cloud.sdk.models import (
|
||||
TaskSubmissionRequest,
|
||||
TaskSubmissionResponse,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from fastapi.responses import JSONResponse
|
||||
from fastapi import APIRouter, HTTPException, Request, status
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from cloud.plugins import PluginRegistry
|
||||
from cloud.pool import DevicePool
|
||||
from cloud.scheduler import TaskConstraints, TaskScheduler
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Principal:
|
||||
"""An authenticated principal. ``anonymous`` for the NullAuthProvider."""
|
||||
|
||||
id: str = "anonymous"
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class AuthProvider(Protocol):
|
||||
"""Returns a Principal if the request is allowed, None to reject."""
|
||||
|
||||
def authenticate(self, request: object) -> Principal | None: ...
|
||||
|
||||
|
||||
class NullAuthProvider:
|
||||
"""Default auth provider: every caller is anonymous-and-allowed."""
|
||||
|
||||
def authenticate(self, request: object) -> Principal | None:
|
||||
return Principal()
|
||||
from cloud.scheduler import TaskScheduler
|
||||
|
||||
|
||||
def create_cloud_router(
|
||||
|
||||
Reference in New Issue
Block a user