feat(cloud): add edge host enrollment
This commit is contained in:
@@ -4,7 +4,10 @@ from dataclasses import dataclass, field
|
||||
from hashlib import sha256
|
||||
from hmac import compare_digest
|
||||
from collections.abc import Iterable
|
||||
from typing import Protocol, runtime_checkable
|
||||
from typing import TYPE_CHECKING, Protocol, runtime_checkable
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from cloud.repository import CloudRepository
|
||||
|
||||
|
||||
TASKS_SUBMIT_SCOPE = "tasks:submit"
|
||||
@@ -78,6 +81,24 @@ class _StoredBearerCredential:
|
||||
token_digest: bytes = field(repr=False)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class EnrollmentCredential:
|
||||
principal_id: str
|
||||
token: str = field(repr=False)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if not self.principal_id.strip():
|
||||
raise ValueError("enrollment credential principal_id must not be empty")
|
||||
if not self.token:
|
||||
raise ValueError("enrollment credential token must not be empty")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class EnrollmentPrincipal:
|
||||
id: str
|
||||
token_digest: str = field(repr=False)
|
||||
|
||||
|
||||
class ConfiguredBearerAuthProvider:
|
||||
"""Authenticate configured bearer tokens without exposing credential text."""
|
||||
|
||||
@@ -107,6 +128,57 @@ class ConfiguredBearerAuthProvider:
|
||||
return matched_principal
|
||||
|
||||
|
||||
class ConfiguredEnrollmentTokenProvider:
|
||||
def __init__(self, credentials: Iterable[EnrollmentCredential]) -> None:
|
||||
self._credentials = tuple(
|
||||
EnrollmentPrincipal(
|
||||
id=credential.principal_id,
|
||||
token_digest=digest_token(credential.token),
|
||||
)
|
||||
for credential in credentials
|
||||
)
|
||||
|
||||
def authenticate(self, request: object) -> EnrollmentPrincipal | None:
|
||||
candidate = bearer_token_digest(request)
|
||||
if candidate is None:
|
||||
return None
|
||||
candidate_bytes = bytes.fromhex(candidate)
|
||||
matched: EnrollmentPrincipal | None = None
|
||||
for credential in self._credentials:
|
||||
if compare_digest(
|
||||
candidate_bytes,
|
||||
bytes.fromhex(credential.token_digest),
|
||||
):
|
||||
matched = credential
|
||||
return matched
|
||||
|
||||
|
||||
class RepositoryHostAuthProvider:
|
||||
def __init__(self, repository: CloudRepository) -> None:
|
||||
self.repository = repository
|
||||
|
||||
def authenticate(self, request: object) -> Principal | None:
|
||||
credential_digest = bearer_token_digest(request)
|
||||
if credential_digest is None:
|
||||
return None
|
||||
host_id = self.repository.authenticate_enrolled_host(credential_digest)
|
||||
if host_id is None:
|
||||
return None
|
||||
return Principal(id=f"enrolled-host:{host_id}", host_id=host_id)
|
||||
|
||||
|
||||
class ChainedAuthProvider:
|
||||
def __init__(self, providers: Iterable[AuthProvider]) -> None:
|
||||
self.providers = tuple(providers)
|
||||
|
||||
def authenticate(self, request: object) -> Principal | None:
|
||||
for provider in self.providers:
|
||||
principal = provider.authenticate(request)
|
||||
if principal is not None:
|
||||
return principal
|
||||
return None
|
||||
|
||||
|
||||
def create_auth_provider(
|
||||
credentials: Iterable[BearerCredential],
|
||||
*,
|
||||
@@ -133,5 +205,14 @@ def _extract_bearer_token(request: object) -> str | None:
|
||||
return token
|
||||
|
||||
|
||||
def bearer_token_digest(request: object) -> str | None:
|
||||
token = _extract_bearer_token(request)
|
||||
return digest_token(token) if token is not None else None
|
||||
|
||||
|
||||
def digest_token(token: str) -> str:
|
||||
return _token_digest(token).hex()
|
||||
|
||||
|
||||
def _token_digest(token: str) -> bytes:
|
||||
return sha256(token.encode("utf-8")).digest()
|
||||
|
||||
Reference in New Issue
Block a user