init
This commit is contained in:
@@ -0,0 +1,218 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user