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