SyntheticMDProductions's picture
Update ADAM safety, UI, and model workflows (#1)
c61c435
Raw
History Blame Contribute Delete
4.21 kB
from __future__ import annotations
import json
import math
from dataclasses import dataclass
from typing import Any
class RemoteApiError(ValueError):
def __init__(self, message: str, *, status: int = 400) -> None:
super().__init__(message)
self.status = status
@dataclass(frozen=True, slots=True)
class RemoteResponse:
status: int
body: bytes
content_type: str = "application/json"
headers: dict[str, str] | None = None
def json_response(payload: dict[str, Any], *, status: int = 200) -> RemoteResponse:
return RemoteResponse(
status=status,
body=json.dumps(payload, separators=(",", ":")).encode("utf-8"),
content_type="application/json",
)
def error_response(message: str, *, status: int = 400) -> RemoteResponse:
return json_response({"ok": False, "error": message}, status=status)
def media_response(body: bytes, content_type: str, *, cache_seconds: int = 86400) -> RemoteResponse:
return RemoteResponse(
status=200,
body=body,
content_type=content_type,
headers={"Cache-Control": f"private, max-age={max(0, int(cache_seconds))}"},
)
def bounded_text(value: Any, *, max_length: int, label: str, required: bool = False) -> str:
if value is None:
value = ""
if not isinstance(value, str):
value = str(value)
text = value.strip()
if required and not text:
raise RemoteApiError(f"{label} is required.")
if len(text) > max_length:
raise RemoteApiError(f"{label} must be {max_length} characters or shorter.")
return text
def bounded_int(
value: Any,
*,
minimum: int,
maximum: int,
default: int,
label: str,
) -> int:
if value in (None, ""):
return default
try:
if isinstance(value, bool):
raise ValueError
number = int(value)
except (TypeError, ValueError) as exc:
raise RemoteApiError(f"{label} must be a whole number.") from exc
if number < minimum or number > maximum:
raise RemoteApiError(f"{label} must be between {minimum} and {maximum}.")
return number
def bounded_float(
value: Any,
*,
minimum: float,
maximum: float,
default: float,
label: str,
) -> float:
if value in (None, ""):
return default
try:
if isinstance(value, bool):
raise ValueError
number = float(value)
except (TypeError, ValueError) as exc:
raise RemoteApiError(f"{label} must be a number.") from exc
if not math.isfinite(number) or number < minimum or number > maximum:
raise RemoteApiError(f"{label} must be between {minimum:g} and {maximum:g}.")
return number
def parse_pagination(query: dict[str, list[str]], *, default_size: int = 24, max_size: int = 80) -> dict[str, int]:
page = bounded_int(
(query.get("page") or ["1"])[0],
minimum=1,
maximum=1_000_000,
default=1,
label="Page",
)
page_size = bounded_int(
(query.get("page_size") or [str(default_size)])[0],
minimum=1,
maximum=max_size,
default=default_size,
label="Page size",
)
return {
"page": page,
"page_size": page_size,
"offset": (page - 1) * page_size,
"limit": page_size,
}
def coerce_json_object(payload: Any) -> dict[str, Any]:
if not isinstance(payload, dict):
raise RemoteApiError("Send a JSON object.")
return payload
def sanitized_arguments(arguments: dict[str, Any]) -> dict[str, Any]:
"""Return client-safe arguments without absolute filesystem paths."""
hidden = {
"dataset_dir",
"output_dir",
"model_path",
"base_model",
"base_model_path",
"resume_from",
"reference_image",
}
clean: dict[str, Any] = {}
for key, value in arguments.items():
if key in hidden:
text = str(value or "")
clean[f"{key}_name"] = text.replace("\\", "/").rstrip("/").rsplit("/", 1)[-1] if text else ""
continue
if isinstance(value, (str, int, float, bool)) or value is None:
clean[key] = value
return clean