219 lines
7.5 KiB
Python
219 lines
7.5 KiB
Python
"""FastAPI application for atomic channel builds."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import hmac
|
|
from typing import Annotated, Any, Dict, List, Optional
|
|
|
|
from fastapi import Depends, FastAPI, Path, Request, Response
|
|
from fastapi.exceptions import RequestValidationError
|
|
from fastapi.responses import JSONResponse
|
|
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
|
from pydantic import BaseModel, ConfigDict, Field, field_validator
|
|
from starlette.exceptions import HTTPException as StarletteHTTPException
|
|
|
|
from .service import ApiError, BuildService, Settings
|
|
|
|
MAX_BODY_BYTES = 64 * 1024
|
|
CHANNEL_PATH = Path(
|
|
...,
|
|
min_length=32,
|
|
max_length=32,
|
|
pattern=r"^[0-9a-z]{32}$",
|
|
description="32 lowercase alphanumeric channel id",
|
|
)
|
|
|
|
bearer_scheme = HTTPBearer(auto_error=False)
|
|
|
|
|
|
class BuildRequest(BaseModel):
|
|
model_config = ConfigDict(extra="forbid")
|
|
|
|
deployment_domains: List[str] = Field(min_length=1, max_length=8)
|
|
reporting_domains: List[str] = Field(min_length=1, max_length=8)
|
|
force: bool = Field(default=False, strict=True)
|
|
|
|
@field_validator("deployment_domains", "reporting_domains")
|
|
@classmethod
|
|
def domains_must_be_nonempty_strings(cls, value: List[str]) -> List[str]:
|
|
for item in value:
|
|
if not isinstance(item, str) or not item.strip():
|
|
raise ValueError("domain entries must be non-empty strings")
|
|
return value
|
|
|
|
|
|
def _error_response(status: int, code: str, message: str) -> JSONResponse:
|
|
return JSONResponse(
|
|
status_code=status,
|
|
content={"error": {"code": code, "message": message}},
|
|
headers={"Cache-Control": "no-store"},
|
|
)
|
|
|
|
|
|
def _validation_message(exc: RequestValidationError) -> str:
|
|
errors = exc.errors()
|
|
if not errors:
|
|
return "request validation failed"
|
|
first = errors[0]
|
|
loc = [str(part) for part in first.get("loc", ()) if part != "body"]
|
|
msg = str(first.get("msg", "invalid"))
|
|
error_type = str(first.get("type", ""))
|
|
|
|
if error_type == "json_invalid" or "JSON decode" in msg:
|
|
return "request body must be valid JSON"
|
|
if error_type == "model_attributes_type":
|
|
return "request body must be a JSON object"
|
|
if error_type == "extra_forbidden" and loc:
|
|
return f"unknown fields: {loc[-1]}"
|
|
if "force" in loc and error_type.startswith("bool"):
|
|
return "force must be a boolean"
|
|
if any(part == "channel_id" for part in first.get("loc", ())):
|
|
return "channel id must be 32 lowercase alphanumeric characters"
|
|
if loc and loc[-1] in {"deployment_domains", "reporting_domains"}:
|
|
field = loc[-1]
|
|
if "too_short" in error_type or "too_long" in error_type:
|
|
return f"{field} must contain 1 to 8 domains"
|
|
return f"{field} contains an invalid domain"
|
|
if loc:
|
|
return f"{'.'.join(loc)}: {msg}"
|
|
return msg
|
|
|
|
|
|
def get_settings(request: Request) -> Settings:
|
|
return request.app.state.settings
|
|
|
|
|
|
def get_service(request: Request) -> BuildService:
|
|
return request.app.state.service
|
|
|
|
|
|
def require_auth(
|
|
credentials: Annotated[
|
|
Optional[HTTPAuthorizationCredentials], Depends(bearer_scheme)
|
|
],
|
|
settings: Annotated[Settings, Depends(get_settings)],
|
|
) -> None:
|
|
token = credentials.credentials if credentials is not None else ""
|
|
if not token or not hmac.compare_digest(token, settings.bearer_token):
|
|
raise ApiError(401, "unauthorized", "valid Bearer token required")
|
|
|
|
|
|
def create_application(
|
|
settings: Optional[Settings] = None,
|
|
service: Optional[BuildService] = None,
|
|
) -> FastAPI:
|
|
resolved_settings = settings or Settings.from_env()
|
|
resolved_service = service or BuildService(resolved_settings)
|
|
|
|
app = FastAPI(
|
|
title="Coruna Build API",
|
|
version="1.0.0",
|
|
docs_url=None,
|
|
redoc_url=None,
|
|
openapi_url=None,
|
|
)
|
|
app.state.settings = resolved_settings
|
|
app.state.service = resolved_service
|
|
|
|
@app.middleware("http")
|
|
async def enforce_body_limit(request: Request, call_next: Any) -> Response:
|
|
content_length = request.headers.get("content-length")
|
|
if content_length is not None:
|
|
try:
|
|
length = int(content_length)
|
|
except ValueError:
|
|
return _error_response(
|
|
422, "validation_error", "valid Content-Length required"
|
|
)
|
|
if length > MAX_BODY_BYTES:
|
|
return _error_response(
|
|
422, "validation_error", "request body size is invalid"
|
|
)
|
|
response = await call_next(request)
|
|
if "cache-control" not in response.headers:
|
|
response.headers["Cache-Control"] = "no-store"
|
|
return response
|
|
|
|
@app.exception_handler(ApiError)
|
|
async def api_error_handler(_: Request, exc: ApiError) -> JSONResponse:
|
|
return _error_response(exc.status, exc.code, exc.message)
|
|
|
|
@app.exception_handler(RequestValidationError)
|
|
async def validation_error_handler(
|
|
_: Request, exc: RequestValidationError
|
|
) -> JSONResponse:
|
|
return _error_response(422, "validation_error", _validation_message(exc))
|
|
|
|
@app.exception_handler(StarletteHTTPException)
|
|
async def http_exception_handler(
|
|
_: Request, exc: StarletteHTTPException
|
|
) -> JSONResponse:
|
|
if exc.status_code == 404:
|
|
return _error_response(404, "not_found", "route not found")
|
|
if exc.status_code == 405:
|
|
return _error_response(405, "method_not_allowed", "method not allowed")
|
|
return _error_response(
|
|
exc.status_code, "http_error", str(exc.detail) or "request failed"
|
|
)
|
|
|
|
@app.exception_handler(Exception)
|
|
async def unhandled_error_handler(_: Request, __: Exception) -> JSONResponse:
|
|
return _error_response(500, "internal_error", "internal server error")
|
|
|
|
@app.get("/health")
|
|
def health() -> Dict[str, str]:
|
|
return {"status": "ok"}
|
|
|
|
@app.post(
|
|
"/v1/channels/{channel_id}/build",
|
|
dependencies=[Depends(require_auth)],
|
|
status_code=201,
|
|
)
|
|
def build_channel(
|
|
channel_id: Annotated[str, CHANNEL_PATH],
|
|
body: BuildRequest,
|
|
service: Annotated[BuildService, Depends(get_service)],
|
|
) -> JSONResponse:
|
|
request_input = {
|
|
"channel_id": channel_id,
|
|
"deployment_domains": BuildService.normalize_domains(
|
|
body.deployment_domains, "deployment_domains"
|
|
),
|
|
"reporting_domains": BuildService.normalize_domains(
|
|
body.reporting_domains, "reporting_domains"
|
|
),
|
|
}
|
|
status, payload, headers = service.build(
|
|
channel_id, request_input, force=body.force
|
|
)
|
|
return JSONResponse(
|
|
status_code=status,
|
|
content=payload,
|
|
headers={"Cache-Control": "no-store", **dict(headers)},
|
|
)
|
|
|
|
@app.get(
|
|
"/v1/channels/{channel_id}",
|
|
dependencies=[Depends(require_auth)],
|
|
)
|
|
def channel_status(
|
|
channel_id: Annotated[str, CHANNEL_PATH],
|
|
service: Annotated[BuildService, Depends(get_service)],
|
|
) -> Dict[str, Any]:
|
|
return service.status(channel_id)
|
|
|
|
@app.delete(
|
|
"/v1/channels/{channel_id}",
|
|
dependencies=[Depends(require_auth)],
|
|
status_code=204,
|
|
response_class=Response,
|
|
)
|
|
def delete_channel(
|
|
channel_id: Annotated[str, CHANNEL_PATH],
|
|
service: Annotated[BuildService, Depends(get_service)],
|
|
) -> Response:
|
|
service.delete(channel_id)
|
|
return Response(status_code=204, headers={"Cache-Control": "no-store"})
|
|
|
|
return app
|