File size: 4,210 Bytes
c61c435 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 | 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
|