Scaffold FastAPI + MySQL + Keycloak service with devcontainer
Sets up the project skeleton: - FastAPI app factory with lifespan, request-id middleware, and RFC 9457 problem+json error handlers - Async SQLAlchemy 2.0 over MySQL (asyncmy), with a constraint naming convention in place before the first migration and async Alembic - Keycloak as a pure resource server: OIDC discovery, cached JWKS with rotation-aware refresh, and require_roles dependencies - Devcontainer running MySQL 8.4 and Keycloak 26.7 as compose siblings, with the realm (clients, roles, test users) imported on first boot - Test suite covering the endpoints plus the token validator itself, exercised against a locally generated RSA keypair - uv packaging, ruff, mypy --strict, pre-commit, Gitea CI, prod Dockerfile Two Keycloak-in-containers traps are handled explicitly and documented in the README: the issuer/internal-URL split (the browser sees localhost:8080, the API sees keycloak:8080) and the audience mapper that stops Keycloak issuing tokens with aud=account. The devices resource is a placeholder proving the routing -> auth -> ORM -> migration path end to end; replace it with the real domain. Co-Authored-By: Claude Opus 5 <[email protected]>
This commit is contained in:
commit
0526d34e42
52 files changed
+4433
No files matched your search
@@ -0,0 +1 @@
|
||||
__version__ = "0.1.0"
|
||||
Whitespace-only changes.
@@ -0,0 +1,9 @@
|
||||
"""Aggregate router for API v1."""
|
||||
|
||||
from fastapi import APIRouter
|
||||
|
||||
from v2x_server.api.v1 import devices, identity
|
||||
|
||||
api_router = APIRouter()
|
||||
api_router.include_router(identity.router)
|
||||
api_router.include_router(devices.router)
|
||||
Whitespace-only changes.
@@ -0,0 +1,97 @@
|
||||
"""Devices resource: the worked example of an authenticated, role-gated CRUD API.
|
||||
|
||||
Read requires `viewer`; writes require `operator`; delete requires `admin`.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, status
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from v2x_server.auth.deps import require_roles
|
||||
from v2x_server.core.errors import ConflictError, NotFoundError
|
||||
from v2x_server.db.session import get_session
|
||||
from v2x_server.models.device import Device, DeviceStatus
|
||||
from v2x_server.repositories.device import DeviceRepository
|
||||
from v2x_server.schemas.device import DeviceCreate, DevicePage, DeviceRead, DeviceUpdate
|
||||
|
||||
router = APIRouter(prefix="/devices", tags=["devices"])
|
||||
|
||||
SessionDep = Annotated[AsyncSession, Depends(get_session)]
|
||||
|
||||
|
||||
def get_repository(session: SessionDep) -> DeviceRepository:
|
||||
return DeviceRepository(session)
|
||||
|
||||
|
||||
RepoDep = Annotated[DeviceRepository, Depends(get_repository)]
|
||||
|
||||
# `admin` is included everywhere so an administrator is never locked out of a
|
||||
# resource they are meant to administer.
|
||||
ReadAccess = Depends(require_roles("viewer", "operator", "admin"))
|
||||
WriteAccess = Depends(require_roles("operator", "admin"))
|
||||
DeleteAccess = Depends(require_roles("admin"))
|
||||
|
||||
|
||||
async def _get_or_404(repo: DeviceRepository, device_id: int) -> Device:
|
||||
device = await repo.get(device_id)
|
||||
if device is None:
|
||||
raise NotFoundError(f"Device {device_id} does not exist")
|
||||
return device
|
||||
|
||||
|
||||
@router.get("", response_model=DevicePage, dependencies=[ReadAccess])
|
||||
async def list_devices(
|
||||
repo: RepoDep,
|
||||
limit: Annotated[int, Query(ge=1, le=200)] = 50,
|
||||
offset: Annotated[int, Query(ge=0)] = 0,
|
||||
status_filter: Annotated[DeviceStatus | None, Query(alias="status")] = None,
|
||||
) -> DevicePage:
|
||||
items, total = await repo.list(limit=limit, offset=offset, status=status_filter)
|
||||
return DevicePage(
|
||||
items=[DeviceRead.model_validate(item) for item in items],
|
||||
total=total,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/{device_id}", response_model=DeviceRead, dependencies=[ReadAccess])
|
||||
async def get_device(device_id: int, repo: RepoDep) -> Device:
|
||||
return await _get_or_404(repo, device_id)
|
||||
|
||||
|
||||
@router.post(
|
||||
"",
|
||||
response_model=DeviceRead,
|
||||
status_code=status.HTTP_201_CREATED,
|
||||
dependencies=[WriteAccess],
|
||||
)
|
||||
async def create_device(payload: DeviceCreate, repo: RepoDep) -> Device:
|
||||
if await repo.get_by_serial(payload.serial):
|
||||
raise ConflictError(f"A device with serial {payload.serial!r} already exists")
|
||||
return await repo.create(payload)
|
||||
|
||||
|
||||
@router.patch("/{device_id}", response_model=DeviceRead, dependencies=[WriteAccess])
|
||||
async def update_device(device_id: int, payload: DeviceUpdate, repo: RepoDep) -> Device:
|
||||
device = await _get_or_404(repo, device_id)
|
||||
|
||||
if payload.serial is not None and payload.serial != device.serial:
|
||||
existing = await repo.get_by_serial(payload.serial)
|
||||
if existing is not None and existing.id != device.id:
|
||||
raise ConflictError(f"A device with serial {payload.serial!r} already exists")
|
||||
|
||||
return await repo.update(device, payload)
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/{device_id}",
|
||||
status_code=status.HTTP_204_NO_CONTENT,
|
||||
dependencies=[DeleteAccess],
|
||||
)
|
||||
async def delete_device(device_id: int, repo: RepoDep) -> None:
|
||||
device = await _get_or_404(repo, device_id)
|
||||
await repo.delete(device)
|
||||
@@ -0,0 +1,54 @@
|
||||
"""Liveness and readiness probes.
|
||||
|
||||
Split deliberately: `/healthz` must not depend on MySQL or Keycloak, or a
|
||||
transient outage in either would make an orchestrator kill healthy pods.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Response, status
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.exc import SQLAlchemyError
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from v2x_server.auth.deps import get_oidc_provider
|
||||
from v2x_server.auth.keycloak import OIDCProvider
|
||||
from v2x_server.db.session import get_session
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(tags=["health"])
|
||||
|
||||
|
||||
@router.get("/healthz", summary="Liveness probe")
|
||||
async def healthz() -> dict[str, str]:
|
||||
"""Is the process up? No dependencies are consulted."""
|
||||
return {"status": "ok"}
|
||||
|
||||
|
||||
@router.get("/readyz", summary="Readiness probe")
|
||||
async def readyz(
|
||||
response: Response,
|
||||
session: Annotated[AsyncSession, Depends(get_session)],
|
||||
provider: Annotated[OIDCProvider, Depends(get_oidc_provider)],
|
||||
) -> dict[str, Any]:
|
||||
"""Can we actually serve traffic -- database reachable, signing keys loadable?"""
|
||||
checks: dict[str, str] = {}
|
||||
|
||||
try:
|
||||
await session.execute(text("SELECT 1"))
|
||||
checks["database"] = "ok"
|
||||
except SQLAlchemyError:
|
||||
logger.warning("Readiness: database check failed", exc_info=True)
|
||||
checks["database"] = "error"
|
||||
|
||||
checks["keycloak"] = "ok" if await provider.healthy() else "error"
|
||||
|
||||
ready = all(value == "ok" for value in checks.values())
|
||||
if not ready:
|
||||
response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE
|
||||
|
||||
return {"status": "ready" if ready else "not ready", "checks": checks}
|
||||
@@ -0,0 +1,22 @@
|
||||
"""Endpoints that describe the caller to themselves."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter
|
||||
|
||||
from v2x_server.auth.deps import CurrentPrincipal
|
||||
|
||||
router = APIRouter(prefix="/me", tags=["identity"])
|
||||
|
||||
|
||||
@router.get("", summary="Echo the caller's identity and roles")
|
||||
async def whoami(principal: CurrentPrincipal) -> dict[str, object]:
|
||||
"""The quickest way to confirm which roles a token actually carries."""
|
||||
return {
|
||||
"subject": principal.subject,
|
||||
"username": principal.username,
|
||||
"email": principal.email,
|
||||
"realm_roles": sorted(principal.realm_roles),
|
||||
"client_roles": sorted(principal.client_roles),
|
||||
"scopes": sorted(principal.scopes),
|
||||
}
|
||||
Whitespace-only changes.
@@ -0,0 +1,102 @@
|
||||
"""FastAPI dependencies for authentication and role-based authorization."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from collections.abc import Callable, Coroutine
|
||||
from typing import Annotated, Any, Literal
|
||||
|
||||
from fastapi import Depends, HTTPException, Request, Security, status
|
||||
from fastapi.security import OAuth2AuthorizationCodeBearer
|
||||
|
||||
from v2x_server.auth.keycloak import OIDCProvider, TokenError
|
||||
from v2x_server.auth.principal import Principal
|
||||
from v2x_server.core.config import Settings, get_settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_settings = get_settings()
|
||||
|
||||
# Drives the Authorize button in /docs. These URLs are the ones the *browser*
|
||||
# uses, hence the public realm URL rather than the internal one.
|
||||
oauth2_scheme = OAuth2AuthorizationCodeBearer(
|
||||
authorizationUrl=f"{_settings.public_realm_url}/protocol/openid-connect/auth",
|
||||
tokenUrl=f"{_settings.public_realm_url}/protocol/openid-connect/token",
|
||||
refreshUrl=f"{_settings.public_realm_url}/protocol/openid-connect/token",
|
||||
scopes={"openid": "OpenID Connect", "profile": "Profile", "email": "Email"},
|
||||
auto_error=False,
|
||||
)
|
||||
|
||||
|
||||
def _unauthorized(detail: str, *, error: str = "invalid_token") -> HTTPException:
|
||||
return HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=detail,
|
||||
headers={"WWW-Authenticate": f'Bearer error="{error}"'},
|
||||
)
|
||||
|
||||
|
||||
def get_oidc_provider(request: Request) -> OIDCProvider:
|
||||
provider = getattr(request.app.state, "oidc_provider", None)
|
||||
if provider is None:
|
||||
raise RuntimeError("OIDC provider not initialised; is the app lifespan running?")
|
||||
return provider # type: ignore[no-any-return]
|
||||
|
||||
|
||||
async def get_current_principal(
|
||||
token: Annotated[str | None, Security(oauth2_scheme)],
|
||||
provider: Annotated[OIDCProvider, Depends(get_oidc_provider)],
|
||||
settings: Annotated[Settings, Depends(get_settings)],
|
||||
) -> Principal:
|
||||
"""Validate the Bearer token and build the caller's `Principal`."""
|
||||
if not token:
|
||||
raise _unauthorized("Missing bearer token", error="invalid_request")
|
||||
|
||||
try:
|
||||
claims = await provider.decode(token)
|
||||
except TokenError as exc:
|
||||
# Log the specific reason, tell the client only that it failed.
|
||||
logger.info("Rejected token: %s", exc)
|
||||
raise _unauthorized("Invalid or expired token") from exc
|
||||
|
||||
try:
|
||||
return Principal.from_claims(claims, client_id=settings.keycloak_audience)
|
||||
except (KeyError, ValueError) as exc:
|
||||
logger.warning("Token validated but claims were unusable: %s", exc)
|
||||
raise _unauthorized("Token is missing required claims") from exc
|
||||
|
||||
|
||||
CurrentPrincipal = Annotated[Principal, Depends(get_current_principal)]
|
||||
|
||||
|
||||
def require_roles(
|
||||
*roles: str,
|
||||
mode: Literal["any", "all"] = "any",
|
||||
) -> Callable[[Principal], Coroutine[Any, Any, Principal]]:
|
||||
"""Dependency factory enforcing role membership.
|
||||
|
||||
Authentication failures are 401 (handled upstream); an authenticated caller
|
||||
who simply lacks the role is 403 -- retrying with the same token won't help.
|
||||
|
||||
@router.post("/", dependencies=[Depends(require_roles("operator"))])
|
||||
"""
|
||||
|
||||
async def _dependency(principal: CurrentPrincipal) -> Principal:
|
||||
granted = (
|
||||
principal.has_all_roles(*roles) if mode == "all" else principal.has_any_role(*roles)
|
||||
)
|
||||
if not granted:
|
||||
logger.info(
|
||||
"Denied %s (roles=%s); requires %s of %s",
|
||||
principal.username or principal.subject,
|
||||
sorted(principal.roles),
|
||||
mode,
|
||||
sorted(roles),
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=f"Requires {mode} of the following roles: {', '.join(sorted(roles))}",
|
||||
)
|
||||
return principal
|
||||
|
||||
return _dependency
|
||||
@@ -0,0 +1,184 @@
|
||||
"""OIDC discovery and JWKS handling for a Keycloak-issued access token.
|
||||
|
||||
Deliberately does not use PyJWT's `PyJWKClient`: it fetches over blocking
|
||||
`urllib`, which stalls the event loop on every cache miss. `httpx` plus
|
||||
`PyJWKSet` is barely more code and stays async throughout.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
import jwt
|
||||
from jwt import PyJWK, PyJWKSet
|
||||
|
||||
from v2x_server.core.config import Settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Explicit allowlist. Never derive the algorithm from the token's own header --
|
||||
# that is what enables `alg: none` and RS256->HS256 confusion attacks.
|
||||
ALLOWED_ALGORITHMS = ["RS256"]
|
||||
|
||||
REQUIRED_CLAIMS = ["exp", "iat", "iss", "sub", "aud"]
|
||||
|
||||
|
||||
class TokenError(Exception):
|
||||
"""Token could not be validated. The message is for logs, not for clients."""
|
||||
|
||||
|
||||
class OIDCProvider:
|
||||
"""Caches the realm's discovery document and signing keys."""
|
||||
|
||||
def __init__(self, settings: Settings, client: httpx.AsyncClient) -> None:
|
||||
self._settings = settings
|
||||
self._client = client
|
||||
# Two locks, not one: fetching the JWKS needs the discovery document, so
|
||||
# a single lock would deadlock (asyncio.Lock is not reentrant).
|
||||
self._metadata_lock = asyncio.Lock()
|
||||
self._jwks_lock = asyncio.Lock()
|
||||
|
||||
self._metadata: dict[str, Any] | None = None
|
||||
self._jwks: PyJWKSet | None = None
|
||||
self._jwks_fetched_at: float = 0.0
|
||||
|
||||
# -- discovery --------------------------------------------------------
|
||||
|
||||
async def metadata(self) -> dict[str, Any]:
|
||||
if self._metadata is None:
|
||||
async with self._metadata_lock:
|
||||
if self._metadata is None:
|
||||
self._metadata = await self._fetch_metadata()
|
||||
return self._metadata
|
||||
|
||||
async def _fetch_metadata(self) -> dict[str, Any]:
|
||||
url = self._settings.discovery_url
|
||||
logger.info("Fetching OIDC discovery document from %s", url)
|
||||
response = await self._client.get(url)
|
||||
response.raise_for_status()
|
||||
metadata: dict[str, Any] = response.json()
|
||||
|
||||
# The discovery document's `issuer` is the *frontend* URL (what tokens
|
||||
# will carry), while we fetched over the backchannel. A mismatch against
|
||||
# the configured issuer means KC_HOSTNAME and KEYCLOAK_ISSUER disagree,
|
||||
# and every token would later fail validation -- so say so loudly now.
|
||||
discovered = str(metadata.get("issuer", "")).rstrip("/")
|
||||
if discovered != self._settings.keycloak_issuer:
|
||||
logger.warning(
|
||||
"Configured issuer %r does not match the issuer advertised by Keycloak %r; "
|
||||
"tokens will be rejected. Check KEYCLOAK_ISSUER against KC_HOSTNAME.",
|
||||
self._settings.keycloak_issuer,
|
||||
discovered,
|
||||
)
|
||||
return metadata
|
||||
|
||||
async def jwks_uri(self) -> str:
|
||||
metadata = await self.metadata()
|
||||
uri = metadata.get("jwks_uri")
|
||||
if not isinstance(uri, str):
|
||||
raise TokenError("Discovery document contains no jwks_uri")
|
||||
return uri
|
||||
|
||||
# -- signing keys -----------------------------------------------------
|
||||
|
||||
async def _fetch_jwks(self) -> PyJWKSet:
|
||||
uri = await self.jwks_uri()
|
||||
response = await self._client.get(uri)
|
||||
response.raise_for_status()
|
||||
jwks = PyJWKSet.from_dict(response.json())
|
||||
self._jwks = jwks
|
||||
self._jwks_fetched_at = time.monotonic()
|
||||
logger.info("Loaded %d signing key(s) from %s", len(jwks.keys), uri)
|
||||
return jwks
|
||||
|
||||
async def get_signing_key(self, kid: str) -> PyJWK:
|
||||
"""Resolve a key id, refreshing once if it is unknown.
|
||||
|
||||
Keycloak rotates realm keys without notice; a cached JWKS that predates
|
||||
a rotation would otherwise reject every new token until restart.
|
||||
"""
|
||||
now = time.monotonic()
|
||||
expired = now - self._jwks_fetched_at > self._settings.keycloak_jwks_ttl_seconds
|
||||
|
||||
if self._jwks is None or expired:
|
||||
async with self._jwks_lock:
|
||||
if self._jwks is None or (
|
||||
time.monotonic() - self._jwks_fetched_at
|
||||
> self._settings.keycloak_jwks_ttl_seconds
|
||||
):
|
||||
await self._fetch_jwks()
|
||||
|
||||
key = self._find_key(kid)
|
||||
if key is not None:
|
||||
return key
|
||||
|
||||
# Unknown kid: refresh, but not more often than the floor allows, so a
|
||||
# flood of junk tokens cannot be amplified into load on Keycloak.
|
||||
async with self._jwks_lock:
|
||||
key = self._find_key(kid)
|
||||
if key is not None:
|
||||
return key
|
||||
since_refresh = time.monotonic() - self._jwks_fetched_at
|
||||
if since_refresh < self._settings.keycloak_jwks_min_refresh_seconds:
|
||||
raise TokenError(f"Unknown signing key {kid!r} (refresh rate-limited)")
|
||||
await self._fetch_jwks()
|
||||
|
||||
key = self._find_key(kid)
|
||||
if key is None:
|
||||
raise TokenError(f"Unknown signing key {kid!r}")
|
||||
return key
|
||||
|
||||
def _find_key(self, kid: str) -> PyJWK | None:
|
||||
if self._jwks is None:
|
||||
return None
|
||||
for key in self._jwks.keys:
|
||||
if key.key_id == kid:
|
||||
return key
|
||||
return None
|
||||
|
||||
# -- validation -------------------------------------------------------
|
||||
|
||||
async def decode(self, token: str) -> dict[str, Any]:
|
||||
"""Verify signature and claims, returning the payload.
|
||||
|
||||
Raises `TokenError` for every failure mode; the caller turns that into
|
||||
a generic 401 so we never leak *why* a token was rejected.
|
||||
"""
|
||||
try:
|
||||
header = jwt.get_unverified_header(token)
|
||||
except jwt.PyJWTError as exc:
|
||||
raise TokenError(f"Malformed token header: {exc}") from exc
|
||||
|
||||
kid = header.get("kid")
|
||||
if not kid:
|
||||
raise TokenError("Token header has no kid")
|
||||
|
||||
signing_key = await self.get_signing_key(kid)
|
||||
|
||||
try:
|
||||
payload: dict[str, Any] = jwt.decode(
|
||||
token,
|
||||
key=signing_key.key,
|
||||
algorithms=ALLOWED_ALGORITHMS,
|
||||
issuer=self._settings.keycloak_issuer,
|
||||
audience=self._settings.keycloak_audience,
|
||||
leeway=self._settings.keycloak_leeway_seconds,
|
||||
options={"require": REQUIRED_CLAIMS, "verify_aud": True},
|
||||
)
|
||||
except jwt.PyJWTError as exc:
|
||||
raise TokenError(f"Token rejected: {exc}") from exc
|
||||
|
||||
return payload
|
||||
|
||||
async def healthy(self) -> bool:
|
||||
"""Readiness probe: can we reach Keycloak and load its keys?"""
|
||||
try:
|
||||
await self._fetch_jwks()
|
||||
except (httpx.HTTPError, TokenError, ValueError):
|
||||
logger.warning("Keycloak readiness check failed", exc_info=True)
|
||||
return False
|
||||
return True
|
||||
@@ -0,0 +1,53 @@
|
||||
"""The authenticated caller, decoupled from Keycloak's claim layout.
|
||||
|
||||
This module is the only place that knows how Keycloak shapes a token. Routes
|
||||
and services depend on `Principal` alone, so swapping the IdP -- or absorbing a
|
||||
claim-format change -- touches one file.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
|
||||
class Principal(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
subject: str
|
||||
username: str | None = None
|
||||
email: str | None = None
|
||||
full_name: str | None = None
|
||||
realm_roles: frozenset[str] = Field(default_factory=frozenset)
|
||||
client_roles: frozenset[str] = Field(default_factory=frozenset)
|
||||
scopes: frozenset[str] = Field(default_factory=frozenset)
|
||||
raw_claims: dict[str, Any] = Field(default_factory=dict, repr=False)
|
||||
|
||||
@property
|
||||
def roles(self) -> frozenset[str]:
|
||||
"""Realm and client roles together; what authorization checks use."""
|
||||
return self.realm_roles | self.client_roles
|
||||
|
||||
def has_any_role(self, *roles: str) -> bool:
|
||||
return bool(self.roles.intersection(roles))
|
||||
|
||||
def has_all_roles(self, *roles: str) -> bool:
|
||||
return set(roles).issubset(self.roles)
|
||||
|
||||
@classmethod
|
||||
def from_claims(cls, claims: dict[str, Any], *, client_id: str) -> Principal:
|
||||
realm_access = claims.get("realm_access") or {}
|
||||
resource_access = claims.get("resource_access") or {}
|
||||
client_access = resource_access.get(client_id) or {}
|
||||
|
||||
return cls(
|
||||
subject=str(claims["sub"]),
|
||||
username=claims.get("preferred_username"),
|
||||
email=claims.get("email"),
|
||||
full_name=claims.get("name"),
|
||||
realm_roles=frozenset(realm_access.get("roles") or ()),
|
||||
client_roles=frozenset(client_access.get("roles") or ()),
|
||||
scopes=frozenset(str(claims.get("scope") or "").split()),
|
||||
raw_claims=claims,
|
||||
)
|
||||
Whitespace-only changes.
@@ -0,0 +1,82 @@
|
||||
"""Application settings, loaded from the environment (and `.env` locally)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from functools import lru_cache
|
||||
from typing import Literal
|
||||
|
||||
from pydantic import Field, field_validator
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
Environment = Literal["local", "test", "staging", "production"]
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
model_config = SettingsConfigDict(
|
||||
env_file=".env",
|
||||
env_file_encoding="utf-8",
|
||||
env_nested_delimiter="__",
|
||||
extra="ignore",
|
||||
)
|
||||
|
||||
# --- application -----------------------------------------------------
|
||||
app_name: str = "v2x-server"
|
||||
app_env: Environment = "local"
|
||||
debug: bool = False
|
||||
log_level: str = "INFO"
|
||||
api_v1_prefix: str = "/api/v1"
|
||||
cors_origins: list[str] = Field(default_factory=list)
|
||||
|
||||
# --- database --------------------------------------------------------
|
||||
database_url: str = "mysql+asyncmy://v2x:v2xpassword@mysql:3306/v2x?charset=utf8mb4"
|
||||
test_database_url: str = "mysql+asyncmy://v2x:v2xpassword@mysql:3306/v2x_test?charset=utf8mb4"
|
||||
db_pool_size: int = 5
|
||||
db_max_overflow: int = 10
|
||||
db_echo: bool = False
|
||||
|
||||
# --- keycloak --------------------------------------------------------
|
||||
#
|
||||
# These are two different URLs on purpose. `keycloak_issuer` is the exact
|
||||
# string compared against a token's `iss` claim, which reflects however the
|
||||
# *client* reached Keycloak (http://localhost:8080 in local development).
|
||||
# `keycloak_internal_url` is how *this process* reaches Keycloak to fetch
|
||||
# discovery and JWKS (http://keycloak:8080 on the compose network).
|
||||
# In deployed environments both point at the same public URL.
|
||||
keycloak_issuer: str = "http://localhost:8080/realms/v2x"
|
||||
keycloak_internal_url: str = "http://keycloak:8080"
|
||||
keycloak_realm: str = "v2x"
|
||||
keycloak_audience: str = "v2x-api"
|
||||
keycloak_swagger_client_id: str = "v2x-swagger"
|
||||
keycloak_jwks_ttl_seconds: int = 3600
|
||||
# Floor between forced JWKS refreshes, so a flood of tokens bearing unknown
|
||||
# key ids cannot be turned into a request amplifier against Keycloak.
|
||||
keycloak_jwks_min_refresh_seconds: int = 30
|
||||
keycloak_leeway_seconds: int = 30
|
||||
keycloak_timeout_seconds: float = 5.0
|
||||
|
||||
@field_validator("keycloak_issuer", "keycloak_internal_url")
|
||||
@classmethod
|
||||
def _strip_trailing_slash(cls, value: str) -> str:
|
||||
return value.rstrip("/")
|
||||
|
||||
@property
|
||||
def is_local(self) -> bool:
|
||||
return self.app_env in ("local", "test")
|
||||
|
||||
@property
|
||||
def discovery_url(self) -> str:
|
||||
"""OIDC discovery document, fetched over the *internal* route."""
|
||||
return (
|
||||
f"{self.keycloak_internal_url}"
|
||||
f"/realms/{self.keycloak_realm}/.well-known/openid-configuration"
|
||||
)
|
||||
|
||||
@property
|
||||
def public_realm_url(self) -> str:
|
||||
"""Realm base URL as a browser sees it; used for Swagger's OAuth flow."""
|
||||
return self.keycloak_issuer
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def get_settings() -> Settings:
|
||||
return Settings()
|
||||
@@ -0,0 +1,128 @@
|
||||
"""Uniform error responses in RFC 9457 `application/problem+json` form."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from collections.abc import Mapping
|
||||
from typing import Any
|
||||
|
||||
from fastapi import FastAPI, HTTPException, Request, status
|
||||
from fastapi.exceptions import RequestValidationError
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from v2x_server.core.logging import request_id_ctx
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
PROBLEM_JSON = "application/problem+json"
|
||||
|
||||
|
||||
class AppError(Exception):
|
||||
"""Base for errors the application raises deliberately."""
|
||||
|
||||
status_code: int = status.HTTP_500_INTERNAL_SERVER_ERROR
|
||||
title: str = "Internal Server Error"
|
||||
|
||||
def __init__(self, detail: str | None = None) -> None:
|
||||
self.detail = detail or self.title
|
||||
super().__init__(self.detail)
|
||||
|
||||
|
||||
class NotFoundError(AppError):
|
||||
status_code = status.HTTP_404_NOT_FOUND
|
||||
title = "Not Found"
|
||||
|
||||
|
||||
class ConflictError(AppError):
|
||||
status_code = status.HTTP_409_CONFLICT
|
||||
title = "Conflict"
|
||||
|
||||
|
||||
def _problem(
|
||||
status_code: int,
|
||||
title: str,
|
||||
detail: str,
|
||||
*,
|
||||
headers: Mapping[str, str] | None = None,
|
||||
**extra: Any,
|
||||
) -> JSONResponse:
|
||||
body: dict[str, Any] = {
|
||||
"type": "about:blank",
|
||||
"title": title,
|
||||
"status": status_code,
|
||||
"detail": detail,
|
||||
"request_id": request_id_ctx.get(),
|
||||
**extra,
|
||||
}
|
||||
return JSONResponse(
|
||||
status_code=status_code,
|
||||
content=body,
|
||||
media_type=PROBLEM_JSON,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
|
||||
def register_exception_handlers(app: FastAPI) -> None:
|
||||
@app.exception_handler(AppError)
|
||||
async def _app_error(_: Request, exc: AppError) -> JSONResponse:
|
||||
return _problem(exc.status_code, exc.title, exc.detail)
|
||||
|
||||
@app.exception_handler(HTTPException)
|
||||
async def _http_error(_: Request, exc: HTTPException) -> JSONResponse:
|
||||
detail = exc.detail if isinstance(exc.detail, str) else str(exc.detail)
|
||||
return _problem(
|
||||
exc.status_code,
|
||||
_title_for(exc.status_code),
|
||||
detail,
|
||||
headers=exc.headers,
|
||||
)
|
||||
|
||||
@app.exception_handler(RequestValidationError)
|
||||
async def _validation_error(_: Request, exc: RequestValidationError) -> JSONResponse:
|
||||
return _problem(
|
||||
# Literal 422 rather than the constant: Starlette renamed it from
|
||||
# HTTP_422_UNPROCESSABLE_ENTITY to ..._CONTENT, so either spelling
|
||||
# ties us to a version range.
|
||||
422,
|
||||
"Validation Error",
|
||||
"The request body or parameters failed validation.",
|
||||
errors=_serializable_errors(exc),
|
||||
)
|
||||
|
||||
@app.exception_handler(Exception)
|
||||
async def _unhandled(request: Request, exc: Exception) -> JSONResponse:
|
||||
# Log the traceback, return an opaque body. The request id is the thread
|
||||
# connecting what the caller saw to what the logs recorded.
|
||||
logger.exception(
|
||||
"Unhandled exception on %s %s", request.method, request.url.path, exc_info=exc
|
||||
)
|
||||
return _problem(
|
||||
status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
"Internal Server Error",
|
||||
"An unexpected error occurred. Quote the request id when reporting this.",
|
||||
)
|
||||
|
||||
|
||||
def _serializable_errors(exc: RequestValidationError) -> list[dict[str, Any]]:
|
||||
"""Strip the non-JSON-serializable `ctx` values pydantic can attach."""
|
||||
cleaned: list[dict[str, Any]] = []
|
||||
for error in exc.errors():
|
||||
item = {k: v for k, v in error.items() if k != "ctx"}
|
||||
item["loc"] = [str(part) for part in error.get("loc", ())]
|
||||
cleaned.append(item)
|
||||
return cleaned
|
||||
|
||||
|
||||
def _title_for(status_code: int) -> str:
|
||||
titles = {
|
||||
400: "Bad Request",
|
||||
401: "Unauthorized",
|
||||
403: "Forbidden",
|
||||
404: "Not Found",
|
||||
405: "Method Not Allowed",
|
||||
409: "Conflict",
|
||||
422: "Validation Error",
|
||||
429: "Too Many Requests",
|
||||
503: "Service Unavailable",
|
||||
}
|
||||
return titles.get(status_code, "Error")
|
||||
@@ -0,0 +1,67 @@
|
||||
"""Structured logging plus a request-scoped correlation id.
|
||||
|
||||
The request id is held in a ContextVar so any log record emitted while handling
|
||||
a request carries it, without threading the value through call signatures.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import sys
|
||||
from contextvars import ContextVar
|
||||
from typing import Any
|
||||
|
||||
from pythonjsonlogger import json as jsonlogger
|
||||
|
||||
request_id_ctx: ContextVar[str | None] = ContextVar("request_id", default=None)
|
||||
|
||||
|
||||
class RequestIdFilter(logging.Filter):
|
||||
"""Stamps the current request id onto every record."""
|
||||
|
||||
def filter(self, record: logging.LogRecord) -> bool:
|
||||
record.request_id = request_id_ctx.get() or "-"
|
||||
return True
|
||||
|
||||
|
||||
class _JsonFormatter(jsonlogger.JsonFormatter):
|
||||
def add_fields(
|
||||
self,
|
||||
log_record: dict[str, Any],
|
||||
record: logging.LogRecord,
|
||||
message_dict: dict[str, Any],
|
||||
) -> None:
|
||||
super().add_fields(log_record, record, message_dict)
|
||||
log_record["level"] = record.levelname
|
||||
log_record["logger"] = record.name
|
||||
log_record.pop("levelname", None)
|
||||
|
||||
|
||||
def configure_logging(level: str = "INFO", *, json_output: bool = True) -> None:
|
||||
"""Install a single stdout handler on the root logger.
|
||||
|
||||
Called once at startup. Uvicorn's own loggers are left to propagate so
|
||||
everything lands in the same format.
|
||||
"""
|
||||
handler = logging.StreamHandler(sys.stdout)
|
||||
handler.addFilter(RequestIdFilter())
|
||||
|
||||
if json_output:
|
||||
handler.setFormatter(
|
||||
_JsonFormatter("%(asctime)s %(levelname)s %(name)s %(message)s %(request_id)s")
|
||||
)
|
||||
else:
|
||||
handler.setFormatter(
|
||||
logging.Formatter("%(asctime)s %(levelname)-8s [%(request_id)s] %(name)s: %(message)s")
|
||||
)
|
||||
|
||||
root = logging.getLogger()
|
||||
root.handlers = [handler]
|
||||
root.setLevel(level.upper())
|
||||
|
||||
# Uvicorn installs its own handlers; drop them so records propagate to root
|
||||
# instead of being emitted twice in two different formats.
|
||||
for name in ("uvicorn", "uvicorn.error", "uvicorn.access"):
|
||||
uvicorn_logger = logging.getLogger(name)
|
||||
uvicorn_logger.handlers = []
|
||||
uvicorn_logger.propagate = True
|
||||
Whitespace-only changes.
@@ -0,0 +1,43 @@
|
||||
"""Declarative base and shared column mixins."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import datetime as dt
|
||||
|
||||
from sqlalchemy import DateTime, MetaData, func
|
||||
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column
|
||||
|
||||
# Deterministic constraint names. This has to be in place *before* the first
|
||||
# migration: MySQL will otherwise invent its own names, and Alembic autogenerate
|
||||
# then produces diffs referencing constraints it cannot reliably drop.
|
||||
NAMING_CONVENTION = {
|
||||
"ix": "ix_%(column_0_label)s",
|
||||
"uq": "uq_%(table_name)s_%(column_0_name)s",
|
||||
"ck": "ck_%(table_name)s_%(constraint_name)s",
|
||||
"fk": "fk_%(table_name)s_%(column_0_name)s_%(referred_table_name)s",
|
||||
"pk": "pk_%(table_name)s",
|
||||
}
|
||||
|
||||
|
||||
class Base(DeclarativeBase):
|
||||
metadata = MetaData(naming_convention=NAMING_CONVENTION)
|
||||
|
||||
|
||||
class TimestampMixin:
|
||||
"""Server-side created/updated timestamps.
|
||||
|
||||
Defaults are rendered by MySQL rather than Python so rows written outside
|
||||
the application (migrations, manual SQL) are stamped consistently.
|
||||
"""
|
||||
|
||||
created_at: Mapped[dt.datetime] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
server_default=func.now(),
|
||||
nullable=False,
|
||||
)
|
||||
updated_at: Mapped[dt.datetime] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
server_default=func.now(),
|
||||
onupdate=func.now(),
|
||||
nullable=False,
|
||||
)
|
||||
@@ -0,0 +1,104 @@
|
||||
"""Async engine, session factory, and the FastAPI session dependency."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import AsyncIterator
|
||||
|
||||
from sqlalchemy import event
|
||||
from sqlalchemy.ext.asyncio import (
|
||||
AsyncEngine,
|
||||
AsyncSession,
|
||||
async_sessionmaker,
|
||||
create_async_engine,
|
||||
)
|
||||
|
||||
from v2x_server.core.config import Settings
|
||||
|
||||
_engine: AsyncEngine | None = None
|
||||
_sessionmaker: async_sessionmaker[AsyncSession] | None = None
|
||||
|
||||
|
||||
def create_engine(settings: Settings) -> AsyncEngine:
|
||||
kwargs: dict[str, object] = {
|
||||
"echo": settings.db_echo,
|
||||
# Verify a pooled connection is alive before handing it out.
|
||||
"pool_pre_ping": True,
|
||||
}
|
||||
|
||||
# SQLite (used only for quick local smoke tests) has no such pool knobs.
|
||||
if not settings.database_url.startswith("sqlite"):
|
||||
kwargs.update(
|
||||
pool_size=settings.db_pool_size,
|
||||
max_overflow=settings.db_max_overflow,
|
||||
# MySQL closes idle connections at `wait_timeout` (8h by default,
|
||||
# often far lower behind a proxy). Recycling below that threshold
|
||||
# is what prevents intermittent "MySQL server has gone away".
|
||||
pool_recycle=1800,
|
||||
)
|
||||
|
||||
engine = create_async_engine(settings.database_url, **kwargs)
|
||||
|
||||
if engine.dialect.name == "mysql":
|
||||
_force_utc(engine)
|
||||
|
||||
return engine
|
||||
|
||||
|
||||
def _force_utc(engine: AsyncEngine) -> None:
|
||||
"""Pin every MySQL connection to UTC.
|
||||
|
||||
MySQL's DATETIME stores no timezone, so SQLAlchemy's `timezone=True` is a
|
||||
no-op there and `CURRENT_TIMESTAMP` follows the *server's* zone. Without
|
||||
this, timestamps silently depend on wherever the database happens to run.
|
||||
"""
|
||||
|
||||
@event.listens_for(engine.sync_engine, "connect")
|
||||
def _set_session_timezone(dbapi_connection: object, _record: object) -> None:
|
||||
cursor = dbapi_connection.cursor() # type: ignore[attr-defined]
|
||||
try:
|
||||
cursor.execute("SET SESSION time_zone = '+00:00'")
|
||||
finally:
|
||||
cursor.close()
|
||||
|
||||
|
||||
def init_engine(settings: Settings) -> AsyncEngine:
|
||||
"""Build the process-wide engine and session factory. Called at startup."""
|
||||
global _engine, _sessionmaker
|
||||
_engine = create_engine(settings)
|
||||
_sessionmaker = async_sessionmaker(
|
||||
bind=_engine,
|
||||
expire_on_commit=False,
|
||||
autoflush=False,
|
||||
)
|
||||
return _engine
|
||||
|
||||
|
||||
async def dispose_engine() -> None:
|
||||
"""Close pooled connections. Called at shutdown."""
|
||||
global _engine, _sessionmaker
|
||||
if _engine is not None:
|
||||
await _engine.dispose()
|
||||
_engine = None
|
||||
_sessionmaker = None
|
||||
|
||||
|
||||
def get_sessionmaker() -> async_sessionmaker[AsyncSession]:
|
||||
if _sessionmaker is None:
|
||||
raise RuntimeError("Database engine not initialised; is the app lifespan running?")
|
||||
return _sessionmaker
|
||||
|
||||
|
||||
async def get_session() -> AsyncIterator[AsyncSession]:
|
||||
"""Request-scoped session: commit on success, roll back on any exception.
|
||||
|
||||
Endpoints therefore never need to call `commit()` themselves, and a raised
|
||||
exception can never leave a partial write behind.
|
||||
"""
|
||||
async with get_sessionmaker()() as session:
|
||||
try:
|
||||
yield session
|
||||
except Exception:
|
||||
await session.rollback()
|
||||
raise
|
||||
else:
|
||||
await session.commit()
|
||||
@@ -0,0 +1,121 @@
|
||||
"""Application factory and ASGI entrypoint."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import uuid
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
import httpx
|
||||
from fastapi import FastAPI, Request, Response
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.middleware.gzip import GZipMiddleware
|
||||
from starlette.middleware.base import RequestResponseEndpoint
|
||||
|
||||
from v2x_server import __version__
|
||||
from v2x_server.api.router import api_router
|
||||
from v2x_server.api.v1 import health
|
||||
from v2x_server.auth.keycloak import OIDCProvider
|
||||
from v2x_server.core.config import Settings, get_settings
|
||||
from v2x_server.core.errors import register_exception_handlers
|
||||
from v2x_server.core.logging import configure_logging, request_id_ctx
|
||||
from v2x_server.db.session import dispose_engine, init_engine
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
REQUEST_ID_HEADER = "X-Request-ID"
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI) -> AsyncIterator[None]:
|
||||
settings: Settings = app.state.settings
|
||||
|
||||
init_engine(settings)
|
||||
|
||||
# One shared client for the process: connection reuse against Keycloak
|
||||
# matters, since JWKS refreshes happen on the request path.
|
||||
client = httpx.AsyncClient(timeout=settings.keycloak_timeout_seconds)
|
||||
app.state.http_client = client
|
||||
app.state.oidc_provider = OIDCProvider(settings, client)
|
||||
|
||||
logger.info(
|
||||
"Started %s (env=%s, issuer=%s)",
|
||||
settings.app_name,
|
||||
settings.app_env,
|
||||
settings.keycloak_issuer,
|
||||
)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
await client.aclose()
|
||||
await dispose_engine()
|
||||
logger.info("Shutdown complete")
|
||||
|
||||
|
||||
def create_app(settings: Settings | None = None) -> FastAPI:
|
||||
settings = settings or get_settings()
|
||||
configure_logging(settings.log_level, json_output=not settings.is_local)
|
||||
|
||||
app = FastAPI(
|
||||
title="v2x-server",
|
||||
version=__version__,
|
||||
description=(
|
||||
"REST API backed by MySQL. Authentication is delegated to Keycloak: "
|
||||
"send a Bearer access token, or use Authorize below to obtain one."
|
||||
),
|
||||
lifespan=lifespan,
|
||||
docs_url="/docs",
|
||||
redoc_url="/redoc",
|
||||
openapi_url="/openapi.json",
|
||||
# Lets the Authorize button run the authorization-code + PKCE flow
|
||||
# against Keycloak, so /docs is usable without pasting tokens by hand.
|
||||
swagger_ui_init_oauth={
|
||||
"clientId": settings.keycloak_swagger_client_id,
|
||||
"usePkceWithAuthorizationCodeGrant": True,
|
||||
"scopes": "openid profile email",
|
||||
},
|
||||
)
|
||||
app.state.settings = settings
|
||||
|
||||
if settings.cors_origins:
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=settings.cors_origins,
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
expose_headers=[REQUEST_ID_HEADER],
|
||||
)
|
||||
|
||||
app.add_middleware(GZipMiddleware, minimum_size=1000)
|
||||
|
||||
@app.middleware("http")
|
||||
async def request_id_middleware(
|
||||
request: Request, call_next: RequestResponseEndpoint
|
||||
) -> Response:
|
||||
"""Adopt the caller's request id or mint one, then echo it back.
|
||||
|
||||
Set before anything else runs so even a 500 from a downstream handler
|
||||
carries an id the client can quote and the logs can be searched by.
|
||||
"""
|
||||
request_id = request.headers.get(REQUEST_ID_HEADER) or uuid.uuid4().hex
|
||||
token = request_id_ctx.set(request_id)
|
||||
try:
|
||||
response = await call_next(request)
|
||||
finally:
|
||||
request_id_ctx.reset(token)
|
||||
response.headers[REQUEST_ID_HEADER] = request_id
|
||||
return response
|
||||
|
||||
register_exception_handlers(app)
|
||||
|
||||
# Probes stay unversioned at the root; orchestrators shouldn't have to know
|
||||
# about an API version to decide whether a container is alive.
|
||||
app.include_router(health.router)
|
||||
app.include_router(api_router, prefix=settings.api_v1_prefix)
|
||||
|
||||
return app
|
||||
|
||||
|
||||
app = create_app()
|
||||
@@ -0,0 +1,10 @@
|
||||
"""Model package.
|
||||
|
||||
Every model must be imported here: Alembic's autogenerate only sees tables that
|
||||
have been registered on `Base.metadata` by import time.
|
||||
"""
|
||||
|
||||
from v2x_server.db.base import Base
|
||||
from v2x_server.models.device import Device, DeviceStatus
|
||||
|
||||
__all__ = ["Base", "Device", "DeviceStatus"]
|
||||
@@ -0,0 +1,46 @@
|
||||
"""Placeholder domain model proving the routing -> auth -> ORM -> migration path.
|
||||
|
||||
Replace with the real V2X domain once it is defined; nothing else depends on it.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import enum
|
||||
|
||||
from sqlalchemy import Enum as SAEnum
|
||||
from sqlalchemy import String
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from v2x_server.db.base import Base, TimestampMixin
|
||||
|
||||
|
||||
class DeviceStatus(enum.StrEnum):
|
||||
ACTIVE = "active"
|
||||
INACTIVE = "inactive"
|
||||
MAINTENANCE = "maintenance"
|
||||
|
||||
|
||||
class Device(Base, TimestampMixin):
|
||||
__tablename__ = "devices"
|
||||
|
||||
id: Mapped[int] = mapped_column(primary_key=True, autoincrement=True)
|
||||
name: Mapped[str] = mapped_column(String(128), nullable=False, index=True)
|
||||
serial: Mapped[str] = mapped_column(String(64), nullable=False, unique=True)
|
||||
status: Mapped[DeviceStatus] = mapped_column(
|
||||
# native_enum=False stores a VARCHAR + CHECK instead of a MySQL ENUM
|
||||
# column: adding a value later is an ordinary migration rather than an
|
||||
# ALTER TABLE that rewrites the whole table.
|
||||
SAEnum(
|
||||
DeviceStatus,
|
||||
native_enum=False,
|
||||
length=32,
|
||||
values_callable=lambda enum_cls: [member.value for member in enum_cls],
|
||||
),
|
||||
nullable=False,
|
||||
default=DeviceStatus.ACTIVE,
|
||||
server_default=DeviceStatus.ACTIVE.value,
|
||||
)
|
||||
description: Mapped[str | None] = mapped_column(String(512), nullable=True)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<Device id={self.id} serial={self.serial!r}>"
|
||||
Whitespace-only changes.
@@ -0,0 +1,59 @@
|
||||
"""Data access for devices. Keeps queries out of route handlers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from v2x_server.models.device import Device, DeviceStatus
|
||||
from v2x_server.schemas.device import DeviceCreate, DeviceUpdate
|
||||
|
||||
|
||||
class DeviceRepository:
|
||||
def __init__(self, session: AsyncSession) -> None:
|
||||
self._session = session
|
||||
|
||||
async def get(self, device_id: int) -> Device | None:
|
||||
return await self._session.get(Device, device_id)
|
||||
|
||||
async def get_by_serial(self, serial: str) -> Device | None:
|
||||
result = await self._session.execute(select(Device).where(Device.serial == serial))
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def list(
|
||||
self,
|
||||
*,
|
||||
limit: int = 50,
|
||||
offset: int = 0,
|
||||
status: DeviceStatus | None = None,
|
||||
) -> tuple[Sequence[Device], int]:
|
||||
filters = [Device.status == status] if status is not None else []
|
||||
|
||||
total = await self._session.scalar(select(func.count()).select_from(Device).where(*filters))
|
||||
result = await self._session.execute(
|
||||
select(Device).where(*filters).order_by(Device.id).limit(limit).offset(offset)
|
||||
)
|
||||
return result.scalars().all(), int(total or 0)
|
||||
|
||||
async def create(self, payload: DeviceCreate) -> Device:
|
||||
device = Device(**payload.model_dump())
|
||||
self._session.add(device)
|
||||
# Flush rather than commit: surfaces constraint violations here, while
|
||||
# leaving the transaction boundary to the session dependency.
|
||||
await self._session.flush()
|
||||
await self._session.refresh(device)
|
||||
return device
|
||||
|
||||
async def update(self, device: Device, payload: DeviceUpdate) -> Device:
|
||||
# exclude_unset keeps a PATCH from nulling fields the caller omitted.
|
||||
for field, value in payload.model_dump(exclude_unset=True).items():
|
||||
setattr(device, field, value)
|
||||
await self._session.flush()
|
||||
await self._session.refresh(device)
|
||||
return device
|
||||
|
||||
async def delete(self, device: Device) -> None:
|
||||
await self._session.delete(device)
|
||||
await self._session.flush()
|
||||
Whitespace-only changes.
@@ -0,0 +1,48 @@
|
||||
"""Request/response schemas for the devices resource."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import datetime as dt
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
from v2x_server.models.device import DeviceStatus
|
||||
|
||||
|
||||
class DeviceBase(BaseModel):
|
||||
name: str = Field(min_length=1, max_length=128)
|
||||
serial: str = Field(min_length=1, max_length=64)
|
||||
status: DeviceStatus = DeviceStatus.ACTIVE
|
||||
description: str | None = Field(default=None, max_length=512)
|
||||
|
||||
|
||||
class DeviceCreate(DeviceBase):
|
||||
pass
|
||||
|
||||
|
||||
class DeviceUpdate(BaseModel):
|
||||
"""All fields optional: this is a PATCH payload.
|
||||
|
||||
`model_dump(exclude_unset=True)` at the call site distinguishes "not
|
||||
supplied" from "explicitly set to null".
|
||||
"""
|
||||
|
||||
name: str | None = Field(default=None, min_length=1, max_length=128)
|
||||
serial: str | None = Field(default=None, min_length=1, max_length=64)
|
||||
status: DeviceStatus | None = None
|
||||
description: str | None = Field(default=None, max_length=512)
|
||||
|
||||
|
||||
class DeviceRead(DeviceBase):
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
id: int
|
||||
created_at: dt.datetime
|
||||
updated_at: dt.datetime
|
||||
|
||||
|
||||
class DevicePage(BaseModel):
|
||||
items: list[DeviceRead]
|
||||
total: int
|
||||
limit: int
|
||||
offset: int
|
||||
Reference in new issue
Block a user