Files
hashbro 9153a4f557 admin
2026-08-09 02:42:45 +08:00

237 lines
8.2 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)
SUPPORT_TEMPLATES = ("test", "blank")
DEFAULT_SUPPORT_TEMPLATE = "test"
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)
support_template: str = Field(default=DEFAULT_SUPPORT_TEMPLATE)
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
@field_validator("support_template")
@classmethod
def support_template_must_be_known(cls, value: str) -> str:
template = value.strip().lower()
if template not in SUPPORT_TEMPLATES:
raise ValueError(
f"support_template must be one of: {', '.join(SUPPORT_TEMPLATES)}"
)
return template
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 "support_template" in loc:
return f"support_template must be one of: {', '.join(SUPPORT_TEMPLATES)}"
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"
),
"support_template": body.support_template,
}
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