comfygit-studio 0.5.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- comfygit_studio/__init__.py +10 -0
- comfygit_studio/api_schema.py +564 -0
- comfygit_studio/embedded.py +295 -0
- comfygit_studio/executor.py +1002 -0
- comfygit_studio/openapi/studio-contract-api.v1.json +1379 -0
- comfygit_studio/runtime.py +3155 -0
- comfygit_studio/state.py +1019 -0
- comfygit_studio/static/assets/geist-cyrillic-wght-normal-CHSlOQsW.woff2 +0 -0
- comfygit_studio/static/assets/geist-latin-ext-wght-normal-DMtmJ5ZE.woff2 +0 -0
- comfygit_studio/static/assets/geist-latin-wght-normal-Dm3htQBi.woff2 +0 -0
- comfygit_studio/static/assets/index-BDmIh9tA.css +1 -0
- comfygit_studio/static/assets/index-BlnDFmNd.js +17 -0
- comfygit_studio/static/index.html +14 -0
- comfygit_studio-0.5.0.dist-info/METADATA +25 -0
- comfygit_studio-0.5.0.dist-info/RECORD +16 -0
- comfygit_studio-0.5.0.dist-info/WHEEL +4 -0
|
@@ -0,0 +1,3155 @@
|
|
|
1
|
+
"""Contract-serving runtime for ComfyGit environments."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import asyncio
|
|
6
|
+
import hashlib
|
|
7
|
+
import json
|
|
8
|
+
import secrets
|
|
9
|
+
import subprocess
|
|
10
|
+
import uuid
|
|
11
|
+
from collections.abc import Callable, Mapping
|
|
12
|
+
from dataclasses import asdict, dataclass
|
|
13
|
+
from importlib import resources
|
|
14
|
+
from pathlib import Path
|
|
15
|
+
from typing import Any, cast
|
|
16
|
+
from urllib.parse import quote
|
|
17
|
+
|
|
18
|
+
import aiohttp
|
|
19
|
+
from aiohttp import web
|
|
20
|
+
from comfygit_core import Environment
|
|
21
|
+
from comfygit_core.models import NamedWorkflowContract
|
|
22
|
+
from comfygit_core.workflow import build_manifest_contract_prompt
|
|
23
|
+
|
|
24
|
+
from .api_schema import studio_contract_api_openapi
|
|
25
|
+
from .executor import (
|
|
26
|
+
PROXY_AUTH_HEADER,
|
|
27
|
+
ComfyGitServeTimeoutError,
|
|
28
|
+
ComfyUIClient,
|
|
29
|
+
ComfyUIExecutionError,
|
|
30
|
+
ComfyUIRequestError,
|
|
31
|
+
LocalComfyExecutor,
|
|
32
|
+
ProxyComfyExecutor,
|
|
33
|
+
RunExecutionRequest,
|
|
34
|
+
RunExecutor,
|
|
35
|
+
StagedUpload,
|
|
36
|
+
_safe_artifact_filename,
|
|
37
|
+
_write_local_artifact,
|
|
38
|
+
artifact_dimensions,
|
|
39
|
+
output_kind,
|
|
40
|
+
workflow_contract_output_from_payload,
|
|
41
|
+
)
|
|
42
|
+
from .state import (
|
|
43
|
+
EphemeralServeStateStore,
|
|
44
|
+
ServeGalleryItem,
|
|
45
|
+
ServeRunOutputSlot,
|
|
46
|
+
ServeRunRecord,
|
|
47
|
+
ServeSession,
|
|
48
|
+
ServeStateStore,
|
|
49
|
+
SQLiteServeStateStore,
|
|
50
|
+
utc_now,
|
|
51
|
+
)
|
|
52
|
+
|
|
53
|
+
DEFAULT_MAX_REQUEST_BYTES = 256 * 1024 * 1024
|
|
54
|
+
DEFAULT_RUN_TIMEOUT_SECONDS = 12 * 60 * 60
|
|
55
|
+
UPLOAD_TOKEN_BYTES = 32
|
|
56
|
+
DEFAULT_UPLOAD_CONTENT_TYPE = "application/octet-stream"
|
|
57
|
+
DEFAULT_UPLOAD_EXTENSION = ".bin"
|
|
58
|
+
SESSION_COOKIE_NAME = "comfygit_studio_session"
|
|
59
|
+
SESSION_HEADER_NAME = "X-ComfyGit-Studio-Session"
|
|
60
|
+
SHARED_GALLERY_SCOPE = "shared"
|
|
61
|
+
MAX_GALLERY_PAGE_LIMIT = 200
|
|
62
|
+
FILE_UPLOAD_CONTRACT_INPUT_TYPES = {"image", "audio", "video", "file"}
|
|
63
|
+
ACTIVE_RUN_STATUSES = {"submitted", "running"}
|
|
64
|
+
TERMINAL_RUN_STATUSES = {"completed", "error", "failed", "cancelled"}
|
|
65
|
+
OUTPUT_REQUEST_HEADERS = ("Range", "If-Range")
|
|
66
|
+
ENVIRONMENT_REF_DIGEST_FILES = ("pyproject.toml",)
|
|
67
|
+
ENVIRONMENT_REF_DIGEST_DIRS = ("workflow_api",)
|
|
68
|
+
GIT_COMMAND_TIMEOUT_SECONDS = 2
|
|
69
|
+
|
|
70
|
+
UPLOAD_FILE_TYPES: tuple[tuple[str, tuple[str, ...]], ...] = (
|
|
71
|
+
("image/jpeg", (".jpg", ".jpeg")),
|
|
72
|
+
("image/png", (".png",)),
|
|
73
|
+
("image/webp", (".webp",)),
|
|
74
|
+
("image/gif", (".gif",)),
|
|
75
|
+
("image/bmp", (".bmp",)),
|
|
76
|
+
("video/mp4", (".mp4",)),
|
|
77
|
+
("video/webm", (".webm",)),
|
|
78
|
+
("video/quicktime", (".mov",)),
|
|
79
|
+
("audio/wav", (".wav",)),
|
|
80
|
+
("audio/mpeg", (".mp3",)),
|
|
81
|
+
)
|
|
82
|
+
UPLOAD_MIME_TYPE_BY_EXTENSION = {
|
|
83
|
+
extension: mime_type for mime_type, extensions in UPLOAD_FILE_TYPES for extension in extensions
|
|
84
|
+
}
|
|
85
|
+
UPLOAD_EXTENSION_BY_MIME_TYPE = {
|
|
86
|
+
mime_type: extensions[0] for mime_type, extensions in UPLOAD_FILE_TYPES
|
|
87
|
+
}
|
|
88
|
+
UPLOAD_MIME_TYPE_ALIASES = {
|
|
89
|
+
"image/jpg": "image/jpeg",
|
|
90
|
+
"audio/mp3": "audio/mpeg",
|
|
91
|
+
"audio/x-wav": "audio/wav",
|
|
92
|
+
}
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
@dataclass(frozen=True)
|
|
96
|
+
class ServeConfig:
|
|
97
|
+
"""Configuration for the local ComfyGit serve adapter."""
|
|
98
|
+
|
|
99
|
+
host: str
|
|
100
|
+
port: int
|
|
101
|
+
comfy_url: str
|
|
102
|
+
max_request_bytes: int = DEFAULT_MAX_REQUEST_BYTES
|
|
103
|
+
run_timeout_seconds: float = DEFAULT_RUN_TIMEOUT_SECONDS
|
|
104
|
+
state: str = "ephemeral"
|
|
105
|
+
gallery: str = "private"
|
|
106
|
+
state_db: Path | None = None
|
|
107
|
+
role: str = "studio"
|
|
108
|
+
executor: str = "local"
|
|
109
|
+
proxy_url: str | None = None
|
|
110
|
+
proxy_token: str | None = None
|
|
111
|
+
callback_url: str | None = None
|
|
112
|
+
callback_token: str | None = None
|
|
113
|
+
artifact_dir: Path | None = None
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
@dataclass
|
|
117
|
+
class UploadRecord:
|
|
118
|
+
upload_id: str
|
|
119
|
+
token: str
|
|
120
|
+
filename: str
|
|
121
|
+
content_type: str
|
|
122
|
+
size: int | None
|
|
123
|
+
path: Path
|
|
124
|
+
comfyui_filename: str
|
|
125
|
+
status: str = "pending"
|
|
126
|
+
|
|
127
|
+
def public_ref(self) -> dict[str, Any]:
|
|
128
|
+
payload: dict[str, Any] = {
|
|
129
|
+
"kind": "file_ref",
|
|
130
|
+
"ref": self.upload_id,
|
|
131
|
+
"filename": self.filename,
|
|
132
|
+
"mime_type": self.content_type,
|
|
133
|
+
}
|
|
134
|
+
if self.size is not None:
|
|
135
|
+
payload["size"] = self.size
|
|
136
|
+
return payload
|
|
137
|
+
|
|
138
|
+
|
|
139
|
+
@dataclass(frozen=True)
|
|
140
|
+
class PreparedContractInputs:
|
|
141
|
+
inputs: dict[str, Any]
|
|
142
|
+
staged_uploads: tuple[StagedUpload, ...] = ()
|
|
143
|
+
|
|
144
|
+
|
|
145
|
+
@dataclass
|
|
146
|
+
class ProxyRuntimeRun:
|
|
147
|
+
prompt_id: str
|
|
148
|
+
status: str
|
|
149
|
+
outputs: tuple[Any, ...]
|
|
150
|
+
raw_result: dict[str, Any]
|
|
151
|
+
error: str | None = None
|
|
152
|
+
callback: ProxyCallbackTarget | None = None
|
|
153
|
+
|
|
154
|
+
|
|
155
|
+
@dataclass(frozen=True)
|
|
156
|
+
class ProxyArtifactRef:
|
|
157
|
+
params: dict[str, str]
|
|
158
|
+
|
|
159
|
+
|
|
160
|
+
@dataclass(frozen=True)
|
|
161
|
+
class ProxyCallbackTarget:
|
|
162
|
+
run_id: str
|
|
163
|
+
url: str
|
|
164
|
+
token: str | None = None
|
|
165
|
+
|
|
166
|
+
|
|
167
|
+
@dataclass(frozen=True)
|
|
168
|
+
class WorkerCallbackUpload:
|
|
169
|
+
field_name: str
|
|
170
|
+
filename: str
|
|
171
|
+
content_type: str
|
|
172
|
+
body: bytes
|
|
173
|
+
|
|
174
|
+
|
|
175
|
+
class ServeState:
|
|
176
|
+
"""Shared state for request handlers."""
|
|
177
|
+
|
|
178
|
+
def __init__(
|
|
179
|
+
self,
|
|
180
|
+
env: Environment,
|
|
181
|
+
config: ServeConfig,
|
|
182
|
+
session: aiohttp.ClientSession,
|
|
183
|
+
state_store: ServeStateStore | None = None,
|
|
184
|
+
) -> None:
|
|
185
|
+
self.env = env
|
|
186
|
+
self.config = config
|
|
187
|
+
self.client = ComfyUIClient(config.comfy_url, session=session)
|
|
188
|
+
self.artifact_dir = _serve_artifact_dir(env, config)
|
|
189
|
+
if config.role == "studio" and config.executor == "proxy":
|
|
190
|
+
if not config.proxy_url:
|
|
191
|
+
raise ValueError("--proxy-url is required when --executor proxy is used.")
|
|
192
|
+
self.executor: RunExecutor = ProxyComfyExecutor(
|
|
193
|
+
config.proxy_url,
|
|
194
|
+
session=session,
|
|
195
|
+
token=config.proxy_token,
|
|
196
|
+
artifact_dir=self.artifact_dir,
|
|
197
|
+
)
|
|
198
|
+
else:
|
|
199
|
+
self.executor = LocalComfyExecutor(self.client, artifact_dir=self.artifact_dir)
|
|
200
|
+
self.uploads: dict[str, UploadRecord] = {}
|
|
201
|
+
self.proxy_runs: dict[str, ProxyRuntimeRun] = {}
|
|
202
|
+
self.proxy_artifacts: dict[str, ProxyArtifactRef] = {}
|
|
203
|
+
self.state_store = state_store or _create_state_store(env, config)
|
|
204
|
+
self.background_tasks: set[asyncio.Task[Any]] = set()
|
|
205
|
+
self.active_run_tasks: dict[str, asyncio.Task[Any]] = {}
|
|
206
|
+
|
|
207
|
+
def manifest_snapshot(self):
|
|
208
|
+
return self.env.get_manifest_snapshot()
|
|
209
|
+
|
|
210
|
+
|
|
211
|
+
SERVE_STATE_KEY = web.AppKey("serve_state", ServeState)
|
|
212
|
+
STUDIO_STATIC_DIR_KEY = web.AppKey("studio_static_dir", Path)
|
|
213
|
+
STUDIO_API_BASE_PATH_KEY = web.AppKey("studio_api_base_path", str)
|
|
214
|
+
|
|
215
|
+
|
|
216
|
+
def serve_environment(env: Environment, config: ServeConfig) -> None:
|
|
217
|
+
"""Run the local contract-serving HTTP server until interrupted."""
|
|
218
|
+
|
|
219
|
+
asyncio.run(_serve_environment_async(env, config))
|
|
220
|
+
|
|
221
|
+
|
|
222
|
+
async def _serve_environment_async(env: Environment, config: ServeConfig) -> None:
|
|
223
|
+
async with aiohttp.ClientSession() as session:
|
|
224
|
+
state = ServeState(env, config, session)
|
|
225
|
+
try:
|
|
226
|
+
app = create_proxy_app(state) if config.role == "proxy" else create_app(state)
|
|
227
|
+
runner = web.AppRunner(app)
|
|
228
|
+
await runner.setup()
|
|
229
|
+
site = web.TCPSite(runner, config.host, config.port)
|
|
230
|
+
await site.start()
|
|
231
|
+
print(f"Serving ComfyGit environment '{env.name}' on http://{config.host}:{config.port}")
|
|
232
|
+
print(f"Serve role: {config.role}")
|
|
233
|
+
if config.role == "proxy":
|
|
234
|
+
print(f"Proxy ComfyUI API target: {config.comfy_url}")
|
|
235
|
+
else:
|
|
236
|
+
print(f"Executor: {config.executor}")
|
|
237
|
+
print(f"ComfyUI API target: {config.comfy_url}")
|
|
238
|
+
if config.executor == "proxy":
|
|
239
|
+
print(f"Proxy runtime target: {config.proxy_url}")
|
|
240
|
+
if config.callback_url:
|
|
241
|
+
print(f"Worker callback base URL: {config.callback_url}")
|
|
242
|
+
print(f"Serve state: {config.state} ({'persistent' if state.state_store.persistent else 'ephemeral'})")
|
|
243
|
+
print("Press Ctrl+C to stop.")
|
|
244
|
+
try:
|
|
245
|
+
await asyncio.Event().wait()
|
|
246
|
+
finally:
|
|
247
|
+
await runner.cleanup()
|
|
248
|
+
finally:
|
|
249
|
+
state.state_store.close()
|
|
250
|
+
|
|
251
|
+
|
|
252
|
+
def create_app(
|
|
253
|
+
state: ServeState,
|
|
254
|
+
*,
|
|
255
|
+
static_dir: Path | None = None,
|
|
256
|
+
api_base_path: str = "",
|
|
257
|
+
) -> web.Application:
|
|
258
|
+
"""Create the aiohttp application for a ComfyGit serve runtime."""
|
|
259
|
+
|
|
260
|
+
app = web.Application(client_max_size=_max_request_bytes(state))
|
|
261
|
+
app[SERVE_STATE_KEY] = state
|
|
262
|
+
register_studio_routes(
|
|
263
|
+
app,
|
|
264
|
+
static_dir=static_dir,
|
|
265
|
+
api_base_path=api_base_path,
|
|
266
|
+
)
|
|
267
|
+
app.on_startup.append(_recover_active_runs_on_startup)
|
|
268
|
+
app.on_cleanup.append(_cleanup_background_tasks)
|
|
269
|
+
return app
|
|
270
|
+
|
|
271
|
+
|
|
272
|
+
def register_studio_routes(
|
|
273
|
+
app: web.Application,
|
|
274
|
+
*,
|
|
275
|
+
static_dir: Path | None = None,
|
|
276
|
+
route_api_prefix: str = "",
|
|
277
|
+
route_ui_prefix: str = "",
|
|
278
|
+
api_base_path: str | None = None,
|
|
279
|
+
include_spa_fallback: bool = True,
|
|
280
|
+
) -> None:
|
|
281
|
+
"""Register Studio UI and contract API routes on an aiohttp application.
|
|
282
|
+
|
|
283
|
+
`route_api_prefix` and `route_ui_prefix` are server-side route prefixes.
|
|
284
|
+
`api_base_path` is the browser-facing API prefix injected into Studio.
|
|
285
|
+
They can differ when a host such as ComfyUI exposes registered routes under
|
|
286
|
+
a public `/api/...` prefix.
|
|
287
|
+
"""
|
|
288
|
+
|
|
289
|
+
static_dir = static_dir or _studio_static_dir()
|
|
290
|
+
app[STUDIO_STATIC_DIR_KEY] = static_dir
|
|
291
|
+
app[STUDIO_API_BASE_PATH_KEY] = _normalize_public_prefix(
|
|
292
|
+
api_base_path if api_base_path is not None else route_api_prefix
|
|
293
|
+
)
|
|
294
|
+
|
|
295
|
+
api_prefix = _normalize_route_prefix(route_api_prefix)
|
|
296
|
+
ui_prefix = _normalize_route_prefix(route_ui_prefix)
|
|
297
|
+
|
|
298
|
+
app.router.add_get(_route_path(ui_prefix, "/"), studio_index_handler)
|
|
299
|
+
if ui_prefix:
|
|
300
|
+
app.router.add_get(ui_prefix, studio_index_handler)
|
|
301
|
+
if (static_dir / "assets").exists():
|
|
302
|
+
app.router.add_static(_route_path(ui_prefix, "/assets/"), static_dir / "assets", append_version=True)
|
|
303
|
+
app.router.add_get(_route_path(ui_prefix, "/favicon.ico"), favicon_handler)
|
|
304
|
+
|
|
305
|
+
app.router.add_get(_route_path(api_prefix, "/openapi.json"), openapi_handler)
|
|
306
|
+
app.router.add_get(_route_path(api_prefix, "/health"), health_handler)
|
|
307
|
+
app.router.add_get(_route_path(api_prefix, "/contracts"), contracts_handler)
|
|
308
|
+
app.router.add_get(
|
|
309
|
+
_route_path(api_prefix, "/contracts/{workflow}/{contract}"),
|
|
310
|
+
single_contract_handler,
|
|
311
|
+
)
|
|
312
|
+
app.router.add_post(_route_path(api_prefix, "/uploads/prepare"), upload_prepare_handler)
|
|
313
|
+
app.router.add_put(_route_path(api_prefix, "/uploads/{upload_id}"), upload_put_handler)
|
|
314
|
+
app.router.add_get(_route_path(api_prefix, "/uploads/{upload_id}/status"), upload_status_handler)
|
|
315
|
+
app.router.add_get(_route_path(api_prefix, "/gallery"), gallery_handler)
|
|
316
|
+
app.router.add_delete(_route_path(api_prefix, "/gallery/{item_id}"), gallery_delete_handler)
|
|
317
|
+
app.router.add_get(_route_path(api_prefix, "/runs"), runs_handler)
|
|
318
|
+
app.router.add_get(_route_path(api_prefix, "/runs/{run_id}"), single_run_handler)
|
|
319
|
+
app.router.add_post(_route_path(api_prefix, "/runs/{run_id}/cancel"), cancel_run_handler)
|
|
320
|
+
app.router.add_post(
|
|
321
|
+
_route_path(api_prefix, "/worker-callback/runs/{run_id}"),
|
|
322
|
+
worker_callback_handler,
|
|
323
|
+
)
|
|
324
|
+
app.router.add_post(
|
|
325
|
+
_route_path(api_prefix, "/contracts/{workflow}/{contract}/run"),
|
|
326
|
+
run_contract_handler,
|
|
327
|
+
)
|
|
328
|
+
app.router.add_get(_route_path(api_prefix, "/outputs/view"), output_view_handler)
|
|
329
|
+
|
|
330
|
+
if include_spa_fallback:
|
|
331
|
+
app.router.add_get(_route_path(ui_prefix, "/{tail:.*}"), studio_index_handler)
|
|
332
|
+
|
|
333
|
+
|
|
334
|
+
def create_proxy_app(state: ServeState) -> web.Application:
|
|
335
|
+
"""Create the compute-only proxy runtime app for remote execution."""
|
|
336
|
+
|
|
337
|
+
app = web.Application(client_max_size=_max_request_bytes(state))
|
|
338
|
+
app[SERVE_STATE_KEY] = state
|
|
339
|
+
app.router.add_get("/proxy/health", proxy_health_handler)
|
|
340
|
+
app.router.add_post("/proxy/runs", proxy_run_create_handler)
|
|
341
|
+
app.router.add_get("/proxy/runs/{prompt_id}", proxy_run_status_handler)
|
|
342
|
+
app.router.add_post("/proxy/runs/{prompt_id}/cancel", proxy_run_cancel_handler)
|
|
343
|
+
app.router.add_get("/proxy/artifacts/{artifact_id}", proxy_artifact_handler)
|
|
344
|
+
app.on_cleanup.append(_cleanup_background_tasks)
|
|
345
|
+
return app
|
|
346
|
+
|
|
347
|
+
|
|
348
|
+
async def _recover_active_runs_on_startup(app: web.Application) -> None:
|
|
349
|
+
await _ensure_active_run_recovery(app[SERVE_STATE_KEY])
|
|
350
|
+
|
|
351
|
+
|
|
352
|
+
async def _cleanup_background_tasks(app: web.Application) -> None:
|
|
353
|
+
state = app[SERVE_STATE_KEY]
|
|
354
|
+
background_tasks = getattr(state, "background_tasks", None)
|
|
355
|
+
if not isinstance(background_tasks, set):
|
|
356
|
+
return
|
|
357
|
+
tasks = [task for task in background_tasks if isinstance(task, asyncio.Task)]
|
|
358
|
+
for task in tasks:
|
|
359
|
+
task.cancel()
|
|
360
|
+
if tasks:
|
|
361
|
+
await asyncio.gather(*tasks, return_exceptions=True)
|
|
362
|
+
background_tasks.clear()
|
|
363
|
+
active_run_tasks = getattr(state, "active_run_tasks", None)
|
|
364
|
+
if isinstance(active_run_tasks, dict):
|
|
365
|
+
active_run_tasks.clear()
|
|
366
|
+
|
|
367
|
+
|
|
368
|
+
async def _ensure_active_run_recovery(state: ServeState) -> None:
|
|
369
|
+
if not _state_supports_active_run_recovery(state):
|
|
370
|
+
return
|
|
371
|
+
if _state_uses_proxy_callbacks(state):
|
|
372
|
+
return
|
|
373
|
+
for run in state.state_store.list_active_runs(ACTIVE_RUN_STATUSES):
|
|
374
|
+
if not run.prompt_id:
|
|
375
|
+
_record_recovery_failure(state, run, "Active run has no ComfyUI prompt id to recover.")
|
|
376
|
+
continue
|
|
377
|
+
task = state.active_run_tasks.get(run.run_id)
|
|
378
|
+
if task is not None and not task.done():
|
|
379
|
+
continue
|
|
380
|
+
outputs = _contract_outputs_for_run(state, run.workflow, run.contract)
|
|
381
|
+
if outputs is None:
|
|
382
|
+
_record_recovery_failure(
|
|
383
|
+
state,
|
|
384
|
+
run,
|
|
385
|
+
f"Cannot recover active run because contract '{run.workflow} / {run.contract}' no longer exists.",
|
|
386
|
+
)
|
|
387
|
+
continue
|
|
388
|
+
session = ServeSession(session_id=run.session_id, scope_key=run.scope_key)
|
|
389
|
+
output_slots = _output_slots_for_run(
|
|
390
|
+
run_id=run.run_id,
|
|
391
|
+
session=session,
|
|
392
|
+
workflow_name=run.workflow,
|
|
393
|
+
contract_name=run.contract,
|
|
394
|
+
outputs=outputs,
|
|
395
|
+
prompt_id=run.prompt_id,
|
|
396
|
+
inputs=run.inputs,
|
|
397
|
+
)
|
|
398
|
+
task = asyncio.create_task(
|
|
399
|
+
_complete_submitted_run(
|
|
400
|
+
state,
|
|
401
|
+
session,
|
|
402
|
+
run_id=run.run_id,
|
|
403
|
+
workflow_name=run.workflow,
|
|
404
|
+
contract_name=run.contract,
|
|
405
|
+
inputs=run.inputs,
|
|
406
|
+
prompt_id=run.prompt_id,
|
|
407
|
+
outputs=outputs,
|
|
408
|
+
timeout_seconds=state.config.run_timeout_seconds,
|
|
409
|
+
poll_interval_seconds=1,
|
|
410
|
+
issues=_run_issues(run),
|
|
411
|
+
output_slots=output_slots,
|
|
412
|
+
created_at=run.created_at,
|
|
413
|
+
)
|
|
414
|
+
)
|
|
415
|
+
_track_active_run_task(state, run.run_id, task)
|
|
416
|
+
|
|
417
|
+
|
|
418
|
+
def _state_supports_active_run_recovery(state: Any) -> bool:
|
|
419
|
+
return (
|
|
420
|
+
isinstance(getattr(state, "active_run_tasks", None), dict)
|
|
421
|
+
and isinstance(getattr(state, "background_tasks", None), set)
|
|
422
|
+
and hasattr(state, "config")
|
|
423
|
+
and hasattr(state, "executor")
|
|
424
|
+
and hasattr(state, "state_store")
|
|
425
|
+
and callable(getattr(state, "manifest_snapshot", None))
|
|
426
|
+
)
|
|
427
|
+
|
|
428
|
+
|
|
429
|
+
def _state_uses_proxy_callbacks(state: Any) -> bool:
|
|
430
|
+
config = getattr(state, "config", None)
|
|
431
|
+
return (
|
|
432
|
+
getattr(config, "executor", "local") == "proxy"
|
|
433
|
+
and bool(getattr(config, "callback_url", None))
|
|
434
|
+
)
|
|
435
|
+
|
|
436
|
+
|
|
437
|
+
def _track_active_run_task(state: ServeState, run_id: str, task: asyncio.Task[Any]) -> None:
|
|
438
|
+
state.background_tasks.add(task)
|
|
439
|
+
state.active_run_tasks[run_id] = task
|
|
440
|
+
|
|
441
|
+
def discard(completed: asyncio.Task[Any]) -> None:
|
|
442
|
+
state.background_tasks.discard(completed)
|
|
443
|
+
if state.active_run_tasks.get(run_id) is completed:
|
|
444
|
+
state.active_run_tasks.pop(run_id, None)
|
|
445
|
+
|
|
446
|
+
task.add_done_callback(discard)
|
|
447
|
+
|
|
448
|
+
|
|
449
|
+
def _contract_outputs_for_run(state: ServeState, workflow_name: str, contract_name: str) -> tuple[Any, ...] | None:
|
|
450
|
+
manifest = state.manifest_snapshot()
|
|
451
|
+
workflow = manifest.workflows.get(workflow_name)
|
|
452
|
+
execution_contract = getattr(workflow, "execution_contract", None) if workflow else None
|
|
453
|
+
contract = execution_contract.contracts.get(contract_name) if execution_contract else None
|
|
454
|
+
if contract is None:
|
|
455
|
+
return None
|
|
456
|
+
return tuple(contract.outputs)
|
|
457
|
+
|
|
458
|
+
|
|
459
|
+
def _run_issues(run: ServeRunRecord) -> list[dict[str, Any]]:
|
|
460
|
+
raw_result = run.raw_result if isinstance(run.raw_result, Mapping) else {}
|
|
461
|
+
issues = raw_result.get("issues")
|
|
462
|
+
return [dict(issue) for issue in issues if isinstance(issue, Mapping)] if isinstance(issues, list) else []
|
|
463
|
+
|
|
464
|
+
|
|
465
|
+
def _record_recovery_failure(state: ServeState, run: ServeRunRecord, message: str) -> None:
|
|
466
|
+
session = ServeSession(session_id=run.session_id, scope_key=run.scope_key)
|
|
467
|
+
payload = {
|
|
468
|
+
"error": "recovery_failed",
|
|
469
|
+
"message": message,
|
|
470
|
+
"run_id": run.run_id,
|
|
471
|
+
"prompt_id": run.prompt_id,
|
|
472
|
+
}
|
|
473
|
+
_record_failed_run(
|
|
474
|
+
state,
|
|
475
|
+
session,
|
|
476
|
+
run.workflow,
|
|
477
|
+
run.contract,
|
|
478
|
+
{"inputs": run.inputs},
|
|
479
|
+
payload,
|
|
480
|
+
run_id=run.run_id,
|
|
481
|
+
prompt_id=run.prompt_id,
|
|
482
|
+
gallery_item_id=f"gallery_{run.run_id}",
|
|
483
|
+
created_at=run.created_at,
|
|
484
|
+
)
|
|
485
|
+
|
|
486
|
+
|
|
487
|
+
def _max_request_bytes(state: ServeState) -> int:
|
|
488
|
+
value = getattr(getattr(state, "config", None), "max_request_bytes", DEFAULT_MAX_REQUEST_BYTES)
|
|
489
|
+
return value if isinstance(value, int) and value > 0 else DEFAULT_MAX_REQUEST_BYTES
|
|
490
|
+
|
|
491
|
+
|
|
492
|
+
def _state(request: web.Request) -> ServeState:
|
|
493
|
+
return request.app[SERVE_STATE_KEY]
|
|
494
|
+
|
|
495
|
+
|
|
496
|
+
def _maybe_state(request: web.Request) -> ServeState | None:
|
|
497
|
+
return request.app.get(SERVE_STATE_KEY)
|
|
498
|
+
|
|
499
|
+
|
|
500
|
+
def _executor_unavailable_payload(
|
|
501
|
+
state: ServeState,
|
|
502
|
+
exc: BaseException,
|
|
503
|
+
*,
|
|
504
|
+
prompt_id: str | None = None,
|
|
505
|
+
) -> dict[str, Any]:
|
|
506
|
+
is_proxy = getattr(state.config, "executor", "local") == "proxy"
|
|
507
|
+
payload: dict[str, Any] = {
|
|
508
|
+
"error": "proxy_unavailable" if is_proxy else "comfyui_unavailable",
|
|
509
|
+
"message": str(exc),
|
|
510
|
+
}
|
|
511
|
+
if is_proxy:
|
|
512
|
+
payload["proxy_url"] = state.config.proxy_url
|
|
513
|
+
else:
|
|
514
|
+
payload["comfy_url"] = state.config.comfy_url
|
|
515
|
+
if prompt_id:
|
|
516
|
+
payload["prompt_id"] = prompt_id
|
|
517
|
+
return payload
|
|
518
|
+
|
|
519
|
+
|
|
520
|
+
def _environment_ref(env: Environment) -> dict[str, Any]:
|
|
521
|
+
cec_path = getattr(env, "cec_path", None)
|
|
522
|
+
cec_path = Path(cec_path) if cec_path is not None else None
|
|
523
|
+
return {
|
|
524
|
+
"environment": getattr(env, "name", None),
|
|
525
|
+
"cec_commit": _git_output(cec_path, "rev-parse", "HEAD") if cec_path else None,
|
|
526
|
+
"cec_dirty": _git_dirty(cec_path) if cec_path else None,
|
|
527
|
+
"contract_digest": _contract_digest(cec_path) if cec_path else None,
|
|
528
|
+
}
|
|
529
|
+
|
|
530
|
+
|
|
531
|
+
def _environment_ref_match(local_ref: Mapping[str, Any], remote_ref: Any) -> bool | None:
|
|
532
|
+
if not isinstance(remote_ref, Mapping):
|
|
533
|
+
return None
|
|
534
|
+
local_digest = local_ref.get("contract_digest")
|
|
535
|
+
remote_digest = remote_ref.get("contract_digest")
|
|
536
|
+
if not isinstance(local_digest, str) or not isinstance(remote_digest, str):
|
|
537
|
+
return None
|
|
538
|
+
return secrets.compare_digest(local_digest, remote_digest)
|
|
539
|
+
|
|
540
|
+
|
|
541
|
+
def _contract_digest(cec_path: Path) -> str | None:
|
|
542
|
+
paths = tuple(_contract_digest_paths(cec_path))
|
|
543
|
+
if not paths:
|
|
544
|
+
return None
|
|
545
|
+
digest = hashlib.sha256()
|
|
546
|
+
for relative_path in paths:
|
|
547
|
+
absolute_path = cec_path / relative_path
|
|
548
|
+
try:
|
|
549
|
+
file_bytes = absolute_path.read_bytes()
|
|
550
|
+
except OSError:
|
|
551
|
+
return None
|
|
552
|
+
digest.update(relative_path.as_posix().encode("utf-8"))
|
|
553
|
+
digest.update(b"\0")
|
|
554
|
+
digest.update(file_bytes)
|
|
555
|
+
digest.update(b"\0")
|
|
556
|
+
return f"sha256:{digest.hexdigest()}"
|
|
557
|
+
|
|
558
|
+
|
|
559
|
+
def _contract_digest_paths(cec_path: Path) -> tuple[Path, ...]:
|
|
560
|
+
paths: list[Path] = []
|
|
561
|
+
for filename in ENVIRONMENT_REF_DIGEST_FILES:
|
|
562
|
+
path = cec_path / filename
|
|
563
|
+
if path.is_file():
|
|
564
|
+
paths.append(Path(filename))
|
|
565
|
+
for dirname in ENVIRONMENT_REF_DIGEST_DIRS:
|
|
566
|
+
directory = cec_path / dirname
|
|
567
|
+
if directory.is_dir():
|
|
568
|
+
paths.extend(
|
|
569
|
+
path.relative_to(cec_path)
|
|
570
|
+
for path in sorted(directory.rglob("*"))
|
|
571
|
+
if path.is_file()
|
|
572
|
+
)
|
|
573
|
+
return tuple(paths)
|
|
574
|
+
|
|
575
|
+
|
|
576
|
+
def _git_dirty(repo_path: Path) -> bool | None:
|
|
577
|
+
status = _git_output(repo_path, "status", "--porcelain")
|
|
578
|
+
if status is None:
|
|
579
|
+
return None
|
|
580
|
+
return bool(status)
|
|
581
|
+
|
|
582
|
+
|
|
583
|
+
def _git_output(repo_path: Path | None, *args: str) -> str | None:
|
|
584
|
+
if repo_path is None or not repo_path.exists():
|
|
585
|
+
return None
|
|
586
|
+
try:
|
|
587
|
+
result = subprocess.run(
|
|
588
|
+
("git", *args),
|
|
589
|
+
cwd=repo_path,
|
|
590
|
+
check=False,
|
|
591
|
+
capture_output=True,
|
|
592
|
+
text=True,
|
|
593
|
+
timeout=GIT_COMMAND_TIMEOUT_SECONDS,
|
|
594
|
+
)
|
|
595
|
+
except (OSError, subprocess.SubprocessError):
|
|
596
|
+
return None
|
|
597
|
+
if result.returncode != 0:
|
|
598
|
+
return None
|
|
599
|
+
return result.stdout.strip()
|
|
600
|
+
|
|
601
|
+
|
|
602
|
+
def _create_state_store(env: Environment, config: ServeConfig) -> ServeStateStore:
|
|
603
|
+
if config.state == "local":
|
|
604
|
+
return SQLiteServeStateStore(config.state_db or _default_state_db_path(env))
|
|
605
|
+
return EphemeralServeStateStore()
|
|
606
|
+
|
|
607
|
+
|
|
608
|
+
def _default_state_db_path(env: Environment) -> Path:
|
|
609
|
+
workspace_paths = getattr(env, "workspace_paths", None)
|
|
610
|
+
metadata_dir = getattr(workspace_paths, "metadata", None)
|
|
611
|
+
if metadata_dir is not None:
|
|
612
|
+
return Path(metadata_dir) / "serve" / "serve.sqlite"
|
|
613
|
+
workspace = getattr(env, "workspace", None)
|
|
614
|
+
workspace_path = getattr(workspace, "path", None)
|
|
615
|
+
if workspace_path is not None:
|
|
616
|
+
return Path(workspace_path) / ".metadata" / "serve" / "serve.sqlite"
|
|
617
|
+
env_path = Path(getattr(env, "path", "."))
|
|
618
|
+
return env_path / ".metadata" / "serve" / "serve.sqlite"
|
|
619
|
+
|
|
620
|
+
|
|
621
|
+
def _serve_artifact_dir(env: Environment, config: ServeConfig) -> Path:
|
|
622
|
+
if config.artifact_dir is not None:
|
|
623
|
+
return config.artifact_dir
|
|
624
|
+
workspace_paths = getattr(env, "workspace_paths", None)
|
|
625
|
+
metadata_dir = getattr(workspace_paths, "metadata", None)
|
|
626
|
+
if metadata_dir is not None:
|
|
627
|
+
return Path(metadata_dir) / "serve" / "artifacts"
|
|
628
|
+
workspace = getattr(env, "workspace", None)
|
|
629
|
+
workspace_path = getattr(workspace, "path", None)
|
|
630
|
+
if workspace_path is not None:
|
|
631
|
+
return Path(workspace_path) / ".metadata" / "serve" / "artifacts"
|
|
632
|
+
env_path = Path(getattr(env, "path", "."))
|
|
633
|
+
return env_path / ".metadata" / "serve" / "artifacts"
|
|
634
|
+
|
|
635
|
+
|
|
636
|
+
def _studio_static_dir() -> Path:
|
|
637
|
+
return Path(str(resources.files("comfygit_studio").joinpath("static")))
|
|
638
|
+
|
|
639
|
+
|
|
640
|
+
def _normalize_route_prefix(prefix: str) -> str:
|
|
641
|
+
normalized = "/" + str(prefix or "").strip("/")
|
|
642
|
+
return "" if normalized == "/" else normalized
|
|
643
|
+
|
|
644
|
+
|
|
645
|
+
def _normalize_public_prefix(prefix: str | None) -> str:
|
|
646
|
+
normalized = "/" + str(prefix or "").strip("/")
|
|
647
|
+
return "" if normalized == "/" else normalized
|
|
648
|
+
|
|
649
|
+
|
|
650
|
+
def _route_path(prefix: str, path: str) -> str:
|
|
651
|
+
if not path.startswith("/"):
|
|
652
|
+
path = f"/{path}"
|
|
653
|
+
if not prefix:
|
|
654
|
+
return path
|
|
655
|
+
if path == "/":
|
|
656
|
+
return f"{prefix}/"
|
|
657
|
+
return f"{prefix}{path}"
|
|
658
|
+
|
|
659
|
+
|
|
660
|
+
async def studio_index_handler(request: web.Request) -> web.StreamResponse:
|
|
661
|
+
static_dir = request.app[STUDIO_STATIC_DIR_KEY]
|
|
662
|
+
index_path = static_dir / "index.html"
|
|
663
|
+
if index_path.exists():
|
|
664
|
+
html = index_path.read_text(encoding="utf-8")
|
|
665
|
+
state = _maybe_state(request)
|
|
666
|
+
if state is None:
|
|
667
|
+
return _studio_not_running_response()
|
|
668
|
+
env_name = getattr(state.env, "name", "Environment")
|
|
669
|
+
config = {
|
|
670
|
+
"apiBasePath": request.app.get(STUDIO_API_BASE_PATH_KEY, ""),
|
|
671
|
+
"authMode": "none",
|
|
672
|
+
"endpointName": env_name if isinstance(env_name, str) else "Environment",
|
|
673
|
+
}
|
|
674
|
+
script = f"<script>window.__COMFYGIT_STUDIO_CONFIG__ = {json.dumps(config)};</script>"
|
|
675
|
+
if "<head>" in html:
|
|
676
|
+
html = html.replace("<head>", f"<head>{script}", 1)
|
|
677
|
+
elif "</head>" in html:
|
|
678
|
+
html = html.replace("</head>", f"{script}</head>", 1)
|
|
679
|
+
else:
|
|
680
|
+
html = f"{script}{html}"
|
|
681
|
+
return web.Response(text=html, content_type="text/html")
|
|
682
|
+
return web.Response(
|
|
683
|
+
text=(
|
|
684
|
+
"<!doctype html><html><head><title>ComfyGit Studio</title></head>"
|
|
685
|
+
"<body><h1>ComfyGit Studio assets are not built.</h1>"
|
|
686
|
+
"<p>Run the contract studio build before packaging or serving the UI.</p>"
|
|
687
|
+
"</body></html>"
|
|
688
|
+
),
|
|
689
|
+
content_type="text/html",
|
|
690
|
+
)
|
|
691
|
+
|
|
692
|
+
|
|
693
|
+
def _studio_not_running_response() -> web.Response:
|
|
694
|
+
return web.Response(
|
|
695
|
+
status=503,
|
|
696
|
+
text=(
|
|
697
|
+
"<!doctype html>"
|
|
698
|
+
"<html>"
|
|
699
|
+
"<head>"
|
|
700
|
+
"<title>ComfyGit Studio Not Running</title>"
|
|
701
|
+
"<meta name=\"viewport\" content=\"width=device-width, initial-scale=1\" />"
|
|
702
|
+
"<style>"
|
|
703
|
+
"body{margin:0;min-height:100vh;display:grid;place-items:center;"
|
|
704
|
+
"background:#151719;color:#f2f2f2;font-family:system-ui,-apple-system,BlinkMacSystemFont,"
|
|
705
|
+
"'Segoe UI',sans-serif;}"
|
|
706
|
+
"main{max-width:560px;padding:32px;border:1px solid #44484f;background:#202226;}"
|
|
707
|
+
"h1{margin:0 0 12px;color:#2bbbf3;font-size:22px;letter-spacing:.02em;}"
|
|
708
|
+
"p{margin:0 0 10px;color:#c7c7c7;line-height:1.5;}"
|
|
709
|
+
"code{color:#f7f7f7;background:#111316;padding:2px 5px;}"
|
|
710
|
+
"</style>"
|
|
711
|
+
"</head>"
|
|
712
|
+
"<body>"
|
|
713
|
+
"<main>"
|
|
714
|
+
"<h1>ComfyGit Studio is not running</h1>"
|
|
715
|
+
"<p>The Studio route is available, but this ComfyUI process has not "
|
|
716
|
+
"started an embedded Studio session yet.</p>"
|
|
717
|
+
"<p>Open Studio from the ComfyGit Manager workflow panel, or call "
|
|
718
|
+
"<code>/api/v2/comfygit/studio/open</code> before loading this page.</p>"
|
|
719
|
+
"</main>"
|
|
720
|
+
"</body>"
|
|
721
|
+
"</html>"
|
|
722
|
+
),
|
|
723
|
+
content_type="text/html",
|
|
724
|
+
)
|
|
725
|
+
|
|
726
|
+
|
|
727
|
+
async def favicon_handler(_request: web.Request) -> web.Response:
|
|
728
|
+
return web.Response(status=204)
|
|
729
|
+
|
|
730
|
+
|
|
731
|
+
async def openapi_handler(_request: web.Request) -> web.Response:
|
|
732
|
+
return web.json_response(studio_contract_api_openapi())
|
|
733
|
+
|
|
734
|
+
|
|
735
|
+
async def health_handler(request: web.Request) -> web.Response:
|
|
736
|
+
state = _state(request)
|
|
737
|
+
local_environment_ref = _environment_ref(state.env)
|
|
738
|
+
payload: dict[str, Any] = {
|
|
739
|
+
"ok": True,
|
|
740
|
+
"environment": state.env.name,
|
|
741
|
+
"environment_ref": local_environment_ref,
|
|
742
|
+
"comfy_url": state.config.comfy_url,
|
|
743
|
+
"executor": getattr(state.config, "executor", "local"),
|
|
744
|
+
"comfyui": {"available": False},
|
|
745
|
+
}
|
|
746
|
+
if getattr(state.config, "executor", "local") == "proxy" and isinstance(state.executor, ProxyComfyExecutor):
|
|
747
|
+
payload["proxy"] = {
|
|
748
|
+
"configured": True,
|
|
749
|
+
"available": None,
|
|
750
|
+
"health_check": "deferred",
|
|
751
|
+
}
|
|
752
|
+
payload["proxy_environment_ref_match"] = None
|
|
753
|
+
payload["comfyui"] = {
|
|
754
|
+
"available": True,
|
|
755
|
+
"mode": "proxy",
|
|
756
|
+
"status": "deferred",
|
|
757
|
+
}
|
|
758
|
+
check_proxy = request.query.get("check_proxy", "").lower() in {"1", "true", "yes", "on"}
|
|
759
|
+
if not check_proxy:
|
|
760
|
+
return web.json_response(payload)
|
|
761
|
+
try:
|
|
762
|
+
proxy_health = await state.executor.check_health()
|
|
763
|
+
proxy_payload = {
|
|
764
|
+
"configured": True,
|
|
765
|
+
"available": bool(proxy_health.get("ok", True)),
|
|
766
|
+
"health_check": "checked",
|
|
767
|
+
}
|
|
768
|
+
proxy_payload.update(proxy_health)
|
|
769
|
+
payload["proxy"] = proxy_payload
|
|
770
|
+
payload["proxy_environment_ref_match"] = _environment_ref_match(
|
|
771
|
+
local_environment_ref,
|
|
772
|
+
proxy_health.get("environment_ref"),
|
|
773
|
+
)
|
|
774
|
+
payload["comfyui"] = proxy_health.get("comfyui", {"available": True})
|
|
775
|
+
except (ComfyUIRequestError, aiohttp.ClientError, asyncio.TimeoutError) as exc:
|
|
776
|
+
payload["proxy"] = {
|
|
777
|
+
"configured": True,
|
|
778
|
+
"available": False,
|
|
779
|
+
"health_check": "checked",
|
|
780
|
+
"error": str(exc),
|
|
781
|
+
}
|
|
782
|
+
payload["proxy_environment_ref_match"] = None
|
|
783
|
+
payload["comfyui"] = {"available": False, "error": str(exc)}
|
|
784
|
+
return web.json_response(payload)
|
|
785
|
+
try:
|
|
786
|
+
await state.client.check_health()
|
|
787
|
+
payload["comfyui"] = {"available": True}
|
|
788
|
+
except (aiohttp.ClientError, asyncio.TimeoutError) as exc:
|
|
789
|
+
payload["comfyui"] = {"available": False, "error": str(exc)}
|
|
790
|
+
return web.json_response(payload)
|
|
791
|
+
|
|
792
|
+
|
|
793
|
+
async def proxy_health_handler(request: web.Request) -> web.Response:
|
|
794
|
+
auth_response = _proxy_auth_response(request)
|
|
795
|
+
if auth_response is not None:
|
|
796
|
+
return auth_response
|
|
797
|
+
state = _state(request)
|
|
798
|
+
payload: dict[str, Any] = {
|
|
799
|
+
"ok": True,
|
|
800
|
+
"role": "proxy",
|
|
801
|
+
"environment": state.env.name,
|
|
802
|
+
"environment_ref": _environment_ref(state.env),
|
|
803
|
+
"comfy_url": state.config.comfy_url,
|
|
804
|
+
"comfyui": {"available": False},
|
|
805
|
+
}
|
|
806
|
+
try:
|
|
807
|
+
await state.client.check_health()
|
|
808
|
+
payload["comfyui"] = {"available": True}
|
|
809
|
+
except (aiohttp.ClientError, asyncio.TimeoutError) as exc:
|
|
810
|
+
payload["comfyui"] = {"available": False, "error": str(exc)}
|
|
811
|
+
return web.json_response(payload)
|
|
812
|
+
|
|
813
|
+
|
|
814
|
+
async def proxy_run_create_handler(request: web.Request) -> web.Response:
|
|
815
|
+
auth_response = _proxy_auth_response(request)
|
|
816
|
+
if auth_response is not None:
|
|
817
|
+
return auth_response
|
|
818
|
+
state = _state(request)
|
|
819
|
+
try:
|
|
820
|
+
payload = await _read_proxy_run_request(request, state)
|
|
821
|
+
prompt = payload.get("prompt")
|
|
822
|
+
if not isinstance(prompt, dict):
|
|
823
|
+
raise ValueError("Proxy run payload must include a prompt object.")
|
|
824
|
+
outputs_payload = payload.get("outputs", [])
|
|
825
|
+
if not isinstance(outputs_payload, list):
|
|
826
|
+
raise ValueError("Proxy run payload outputs must be a list.")
|
|
827
|
+
outputs = tuple(
|
|
828
|
+
workflow_contract_output_from_payload(output)
|
|
829
|
+
for output in outputs_payload
|
|
830
|
+
if isinstance(output, Mapping)
|
|
831
|
+
)
|
|
832
|
+
timeout_seconds = float(payload.get("timeout_seconds", state.config.run_timeout_seconds))
|
|
833
|
+
poll_interval_seconds = float(payload.get("poll_interval_seconds", 1))
|
|
834
|
+
cache_token = str(payload.get("cache_token") or uuid.uuid4().hex[:10])
|
|
835
|
+
callback = _proxy_callback_target(payload.get("callback"))
|
|
836
|
+
|
|
837
|
+
async def record_submitted(prompt_id: str) -> None:
|
|
838
|
+
state.proxy_runs[prompt_id] = ProxyRuntimeRun(
|
|
839
|
+
prompt_id=prompt_id,
|
|
840
|
+
status="submitted",
|
|
841
|
+
outputs=outputs,
|
|
842
|
+
raw_result={"status": "submitted", "prompt_id": prompt_id},
|
|
843
|
+
callback=callback,
|
|
844
|
+
)
|
|
845
|
+
|
|
846
|
+
execution = await state.executor.execute(
|
|
847
|
+
RunExecutionRequest(
|
|
848
|
+
prompt=prompt,
|
|
849
|
+
outputs=outputs,
|
|
850
|
+
wait=False,
|
|
851
|
+
timeout_seconds=timeout_seconds,
|
|
852
|
+
poll_interval_seconds=poll_interval_seconds,
|
|
853
|
+
cache_token=cache_token,
|
|
854
|
+
on_submitted=record_submitted,
|
|
855
|
+
)
|
|
856
|
+
)
|
|
857
|
+
task = asyncio.create_task(
|
|
858
|
+
_complete_proxy_runtime_run(
|
|
859
|
+
state,
|
|
860
|
+
execution.prompt_id,
|
|
861
|
+
outputs,
|
|
862
|
+
timeout_seconds=timeout_seconds,
|
|
863
|
+
poll_interval_seconds=poll_interval_seconds,
|
|
864
|
+
callback=callback,
|
|
865
|
+
)
|
|
866
|
+
)
|
|
867
|
+
_track_proxy_task(state, task)
|
|
868
|
+
return web.json_response({"status": execution.status, "prompt_id": execution.prompt_id})
|
|
869
|
+
except ValueError as exc:
|
|
870
|
+
return web.json_response({"error": "bad_request", "message": str(exc)}, status=400)
|
|
871
|
+
except ComfyUIRequestError as exc:
|
|
872
|
+
return web.json_response(
|
|
873
|
+
{
|
|
874
|
+
"error": "comfyui_rejected_request",
|
|
875
|
+
"message": str(exc),
|
|
876
|
+
"comfy_status": exc.status,
|
|
877
|
+
"comfy_url": exc.url,
|
|
878
|
+
"comfyui": exc.payload,
|
|
879
|
+
},
|
|
880
|
+
status=400 if exc.status == 400 else 502,
|
|
881
|
+
)
|
|
882
|
+
except (aiohttp.ClientError, asyncio.TimeoutError) as exc:
|
|
883
|
+
return web.json_response(
|
|
884
|
+
{
|
|
885
|
+
"error": "comfyui_unavailable",
|
|
886
|
+
"message": str(exc),
|
|
887
|
+
"comfy_url": state.config.comfy_url,
|
|
888
|
+
},
|
|
889
|
+
status=502,
|
|
890
|
+
)
|
|
891
|
+
|
|
892
|
+
|
|
893
|
+
async def proxy_run_status_handler(request: web.Request) -> web.Response:
|
|
894
|
+
auth_response = _proxy_auth_response(request)
|
|
895
|
+
if auth_response is not None:
|
|
896
|
+
return auth_response
|
|
897
|
+
record = _state(request).proxy_runs.get(request.match_info["prompt_id"])
|
|
898
|
+
if record is None:
|
|
899
|
+
return web.json_response({"error": "not_found", "message": "Unknown proxy run."}, status=404)
|
|
900
|
+
return web.json_response(_proxy_run_payload(record))
|
|
901
|
+
|
|
902
|
+
|
|
903
|
+
async def proxy_run_cancel_handler(request: web.Request) -> web.Response:
|
|
904
|
+
auth_response = _proxy_auth_response(request)
|
|
905
|
+
if auth_response is not None:
|
|
906
|
+
return auth_response
|
|
907
|
+
state = _state(request)
|
|
908
|
+
prompt_id = request.match_info["prompt_id"]
|
|
909
|
+
record = state.proxy_runs.get(prompt_id)
|
|
910
|
+
if record is None:
|
|
911
|
+
return web.json_response({"error": "not_found", "message": "Unknown proxy run."}, status=404)
|
|
912
|
+
try:
|
|
913
|
+
await state.executor.cancel(prompt_id)
|
|
914
|
+
except ComfyUIRequestError as exc:
|
|
915
|
+
return web.json_response({"error": "comfyui_rejected_cancel", "message": str(exc)}, status=400)
|
|
916
|
+
record.status = "cancelled"
|
|
917
|
+
record.raw_result = {"status": "cancelled", "prompt_id": prompt_id}
|
|
918
|
+
return web.json_response(_proxy_run_payload(record))
|
|
919
|
+
|
|
920
|
+
|
|
921
|
+
async def proxy_artifact_handler(request: web.Request) -> web.StreamResponse:
|
|
922
|
+
auth_response = _proxy_auth_response(request)
|
|
923
|
+
if auth_response is not None:
|
|
924
|
+
return auth_response
|
|
925
|
+
state = _state(request)
|
|
926
|
+
artifact = state.proxy_artifacts.get(request.match_info["artifact_id"])
|
|
927
|
+
if artifact is None:
|
|
928
|
+
return web.json_response({"error": "not_found", "message": "Unknown proxy artifact."}, status=404)
|
|
929
|
+
try:
|
|
930
|
+
output_response = await state.client.fetch_output(
|
|
931
|
+
artifact.params,
|
|
932
|
+
request_headers=_output_request_headers(request),
|
|
933
|
+
)
|
|
934
|
+
except (aiohttp.ClientError, asyncio.TimeoutError) as exc:
|
|
935
|
+
return web.json_response(
|
|
936
|
+
{
|
|
937
|
+
"error": "comfyui_unavailable",
|
|
938
|
+
"message": str(exc),
|
|
939
|
+
"comfy_url": state.config.comfy_url,
|
|
940
|
+
},
|
|
941
|
+
status=502,
|
|
942
|
+
)
|
|
943
|
+
headers = {"Content-Type": output_response.content_type}
|
|
944
|
+
if output_response.disposition:
|
|
945
|
+
headers["Content-Disposition"] = output_response.disposition
|
|
946
|
+
headers.update(output_response.headers)
|
|
947
|
+
return web.Response(body=output_response.body, status=output_response.status, headers=headers)
|
|
948
|
+
|
|
949
|
+
|
|
950
|
+
async def contracts_handler(request: web.Request) -> web.Response:
|
|
951
|
+
return web.json_response(_contracts_payload(_state(request)))
|
|
952
|
+
|
|
953
|
+
|
|
954
|
+
async def single_contract_handler(request: web.Request) -> web.Response:
|
|
955
|
+
try:
|
|
956
|
+
payload = _single_contract_payload(
|
|
957
|
+
_state(request),
|
|
958
|
+
request.match_info["workflow"],
|
|
959
|
+
request.match_info["contract"],
|
|
960
|
+
)
|
|
961
|
+
return web.json_response(payload)
|
|
962
|
+
except ValueError as exc:
|
|
963
|
+
return web.json_response({"error": "bad_request", "message": str(exc)}, status=400)
|
|
964
|
+
|
|
965
|
+
|
|
966
|
+
async def upload_prepare_handler(request: web.Request) -> web.Response:
|
|
967
|
+
try:
|
|
968
|
+
body = await _read_json_body(request)
|
|
969
|
+
if not isinstance(body, Mapping):
|
|
970
|
+
raise ValueError("Upload prepare body must be a JSON object.")
|
|
971
|
+
record = _prepare_upload_slot(_state(request), body)
|
|
972
|
+
return web.json_response(
|
|
973
|
+
{
|
|
974
|
+
"kind": "upload_slot",
|
|
975
|
+
"upload_id": record.upload_id,
|
|
976
|
+
"ref": record.upload_id,
|
|
977
|
+
"upload_url": f"/uploads/{record.upload_id}?token={record.token}",
|
|
978
|
+
"method": "PUT",
|
|
979
|
+
"headers": {"content-type": record.content_type},
|
|
980
|
+
"destination": "input",
|
|
981
|
+
"max_size": _max_request_bytes(_state(request)),
|
|
982
|
+
"file_ref": record.public_ref(),
|
|
983
|
+
}
|
|
984
|
+
)
|
|
985
|
+
except ValueError as exc:
|
|
986
|
+
return web.json_response({"error": "bad_request", "message": str(exc)}, status=400)
|
|
987
|
+
|
|
988
|
+
|
|
989
|
+
async def upload_put_handler(request: web.Request) -> web.Response:
|
|
990
|
+
state = _state(request)
|
|
991
|
+
upload_id = request.match_info["upload_id"]
|
|
992
|
+
record = state.uploads.get(upload_id)
|
|
993
|
+
if record is None:
|
|
994
|
+
return web.json_response({"error": "not_found", "message": "Unknown upload id."}, status=404)
|
|
995
|
+
if not secrets.compare_digest(request.query.get("token", ""), record.token):
|
|
996
|
+
return web.json_response({"error": "forbidden", "message": "Upload token is invalid."}, status=403)
|
|
997
|
+
|
|
998
|
+
content_length = request.content_length
|
|
999
|
+
max_bytes = _max_request_bytes(state)
|
|
1000
|
+
if content_length is not None and content_length > max_bytes:
|
|
1001
|
+
return _upload_too_large_response(max_bytes)
|
|
1002
|
+
|
|
1003
|
+
record.path.parent.mkdir(parents=True, exist_ok=True)
|
|
1004
|
+
temp_path = record.path.with_name(f".{record.path.name}.tmp")
|
|
1005
|
+
bytes_written = 0
|
|
1006
|
+
try:
|
|
1007
|
+
with temp_path.open("wb") as handle:
|
|
1008
|
+
try:
|
|
1009
|
+
async for chunk in request.content.iter_chunked(1024 * 1024):
|
|
1010
|
+
bytes_written += len(chunk)
|
|
1011
|
+
if bytes_written > max_bytes:
|
|
1012
|
+
handle.close()
|
|
1013
|
+
temp_path.unlink(missing_ok=True)
|
|
1014
|
+
return _upload_too_large_response(max_bytes)
|
|
1015
|
+
handle.write(chunk)
|
|
1016
|
+
except web.HTTPRequestEntityTooLarge:
|
|
1017
|
+
handle.close()
|
|
1018
|
+
temp_path.unlink(missing_ok=True)
|
|
1019
|
+
return _upload_too_large_response(max_bytes)
|
|
1020
|
+
temp_path.replace(record.path)
|
|
1021
|
+
finally:
|
|
1022
|
+
temp_path.unlink(missing_ok=True)
|
|
1023
|
+
|
|
1024
|
+
record.size = bytes_written
|
|
1025
|
+
record.status = "ready"
|
|
1026
|
+
return web.json_response({"status": "ready", "file_ref": record.public_ref()})
|
|
1027
|
+
|
|
1028
|
+
|
|
1029
|
+
async def upload_status_handler(request: web.Request) -> web.Response:
|
|
1030
|
+
record = _state(request).uploads.get(request.match_info["upload_id"])
|
|
1031
|
+
if record is None:
|
|
1032
|
+
return web.json_response({"error": "not_found", "message": "Unknown upload id."}, status=404)
|
|
1033
|
+
return web.json_response({"status": record.status, "file_ref": record.public_ref()})
|
|
1034
|
+
|
|
1035
|
+
|
|
1036
|
+
async def gallery_handler(request: web.Request) -> web.Response:
|
|
1037
|
+
state = _state(request)
|
|
1038
|
+
try:
|
|
1039
|
+
limit = _optional_positive_int_query(
|
|
1040
|
+
request,
|
|
1041
|
+
"limit",
|
|
1042
|
+
max_value=MAX_GALLERY_PAGE_LIMIT,
|
|
1043
|
+
)
|
|
1044
|
+
except ValueError as exc:
|
|
1045
|
+
session = _serve_session(request)
|
|
1046
|
+
return _json_response_for_session(
|
|
1047
|
+
{"error": "bad_request", "message": str(exc)},
|
|
1048
|
+
session,
|
|
1049
|
+
status=400,
|
|
1050
|
+
request=request,
|
|
1051
|
+
)
|
|
1052
|
+
|
|
1053
|
+
await _ensure_active_run_recovery(state)
|
|
1054
|
+
session = _serve_session(request)
|
|
1055
|
+
try:
|
|
1056
|
+
page = state.state_store.list_gallery_page(
|
|
1057
|
+
session.scope_key,
|
|
1058
|
+
limit=limit,
|
|
1059
|
+
cursor=request.query.get("cursor"),
|
|
1060
|
+
)
|
|
1061
|
+
except ValueError as exc:
|
|
1062
|
+
return _json_response_for_session(
|
|
1063
|
+
{"error": "bad_request", "message": str(exc)},
|
|
1064
|
+
session,
|
|
1065
|
+
status=400,
|
|
1066
|
+
request=request,
|
|
1067
|
+
)
|
|
1068
|
+
payload = {
|
|
1069
|
+
"state": state.config.state,
|
|
1070
|
+
"gallery": state.config.gallery,
|
|
1071
|
+
"session_id": session.session_id,
|
|
1072
|
+
"items": page.items,
|
|
1073
|
+
"next_cursor": page.next_cursor,
|
|
1074
|
+
"has_more": page.has_more,
|
|
1075
|
+
"limit": page.limit,
|
|
1076
|
+
}
|
|
1077
|
+
return _json_response_for_session(payload, session, request=request)
|
|
1078
|
+
|
|
1079
|
+
|
|
1080
|
+
async def gallery_delete_handler(request: web.Request) -> web.Response:
|
|
1081
|
+
session = _serve_session(request)
|
|
1082
|
+
deleted = _state(request).state_store.delete_gallery_item(
|
|
1083
|
+
session.scope_key,
|
|
1084
|
+
request.match_info["item_id"],
|
|
1085
|
+
)
|
|
1086
|
+
status = 200 if deleted else 404
|
|
1087
|
+
return _json_response_for_session({"deleted": deleted}, session, status=status)
|
|
1088
|
+
|
|
1089
|
+
|
|
1090
|
+
async def runs_handler(request: web.Request) -> web.Response:
|
|
1091
|
+
await _ensure_active_run_recovery(_state(request))
|
|
1092
|
+
session = _serve_session(request)
|
|
1093
|
+
statuses = None
|
|
1094
|
+
if request.query.get("active") == "true":
|
|
1095
|
+
statuses = ACTIVE_RUN_STATUSES
|
|
1096
|
+
payload = {
|
|
1097
|
+
"state": _state(request).config.state,
|
|
1098
|
+
"session_id": session.session_id,
|
|
1099
|
+
"runs": _state(request).state_store.list_runs(session.scope_key, statuses=statuses),
|
|
1100
|
+
}
|
|
1101
|
+
return _json_response_for_session(payload, session, request=request)
|
|
1102
|
+
|
|
1103
|
+
|
|
1104
|
+
async def single_run_handler(request: web.Request) -> web.Response:
|
|
1105
|
+
await _ensure_active_run_recovery(_state(request))
|
|
1106
|
+
session = _serve_session(request)
|
|
1107
|
+
state = _state(request)
|
|
1108
|
+
run_id = request.match_info["run_id"]
|
|
1109
|
+
run = state.state_store.get_run(session.scope_key, run_id)
|
|
1110
|
+
if run is None:
|
|
1111
|
+
return _json_response_for_session(
|
|
1112
|
+
{"error": "not_found", "message": f"Run '{run_id}' was not found."},
|
|
1113
|
+
session,
|
|
1114
|
+
status=404,
|
|
1115
|
+
)
|
|
1116
|
+
payload = {
|
|
1117
|
+
"state": state.config.state,
|
|
1118
|
+
"session_id": session.session_id,
|
|
1119
|
+
"run": run,
|
|
1120
|
+
"output_slots": state.state_store.list_output_slots(session.scope_key, run_id),
|
|
1121
|
+
"gallery_items": state.state_store.list_gallery_items_for_run(session.scope_key, run_id),
|
|
1122
|
+
}
|
|
1123
|
+
return _json_response_for_session(payload, session, request=request)
|
|
1124
|
+
|
|
1125
|
+
|
|
1126
|
+
async def cancel_run_handler(request: web.Request) -> web.Response:
|
|
1127
|
+
session = _serve_session(request)
|
|
1128
|
+
state = _state(request)
|
|
1129
|
+
run_id = request.match_info["run_id"]
|
|
1130
|
+
run = state.state_store.get_run(session.scope_key, run_id)
|
|
1131
|
+
if run is None:
|
|
1132
|
+
return _json_response_for_session(
|
|
1133
|
+
{"error": "not_found", "message": f"Run '{run_id}' was not found."},
|
|
1134
|
+
session,
|
|
1135
|
+
status=404,
|
|
1136
|
+
)
|
|
1137
|
+
|
|
1138
|
+
run_status = str(run.get("status") or "")
|
|
1139
|
+
if run_status == "cancelled":
|
|
1140
|
+
return _json_response_for_session(_cancelled_run_payload(state, session, run_id), session, request=request)
|
|
1141
|
+
if run_status in TERMINAL_RUN_STATUSES:
|
|
1142
|
+
return _json_response_for_session(
|
|
1143
|
+
{
|
|
1144
|
+
"error": "run_not_cancellable",
|
|
1145
|
+
"message": f"Run '{run_id}' is already {run_status}.",
|
|
1146
|
+
"run": run,
|
|
1147
|
+
},
|
|
1148
|
+
session,
|
|
1149
|
+
status=409,
|
|
1150
|
+
)
|
|
1151
|
+
|
|
1152
|
+
prompt_id = run.get("prompt_id")
|
|
1153
|
+
cancel_warning: dict[str, Any] | None = None
|
|
1154
|
+
if isinstance(prompt_id, str) and prompt_id:
|
|
1155
|
+
try:
|
|
1156
|
+
await asyncio.wait_for(state.executor.cancel(prompt_id), timeout=10)
|
|
1157
|
+
except ComfyUIRequestError as exc:
|
|
1158
|
+
cancel_warning = {
|
|
1159
|
+
"error": "comfyui_rejected_cancel",
|
|
1160
|
+
"message": str(exc),
|
|
1161
|
+
"comfy_status": exc.status,
|
|
1162
|
+
"comfy_url": exc.url,
|
|
1163
|
+
"comfyui": exc.payload,
|
|
1164
|
+
}
|
|
1165
|
+
except (aiohttp.ClientError, asyncio.TimeoutError) as exc:
|
|
1166
|
+
cancel_warning = {
|
|
1167
|
+
"error": "comfyui_unavailable",
|
|
1168
|
+
"message": str(exc),
|
|
1169
|
+
"comfy_url": state.config.comfy_url,
|
|
1170
|
+
}
|
|
1171
|
+
else:
|
|
1172
|
+
cancel_warning = {
|
|
1173
|
+
"error": "prompt_id_missing",
|
|
1174
|
+
"message": f"Run '{run_id}' did not have a remote prompt id when local cancellation was requested.",
|
|
1175
|
+
}
|
|
1176
|
+
message = "Generation cancelled."
|
|
1177
|
+
raw_result = {
|
|
1178
|
+
"status": "cancelled",
|
|
1179
|
+
"run_id": run_id,
|
|
1180
|
+
"prompt_id": prompt_id,
|
|
1181
|
+
"message": message,
|
|
1182
|
+
}
|
|
1183
|
+
if cancel_warning is not None:
|
|
1184
|
+
raw_result["remote_cancel"] = cancel_warning
|
|
1185
|
+
if not state.state_store.cancel_run(session.scope_key, run_id, raw_result=raw_result, error=message):
|
|
1186
|
+
return _json_response_for_session(
|
|
1187
|
+
{
|
|
1188
|
+
"error": "run_not_cancellable",
|
|
1189
|
+
"message": f"Run '{run_id}' is no longer active.",
|
|
1190
|
+
"run": state.state_store.get_run(session.scope_key, run_id),
|
|
1191
|
+
},
|
|
1192
|
+
session,
|
|
1193
|
+
status=409,
|
|
1194
|
+
)
|
|
1195
|
+
|
|
1196
|
+
task = state.active_run_tasks.pop(run_id, None)
|
|
1197
|
+
if task is not None and not task.done():
|
|
1198
|
+
task.cancel()
|
|
1199
|
+
|
|
1200
|
+
payload = _cancelled_run_payload(state, session, run_id)
|
|
1201
|
+
if cancel_warning is not None:
|
|
1202
|
+
payload["remote_cancel"] = cancel_warning
|
|
1203
|
+
return _json_response_for_session(payload, session)
|
|
1204
|
+
|
|
1205
|
+
|
|
1206
|
+
async def worker_callback_handler(request: web.Request) -> web.Response:
|
|
1207
|
+
auth_response = _worker_callback_auth_response(request)
|
|
1208
|
+
if auth_response is not None:
|
|
1209
|
+
return auth_response
|
|
1210
|
+
state = _state(request)
|
|
1211
|
+
route_run_id = request.match_info["run_id"]
|
|
1212
|
+
try:
|
|
1213
|
+
payload, uploads = await _read_worker_callback_request(request)
|
|
1214
|
+
except ValueError as exc:
|
|
1215
|
+
return web.json_response({"error": "bad_request", "message": str(exc)}, status=400)
|
|
1216
|
+
|
|
1217
|
+
payload_run_id = payload.get("run_id")
|
|
1218
|
+
if payload_run_id is not None and payload_run_id != route_run_id:
|
|
1219
|
+
return web.json_response(
|
|
1220
|
+
{"error": "bad_request", "message": "Callback run_id does not match route run_id."},
|
|
1221
|
+
status=400,
|
|
1222
|
+
)
|
|
1223
|
+
run = state.state_store.get_run_record(route_run_id)
|
|
1224
|
+
if run is None:
|
|
1225
|
+
return web.json_response({"error": "not_found", "message": "Unknown coordinator run."}, status=404)
|
|
1226
|
+
|
|
1227
|
+
status = str(payload.get("status") or "").lower()
|
|
1228
|
+
if run.status in TERMINAL_RUN_STATUSES and status in TERMINAL_RUN_STATUSES:
|
|
1229
|
+
session = ServeSession(session_id=run.session_id, scope_key=run.scope_key)
|
|
1230
|
+
return web.json_response({"status": run.status, "duplicate": True, **_callback_run_snapshot(state, session, run.run_id)})
|
|
1231
|
+
if status == "running":
|
|
1232
|
+
slots = _output_slots_for_callback_run(state, run)
|
|
1233
|
+
_record_worker_running_callback(state, run, payload, slots)
|
|
1234
|
+
session = ServeSession(session_id=run.session_id, scope_key=run.scope_key)
|
|
1235
|
+
return web.json_response({"status": "running", **_callback_run_snapshot(state, session, run.run_id)})
|
|
1236
|
+
if status == "completed":
|
|
1237
|
+
outputs = payload.get("outputs")
|
|
1238
|
+
if not isinstance(outputs, list):
|
|
1239
|
+
return web.json_response(
|
|
1240
|
+
{"error": "bad_request", "message": "Completed callback payload must include outputs."},
|
|
1241
|
+
status=400,
|
|
1242
|
+
)
|
|
1243
|
+
slots = _output_slots_for_callback_run(state, run)
|
|
1244
|
+
output_payloads = _localize_worker_callback_outputs(
|
|
1245
|
+
state,
|
|
1246
|
+
str(payload.get("prompt_id") or run.prompt_id or run.run_id),
|
|
1247
|
+
[dict(output) for output in outputs if isinstance(output, Mapping)],
|
|
1248
|
+
uploads,
|
|
1249
|
+
)
|
|
1250
|
+
response = {
|
|
1251
|
+
"status": "completed",
|
|
1252
|
+
"run_id": run.run_id,
|
|
1253
|
+
"prompt_id": str(payload.get("prompt_id") or run.prompt_id or ""),
|
|
1254
|
+
"issues": payload.get("issues") if isinstance(payload.get("issues"), list) else [],
|
|
1255
|
+
"outputs": output_payloads,
|
|
1256
|
+
}
|
|
1257
|
+
session = ServeSession(session_id=run.session_id, scope_key=run.scope_key)
|
|
1258
|
+
_record_completed_run_response(
|
|
1259
|
+
state,
|
|
1260
|
+
session,
|
|
1261
|
+
workflow_name=run.workflow,
|
|
1262
|
+
contract_name=run.contract,
|
|
1263
|
+
inputs=run.inputs,
|
|
1264
|
+
response=response,
|
|
1265
|
+
output_slots=slots,
|
|
1266
|
+
created_at=run.created_at,
|
|
1267
|
+
)
|
|
1268
|
+
return web.json_response({"status": "completed", **_callback_run_snapshot(state, session, run.run_id)})
|
|
1269
|
+
|
|
1270
|
+
if status in {"error", "failed", "timeout", "cancelled"}:
|
|
1271
|
+
slots = _output_slots_for_callback_run(state, run)
|
|
1272
|
+
error_payload = dict(payload)
|
|
1273
|
+
error_payload.setdefault("status", "error" if status != "cancelled" else "cancelled")
|
|
1274
|
+
error_payload.setdefault("run_id", run.run_id)
|
|
1275
|
+
error_payload.setdefault("prompt_id", run.prompt_id)
|
|
1276
|
+
if status == "cancelled":
|
|
1277
|
+
state.state_store.cancel_run(
|
|
1278
|
+
run.scope_key,
|
|
1279
|
+
run.run_id,
|
|
1280
|
+
raw_result=error_payload,
|
|
1281
|
+
error=str(error_payload.get("message") or "Generation cancelled."),
|
|
1282
|
+
)
|
|
1283
|
+
else:
|
|
1284
|
+
_record_failed_run(
|
|
1285
|
+
state,
|
|
1286
|
+
ServeSession(session_id=run.session_id, scope_key=run.scope_key),
|
|
1287
|
+
run.workflow,
|
|
1288
|
+
run.contract,
|
|
1289
|
+
{"inputs": run.inputs},
|
|
1290
|
+
error_payload,
|
|
1291
|
+
run_id=run.run_id,
|
|
1292
|
+
prompt_id=str(error_payload.get("prompt_id") or run.prompt_id or "") or None,
|
|
1293
|
+
output_slots=slots,
|
|
1294
|
+
created_at=run.created_at,
|
|
1295
|
+
)
|
|
1296
|
+
session = ServeSession(session_id=run.session_id, scope_key=run.scope_key)
|
|
1297
|
+
return web.json_response({"status": error_payload["status"], **_callback_run_snapshot(state, session, run.run_id)})
|
|
1298
|
+
|
|
1299
|
+
return web.json_response(
|
|
1300
|
+
{"error": "bad_request", "message": "Callback status must be running, completed, error, failed, timeout, or cancelled."},
|
|
1301
|
+
status=400,
|
|
1302
|
+
)
|
|
1303
|
+
|
|
1304
|
+
|
|
1305
|
+
def _cancelled_run_payload(state: ServeState, session: ServeSession, run_id: str) -> dict[str, Any]:
|
|
1306
|
+
return {
|
|
1307
|
+
"status": "cancelled",
|
|
1308
|
+
"run_id": run_id,
|
|
1309
|
+
"run": state.state_store.get_run(session.scope_key, run_id),
|
|
1310
|
+
"output_slots": state.state_store.list_output_slots(session.scope_key, run_id),
|
|
1311
|
+
"gallery_items": state.state_store.list_gallery_items_for_run(session.scope_key, run_id),
|
|
1312
|
+
}
|
|
1313
|
+
|
|
1314
|
+
|
|
1315
|
+
def _worker_callback_url(state: ServeState, run_id: str, *, wait: bool) -> str | None:
|
|
1316
|
+
if wait or state.config.executor != "proxy":
|
|
1317
|
+
return None
|
|
1318
|
+
base_url = getattr(state.config, "callback_url", None)
|
|
1319
|
+
if not base_url:
|
|
1320
|
+
return None
|
|
1321
|
+
return f"{base_url.rstrip('/')}/worker-callback/runs/{run_id}"
|
|
1322
|
+
|
|
1323
|
+
|
|
1324
|
+
def _callback_auth_token(state: ServeState) -> str | None:
|
|
1325
|
+
token = getattr(state.config, "callback_token", None) or getattr(state.config, "proxy_token", None)
|
|
1326
|
+
return str(token) if token else None
|
|
1327
|
+
|
|
1328
|
+
|
|
1329
|
+
def _worker_callback_auth_response(request: web.Request) -> web.Response | None:
|
|
1330
|
+
token = _callback_auth_token(_state(request))
|
|
1331
|
+
if not token:
|
|
1332
|
+
return None
|
|
1333
|
+
expected = f"Bearer {token}"
|
|
1334
|
+
received = request.headers.get(PROXY_AUTH_HEADER, "")
|
|
1335
|
+
if secrets.compare_digest(received, expected):
|
|
1336
|
+
return None
|
|
1337
|
+
return web.json_response({"error": "forbidden", "message": "Worker callback token is invalid."}, status=403)
|
|
1338
|
+
|
|
1339
|
+
|
|
1340
|
+
async def _read_worker_callback_request(
|
|
1341
|
+
request: web.Request,
|
|
1342
|
+
) -> tuple[dict[str, Any], dict[str, WorkerCallbackUpload]]:
|
|
1343
|
+
if not request.content_type.lower().startswith("multipart/"):
|
|
1344
|
+
return await _read_json_body(request), {}
|
|
1345
|
+
|
|
1346
|
+
reader = await request.multipart()
|
|
1347
|
+
payload: dict[str, Any] | None = None
|
|
1348
|
+
uploads: dict[str, WorkerCallbackUpload] = {}
|
|
1349
|
+
while raw_part := await reader.next():
|
|
1350
|
+
if not isinstance(raw_part, aiohttp.BodyPartReader):
|
|
1351
|
+
raise ValueError("Nested multipart callback uploads are not supported.")
|
|
1352
|
+
part = raw_part
|
|
1353
|
+
if part.name == "payload":
|
|
1354
|
+
try:
|
|
1355
|
+
payload_data = json.loads(await part.text())
|
|
1356
|
+
except json.JSONDecodeError as exc:
|
|
1357
|
+
raise ValueError(f"Invalid callback payload JSON: {exc}") from exc
|
|
1358
|
+
if not isinstance(payload_data, dict):
|
|
1359
|
+
raise ValueError("Callback payload must be a JSON object.")
|
|
1360
|
+
payload = payload_data
|
|
1361
|
+
continue
|
|
1362
|
+
field_name = str(part.name or "")
|
|
1363
|
+
if not field_name:
|
|
1364
|
+
raise ValueError("Callback upload part is missing a field name.")
|
|
1365
|
+
uploads[field_name] = WorkerCallbackUpload(
|
|
1366
|
+
field_name=field_name,
|
|
1367
|
+
filename=_safe_artifact_filename(part.filename, fallback=f"{field_name}.bin"),
|
|
1368
|
+
content_type=part.headers.get("Content-Type") or DEFAULT_UPLOAD_CONTENT_TYPE,
|
|
1369
|
+
body=await part.read(),
|
|
1370
|
+
)
|
|
1371
|
+
|
|
1372
|
+
if payload is None:
|
|
1373
|
+
raise ValueError("Worker callback request is missing a payload field.")
|
|
1374
|
+
return payload, uploads
|
|
1375
|
+
|
|
1376
|
+
|
|
1377
|
+
def _callback_run_snapshot(state: ServeState, session: ServeSession, run_id: str) -> dict[str, Any]:
|
|
1378
|
+
return {
|
|
1379
|
+
"run_id": run_id,
|
|
1380
|
+
"run": state.state_store.get_run(session.scope_key, run_id),
|
|
1381
|
+
"output_slots": state.state_store.list_output_slots(session.scope_key, run_id),
|
|
1382
|
+
"gallery_items": state.state_store.list_gallery_items_for_run(session.scope_key, run_id),
|
|
1383
|
+
}
|
|
1384
|
+
|
|
1385
|
+
|
|
1386
|
+
def _output_slots_for_callback_run(state: ServeState, run: ServeRunRecord) -> list[ServeRunOutputSlot]:
|
|
1387
|
+
slots = [_output_slot_from_public_dict(run, slot) for slot in state.state_store.list_output_slots(run.scope_key, run.run_id)]
|
|
1388
|
+
if slots:
|
|
1389
|
+
return slots
|
|
1390
|
+
outputs = _contract_outputs_for_run(state, run.workflow, run.contract) or ()
|
|
1391
|
+
return _output_slots_for_run(
|
|
1392
|
+
run_id=run.run_id,
|
|
1393
|
+
session=ServeSession(session_id=run.session_id, scope_key=run.scope_key),
|
|
1394
|
+
workflow_name=run.workflow,
|
|
1395
|
+
contract_name=run.contract,
|
|
1396
|
+
outputs=outputs,
|
|
1397
|
+
prompt_id=run.prompt_id,
|
|
1398
|
+
inputs=run.inputs,
|
|
1399
|
+
)
|
|
1400
|
+
|
|
1401
|
+
|
|
1402
|
+
def _output_slot_from_public_dict(run: ServeRunRecord, payload: Mapping[str, Any]) -> ServeRunOutputSlot:
|
|
1403
|
+
raw_result = payload.get("rawResult")
|
|
1404
|
+
return ServeRunOutputSlot(
|
|
1405
|
+
slot_id=str(payload.get("slot_id") or payload.get("slotId") or f"slot_{run.run_id}_0_result"),
|
|
1406
|
+
run_id=run.run_id,
|
|
1407
|
+
session_id=run.session_id,
|
|
1408
|
+
scope_key=run.scope_key,
|
|
1409
|
+
workflow=run.workflow,
|
|
1410
|
+
contract=run.contract,
|
|
1411
|
+
output_name=str(payload.get("outputName") or "result"),
|
|
1412
|
+
output_type=_slot_output_type(str(payload.get("type") or "json")),
|
|
1413
|
+
status=str(payload.get("status") or "pending"),
|
|
1414
|
+
prompt_id=str(payload.get("promptId") or run.prompt_id or "") or None,
|
|
1415
|
+
width=_optional_positive_int(payload.get("width")),
|
|
1416
|
+
height=_optional_positive_int(payload.get("height")),
|
|
1417
|
+
error=str(payload.get("error")) if payload.get("error") is not None else None,
|
|
1418
|
+
raw_result={str(key): value for key, value in raw_result.items()} if isinstance(raw_result, Mapping) else None,
|
|
1419
|
+
created_at=str(payload.get("createdAt") or run.created_at),
|
|
1420
|
+
updated_at=str(payload.get("updatedAt") or run.updated_at),
|
|
1421
|
+
)
|
|
1422
|
+
|
|
1423
|
+
|
|
1424
|
+
def _record_worker_running_callback(
|
|
1425
|
+
state: ServeState,
|
|
1426
|
+
run: ServeRunRecord,
|
|
1427
|
+
payload: Mapping[str, Any],
|
|
1428
|
+
slots: list[ServeRunOutputSlot],
|
|
1429
|
+
) -> None:
|
|
1430
|
+
prompt_id = str(payload.get("prompt_id") or run.prompt_id or "") or None
|
|
1431
|
+
raw_result = {"status": "running", "run_id": run.run_id, **dict(payload)}
|
|
1432
|
+
state.state_store.record_run(
|
|
1433
|
+
ServeRunRecord(
|
|
1434
|
+
run_id=run.run_id,
|
|
1435
|
+
session_id=run.session_id,
|
|
1436
|
+
scope_key=run.scope_key,
|
|
1437
|
+
workflow=run.workflow,
|
|
1438
|
+
contract=run.contract,
|
|
1439
|
+
status="running",
|
|
1440
|
+
inputs=run.inputs,
|
|
1441
|
+
prompt_id=prompt_id,
|
|
1442
|
+
raw_result=raw_result,
|
|
1443
|
+
created_at=run.created_at,
|
|
1444
|
+
)
|
|
1445
|
+
)
|
|
1446
|
+
state.state_store.record_output_slots(
|
|
1447
|
+
[_copy_output_slot(slot, status="running", prompt_id=prompt_id, raw_result=raw_result) for slot in slots]
|
|
1448
|
+
)
|
|
1449
|
+
|
|
1450
|
+
|
|
1451
|
+
def _localize_worker_callback_outputs(
|
|
1452
|
+
state: ServeState,
|
|
1453
|
+
prompt_id: str,
|
|
1454
|
+
outputs: list[dict[str, Any]],
|
|
1455
|
+
uploads: Mapping[str, WorkerCallbackUpload],
|
|
1456
|
+
) -> list[dict[str, Any]]:
|
|
1457
|
+
output_payloads = [dict(output) for output in outputs]
|
|
1458
|
+
for output_index, output in enumerate(output_payloads):
|
|
1459
|
+
artifacts = output.get("artifacts")
|
|
1460
|
+
if not isinstance(artifacts, list):
|
|
1461
|
+
continue
|
|
1462
|
+
for artifact_index, artifact in enumerate(artifacts):
|
|
1463
|
+
if not isinstance(artifact, dict):
|
|
1464
|
+
continue
|
|
1465
|
+
artifact = cast(dict[str, Any], artifact)
|
|
1466
|
+
field_name = str(artifact.get("upload_field") or f"artifact_{output_index}_{artifact_index}")
|
|
1467
|
+
upload = uploads.get(field_name)
|
|
1468
|
+
if upload is None:
|
|
1469
|
+
continue
|
|
1470
|
+
filename = _safe_artifact_filename(
|
|
1471
|
+
artifact.get("filename") or upload.filename,
|
|
1472
|
+
fallback=upload.filename,
|
|
1473
|
+
)
|
|
1474
|
+
ref = _write_local_artifact(state.artifact_dir, prompt_id, filename, upload.body)
|
|
1475
|
+
artifact["serve_artifact"] = ref
|
|
1476
|
+
artifact["url"] = f"/outputs/view?serve_artifact={quote(ref, safe='/')}"
|
|
1477
|
+
artifact["content_type"] = upload.content_type
|
|
1478
|
+
artifact.pop("upload_field", None)
|
|
1479
|
+
return output_payloads
|
|
1480
|
+
|
|
1481
|
+
|
|
1482
|
+
async def run_contract_handler(request: web.Request) -> web.Response:
|
|
1483
|
+
body: dict[str, Any] = {}
|
|
1484
|
+
session = _serve_session(request)
|
|
1485
|
+
try:
|
|
1486
|
+
body = await _read_json_body(request)
|
|
1487
|
+
payload = await _run_contract(
|
|
1488
|
+
_state(request),
|
|
1489
|
+
session,
|
|
1490
|
+
request.match_info["workflow"],
|
|
1491
|
+
request.match_info["contract"],
|
|
1492
|
+
body,
|
|
1493
|
+
)
|
|
1494
|
+
status = 400 if payload.get("status") == "invalid_request" else 200
|
|
1495
|
+
return _json_response_for_session(payload, session, status=status, request=request)
|
|
1496
|
+
except web.HTTPRequestEntityTooLarge:
|
|
1497
|
+
state = _state(request)
|
|
1498
|
+
max_mib = _max_request_bytes(state) // (1024 * 1024)
|
|
1499
|
+
return _json_response_for_session(
|
|
1500
|
+
{
|
|
1501
|
+
"error": "request_too_large",
|
|
1502
|
+
"message": (
|
|
1503
|
+
f"Request body is too large. This cg serve instance accepts "
|
|
1504
|
+
f"contract requests up to {max_mib} MiB."
|
|
1505
|
+
),
|
|
1506
|
+
},
|
|
1507
|
+
session,
|
|
1508
|
+
status=413,
|
|
1509
|
+
)
|
|
1510
|
+
except ComfyGitServeTimeoutError as exc:
|
|
1511
|
+
payload = {"error": "timeout", "message": str(exc)}
|
|
1512
|
+
payload.update(_record_failed_run(_state(request), session, request.match_info["workflow"], request.match_info["contract"], body, payload))
|
|
1513
|
+
return _json_response_for_session(payload, session, status=504)
|
|
1514
|
+
except ComfyUIRequestError as exc:
|
|
1515
|
+
payload = {
|
|
1516
|
+
"error": "comfyui_rejected_request",
|
|
1517
|
+
"message": str(exc),
|
|
1518
|
+
"comfy_status": exc.status,
|
|
1519
|
+
"comfy_url": exc.url,
|
|
1520
|
+
"comfyui": exc.payload,
|
|
1521
|
+
}
|
|
1522
|
+
payload.update(_record_failed_run(_state(request), session, request.match_info["workflow"], request.match_info["contract"], body, payload))
|
|
1523
|
+
return _json_response_for_session(payload, session, status=400 if exc.status == 400 else 502)
|
|
1524
|
+
except ComfyUIExecutionError as exc:
|
|
1525
|
+
payload = {
|
|
1526
|
+
"error": "comfyui_execution_failed",
|
|
1527
|
+
"message": str(exc),
|
|
1528
|
+
"prompt_id": exc.prompt_id,
|
|
1529
|
+
"comfyui": exc.payload,
|
|
1530
|
+
}
|
|
1531
|
+
payload.update(_record_failed_run(_state(request), session, request.match_info["workflow"], request.match_info["contract"], body, payload))
|
|
1532
|
+
return _json_response_for_session(payload, session, status=500)
|
|
1533
|
+
except (aiohttp.ClientError, asyncio.TimeoutError) as exc:
|
|
1534
|
+
state = _state(request)
|
|
1535
|
+
payload = _executor_unavailable_payload(state, exc)
|
|
1536
|
+
payload.update(_record_failed_run(state, session, request.match_info["workflow"], request.match_info["contract"], body, payload))
|
|
1537
|
+
return _json_response_for_session(payload, session, status=502)
|
|
1538
|
+
except ValueError as exc:
|
|
1539
|
+
payload = {"error": "bad_request", "message": str(exc)}
|
|
1540
|
+
payload.update(_record_failed_run(_state(request), session, request.match_info["workflow"], request.match_info["contract"], body, payload))
|
|
1541
|
+
return _json_response_for_session(payload, session, status=400)
|
|
1542
|
+
except Exception as exc:
|
|
1543
|
+
payload = {"error": "internal_error", "message": str(exc)}
|
|
1544
|
+
payload.update(_record_failed_run(_state(request), session, request.match_info["workflow"], request.match_info["contract"], body, payload))
|
|
1545
|
+
return _json_response_for_session(payload, session, status=500)
|
|
1546
|
+
|
|
1547
|
+
|
|
1548
|
+
async def output_view_handler(request: web.Request) -> web.StreamResponse:
|
|
1549
|
+
serve_artifact = request.query.get("serve_artifact")
|
|
1550
|
+
if serve_artifact:
|
|
1551
|
+
return _serve_local_artifact_response(_state(request), serve_artifact)
|
|
1552
|
+
|
|
1553
|
+
filename = request.query.get("filename")
|
|
1554
|
+
if not filename:
|
|
1555
|
+
return web.json_response(
|
|
1556
|
+
{"error": "bad_request", "message": "'filename' query parameter is required."},
|
|
1557
|
+
status=400,
|
|
1558
|
+
)
|
|
1559
|
+
params = {
|
|
1560
|
+
"filename": filename,
|
|
1561
|
+
"subfolder": request.query.get("subfolder", ""),
|
|
1562
|
+
"type": request.query.get("type", "output"),
|
|
1563
|
+
}
|
|
1564
|
+
try:
|
|
1565
|
+
output_response = await _state(request).client.fetch_output(
|
|
1566
|
+
params,
|
|
1567
|
+
request_headers=_output_request_headers(request),
|
|
1568
|
+
)
|
|
1569
|
+
except aiohttp.ClientResponseError as exc:
|
|
1570
|
+
if exc.status == 404:
|
|
1571
|
+
return web.json_response(
|
|
1572
|
+
{
|
|
1573
|
+
"error": "output_not_found",
|
|
1574
|
+
"message": "ComfyUI no longer has this output artifact.",
|
|
1575
|
+
"filename": filename,
|
|
1576
|
+
"type": params["type"],
|
|
1577
|
+
},
|
|
1578
|
+
status=404,
|
|
1579
|
+
)
|
|
1580
|
+
state = _state(request)
|
|
1581
|
+
return web.json_response(
|
|
1582
|
+
{
|
|
1583
|
+
"error": "comfyui_unavailable",
|
|
1584
|
+
"message": str(exc),
|
|
1585
|
+
"comfy_url": state.config.comfy_url,
|
|
1586
|
+
},
|
|
1587
|
+
status=502,
|
|
1588
|
+
)
|
|
1589
|
+
except (aiohttp.ClientError, asyncio.TimeoutError) as exc:
|
|
1590
|
+
state = _state(request)
|
|
1591
|
+
return web.json_response(
|
|
1592
|
+
{
|
|
1593
|
+
"error": "comfyui_unavailable",
|
|
1594
|
+
"message": str(exc),
|
|
1595
|
+
"comfy_url": state.config.comfy_url,
|
|
1596
|
+
},
|
|
1597
|
+
status=502,
|
|
1598
|
+
)
|
|
1599
|
+
headers = {"Content-Type": output_response.content_type}
|
|
1600
|
+
if output_response.disposition:
|
|
1601
|
+
headers["Content-Disposition"] = output_response.disposition
|
|
1602
|
+
headers.update(output_response.headers)
|
|
1603
|
+
return web.Response(body=output_response.body, status=output_response.status, headers=headers)
|
|
1604
|
+
|
|
1605
|
+
|
|
1606
|
+
def _output_request_headers(request: web.Request) -> dict[str, str]:
|
|
1607
|
+
return {
|
|
1608
|
+
header_name: request.headers[header_name]
|
|
1609
|
+
for header_name in OUTPUT_REQUEST_HEADERS
|
|
1610
|
+
if header_name in request.headers
|
|
1611
|
+
}
|
|
1612
|
+
|
|
1613
|
+
|
|
1614
|
+
def _serve_local_artifact_response(state: ServeState, artifact_ref: str) -> web.StreamResponse:
|
|
1615
|
+
path = _local_artifact_path(state, artifact_ref)
|
|
1616
|
+
if path is None or not path.is_file():
|
|
1617
|
+
return web.json_response({"error": "not_found", "message": "Unknown serve artifact."}, status=404)
|
|
1618
|
+
return web.FileResponse(path)
|
|
1619
|
+
|
|
1620
|
+
|
|
1621
|
+
def _local_artifact_path(state: ServeState, artifact_ref: str) -> Path | None:
|
|
1622
|
+
ref_path = Path(artifact_ref)
|
|
1623
|
+
if ref_path.is_absolute() or ".." in ref_path.parts:
|
|
1624
|
+
return None
|
|
1625
|
+
if len(ref_path.parts) != 2:
|
|
1626
|
+
return None
|
|
1627
|
+
if any(_safe_token(part) != part for part in ref_path.parts):
|
|
1628
|
+
filename = ref_path.parts[-1]
|
|
1629
|
+
if _safe_upload_filename(filename) != filename:
|
|
1630
|
+
return None
|
|
1631
|
+
resolved = (state.artifact_dir / ref_path).resolve()
|
|
1632
|
+
artifact_root = state.artifact_dir.resolve()
|
|
1633
|
+
if artifact_root not in resolved.parents:
|
|
1634
|
+
return None
|
|
1635
|
+
return resolved
|
|
1636
|
+
|
|
1637
|
+
|
|
1638
|
+
async def _read_json_body(request: web.Request) -> dict[str, Any]:
|
|
1639
|
+
if request.can_read_body is False:
|
|
1640
|
+
return {}
|
|
1641
|
+
try:
|
|
1642
|
+
data = await request.json()
|
|
1643
|
+
except ValueError as exc:
|
|
1644
|
+
raise ValueError(f"Invalid JSON body: {exc}") from exc
|
|
1645
|
+
if data is None:
|
|
1646
|
+
return {}
|
|
1647
|
+
if not isinstance(data, dict):
|
|
1648
|
+
raise ValueError("Request body must be a JSON object.")
|
|
1649
|
+
return data
|
|
1650
|
+
|
|
1651
|
+
|
|
1652
|
+
async def _read_proxy_run_request(request: web.Request, state: ServeState) -> dict[str, Any]:
|
|
1653
|
+
if not request.content_type.lower().startswith("multipart/"):
|
|
1654
|
+
return await _read_json_body(request)
|
|
1655
|
+
|
|
1656
|
+
reader = await request.multipart()
|
|
1657
|
+
payload: dict[str, Any] | None = None
|
|
1658
|
+
staged_fields: set[str] = set()
|
|
1659
|
+
while raw_part := await reader.next():
|
|
1660
|
+
if not isinstance(raw_part, aiohttp.BodyPartReader):
|
|
1661
|
+
raise ValueError("Nested multipart proxy uploads are not supported.")
|
|
1662
|
+
part = raw_part
|
|
1663
|
+
if part.name == "payload":
|
|
1664
|
+
try:
|
|
1665
|
+
payload_data = json.loads(await part.text())
|
|
1666
|
+
except json.JSONDecodeError as exc:
|
|
1667
|
+
raise ValueError(f"Invalid proxy payload JSON: {exc}") from exc
|
|
1668
|
+
if not isinstance(payload_data, dict):
|
|
1669
|
+
raise ValueError("Proxy payload must be a JSON object.")
|
|
1670
|
+
payload = payload_data
|
|
1671
|
+
continue
|
|
1672
|
+
if payload is None:
|
|
1673
|
+
raise ValueError("Proxy multipart payload must be sent before file parts.")
|
|
1674
|
+
upload = _proxy_upload_by_field(payload, str(part.name or ""))
|
|
1675
|
+
if upload is None:
|
|
1676
|
+
raise ValueError(f"Unexpected proxy upload field '{part.name}'.")
|
|
1677
|
+
await _stage_proxy_upload_part(state, part, upload)
|
|
1678
|
+
staged_fields.add(str(part.name))
|
|
1679
|
+
|
|
1680
|
+
if payload is None:
|
|
1681
|
+
raise ValueError("Proxy run request is missing a payload field.")
|
|
1682
|
+
expected_fields = {
|
|
1683
|
+
str(upload.get("field_name") or "")
|
|
1684
|
+
for upload in payload.get("uploads", [])
|
|
1685
|
+
if isinstance(upload, Mapping)
|
|
1686
|
+
}
|
|
1687
|
+
missing_fields = {field for field in expected_fields if field and field not in staged_fields}
|
|
1688
|
+
if missing_fields:
|
|
1689
|
+
raise ValueError(f"Proxy run request is missing upload file part(s): {', '.join(sorted(missing_fields))}.")
|
|
1690
|
+
return payload
|
|
1691
|
+
|
|
1692
|
+
|
|
1693
|
+
def _proxy_upload_by_field(payload: Mapping[str, Any], field_name: str) -> Mapping[str, Any] | None:
|
|
1694
|
+
uploads = payload.get("uploads")
|
|
1695
|
+
if not isinstance(uploads, list):
|
|
1696
|
+
return None
|
|
1697
|
+
for upload in uploads:
|
|
1698
|
+
if isinstance(upload, Mapping) and upload.get("field_name") == field_name:
|
|
1699
|
+
return upload
|
|
1700
|
+
return None
|
|
1701
|
+
|
|
1702
|
+
|
|
1703
|
+
async def _stage_proxy_upload_part(
|
|
1704
|
+
state: ServeState,
|
|
1705
|
+
part: Any,
|
|
1706
|
+
upload: Mapping[str, Any],
|
|
1707
|
+
) -> None:
|
|
1708
|
+
filename = _safe_upload_filename(upload.get("comfyui_filename") or part.filename)
|
|
1709
|
+
expected_size = _optional_positive_int(upload.get("size"))
|
|
1710
|
+
max_bytes = _max_request_bytes(state)
|
|
1711
|
+
if expected_size is not None and expected_size > max_bytes:
|
|
1712
|
+
raise ValueError(f"Proxy upload '{filename}' is too large. Limit: {max_bytes} bytes.")
|
|
1713
|
+
|
|
1714
|
+
input_dir = _comfyui_input_dir(state.env)
|
|
1715
|
+
input_dir.mkdir(parents=True, exist_ok=True)
|
|
1716
|
+
target_path = input_dir / filename
|
|
1717
|
+
temp_path = target_path.with_name(f".{target_path.name}.tmp")
|
|
1718
|
+
bytes_written = 0
|
|
1719
|
+
try:
|
|
1720
|
+
with temp_path.open("wb") as handle:
|
|
1721
|
+
while chunk := await part.read_chunk(1024 * 1024):
|
|
1722
|
+
bytes_written += len(chunk)
|
|
1723
|
+
if bytes_written > max_bytes:
|
|
1724
|
+
handle.close()
|
|
1725
|
+
temp_path.unlink(missing_ok=True)
|
|
1726
|
+
raise ValueError(f"Proxy upload '{filename}' is too large. Limit: {max_bytes} bytes.")
|
|
1727
|
+
handle.write(chunk)
|
|
1728
|
+
temp_path.replace(target_path)
|
|
1729
|
+
finally:
|
|
1730
|
+
temp_path.unlink(missing_ok=True)
|
|
1731
|
+
|
|
1732
|
+
|
|
1733
|
+
def _track_proxy_task(state: ServeState, task: asyncio.Task[Any]) -> None:
|
|
1734
|
+
state.background_tasks.add(task)
|
|
1735
|
+
task.add_done_callback(state.background_tasks.discard)
|
|
1736
|
+
|
|
1737
|
+
|
|
1738
|
+
def _proxy_callback_target(value: Any) -> ProxyCallbackTarget | None:
|
|
1739
|
+
if not isinstance(value, Mapping):
|
|
1740
|
+
return None
|
|
1741
|
+
run_id = value.get("run_id")
|
|
1742
|
+
url = value.get("url")
|
|
1743
|
+
if not isinstance(run_id, str) or not run_id:
|
|
1744
|
+
raise ValueError("Proxy callback payload is missing a run_id.")
|
|
1745
|
+
if not isinstance(url, str) or not url:
|
|
1746
|
+
raise ValueError("Proxy callback payload is missing a url.")
|
|
1747
|
+
token = value.get("token")
|
|
1748
|
+
return ProxyCallbackTarget(
|
|
1749
|
+
run_id=run_id,
|
|
1750
|
+
url=url,
|
|
1751
|
+
token=str(token) if token else None,
|
|
1752
|
+
)
|
|
1753
|
+
|
|
1754
|
+
|
|
1755
|
+
async def _complete_proxy_runtime_run(
|
|
1756
|
+
state: ServeState,
|
|
1757
|
+
prompt_id: str,
|
|
1758
|
+
outputs: tuple[Any, ...],
|
|
1759
|
+
*,
|
|
1760
|
+
timeout_seconds: float,
|
|
1761
|
+
poll_interval_seconds: float,
|
|
1762
|
+
callback: ProxyCallbackTarget | None = None,
|
|
1763
|
+
) -> None:
|
|
1764
|
+
record = state.proxy_runs.get(prompt_id)
|
|
1765
|
+
if record is None:
|
|
1766
|
+
return
|
|
1767
|
+
record.status = "running"
|
|
1768
|
+
record.raw_result = {"status": "running", "prompt_id": prompt_id}
|
|
1769
|
+
await _post_worker_status_callback(callback, {"status": "running", "run_id": callback.run_id, "prompt_id": prompt_id} if callback else {})
|
|
1770
|
+
try:
|
|
1771
|
+
execution = await state.executor.complete_submitted(
|
|
1772
|
+
prompt_id,
|
|
1773
|
+
outputs,
|
|
1774
|
+
timeout_seconds=timeout_seconds,
|
|
1775
|
+
poll_interval_seconds=poll_interval_seconds,
|
|
1776
|
+
)
|
|
1777
|
+
except ComfyGitServeTimeoutError as exc:
|
|
1778
|
+
payload = {"error": "timeout", "message": str(exc), "prompt_id": prompt_id}
|
|
1779
|
+
_record_proxy_run_error(record, payload)
|
|
1780
|
+
await _post_worker_status_callback(callback, {"status": "error", "run_id": callback.run_id, **payload} if callback else {})
|
|
1781
|
+
return
|
|
1782
|
+
except ComfyUIRequestError as exc:
|
|
1783
|
+
payload = {
|
|
1784
|
+
"error": "comfyui_rejected_request",
|
|
1785
|
+
"message": str(exc),
|
|
1786
|
+
"comfy_status": exc.status,
|
|
1787
|
+
"comfy_url": exc.url,
|
|
1788
|
+
"comfyui": exc.payload,
|
|
1789
|
+
"prompt_id": prompt_id,
|
|
1790
|
+
}
|
|
1791
|
+
_record_proxy_run_error(record, payload)
|
|
1792
|
+
await _post_worker_status_callback(callback, {"status": "error", "run_id": callback.run_id, **payload} if callback else {})
|
|
1793
|
+
return
|
|
1794
|
+
except ComfyUIExecutionError as exc:
|
|
1795
|
+
payload = {
|
|
1796
|
+
"error": "comfyui_execution_failed",
|
|
1797
|
+
"message": str(exc),
|
|
1798
|
+
"prompt_id": exc.prompt_id,
|
|
1799
|
+
"comfyui": exc.payload,
|
|
1800
|
+
}
|
|
1801
|
+
_record_proxy_run_error(record, payload)
|
|
1802
|
+
await _post_worker_status_callback(callback, {"status": "error", "run_id": callback.run_id, **payload} if callback else {})
|
|
1803
|
+
return
|
|
1804
|
+
except (aiohttp.ClientError, asyncio.TimeoutError) as exc:
|
|
1805
|
+
payload = _executor_unavailable_payload(state, exc, prompt_id=prompt_id)
|
|
1806
|
+
_record_proxy_run_error(record, payload)
|
|
1807
|
+
await _post_worker_status_callback(callback, {"status": "error", "run_id": callback.run_id, **payload} if callback else {})
|
|
1808
|
+
return
|
|
1809
|
+
except Exception as exc:
|
|
1810
|
+
payload = {"error": "internal_error", "message": str(exc), "prompt_id": prompt_id}
|
|
1811
|
+
_record_proxy_run_error(record, payload)
|
|
1812
|
+
await _post_worker_status_callback(callback, {"status": "error", "run_id": callback.run_id, **payload} if callback else {})
|
|
1813
|
+
return
|
|
1814
|
+
|
|
1815
|
+
output_payloads = _register_proxy_artifacts(state, execution.outputs)
|
|
1816
|
+
record.status = "completed"
|
|
1817
|
+
record.raw_result = {
|
|
1818
|
+
"status": "completed",
|
|
1819
|
+
"prompt_id": execution.prompt_id,
|
|
1820
|
+
"outputs": output_payloads,
|
|
1821
|
+
}
|
|
1822
|
+
if callback is not None:
|
|
1823
|
+
try:
|
|
1824
|
+
callback_outputs, uploads = await _worker_callback_outputs_and_uploads(state, execution.outputs)
|
|
1825
|
+
await _post_worker_completion_callback(
|
|
1826
|
+
callback,
|
|
1827
|
+
{
|
|
1828
|
+
"status": "completed",
|
|
1829
|
+
"run_id": callback.run_id,
|
|
1830
|
+
"prompt_id": execution.prompt_id,
|
|
1831
|
+
"outputs": callback_outputs,
|
|
1832
|
+
},
|
|
1833
|
+
uploads,
|
|
1834
|
+
)
|
|
1835
|
+
except Exception as exc:
|
|
1836
|
+
_record_proxy_run_error(
|
|
1837
|
+
record,
|
|
1838
|
+
{
|
|
1839
|
+
"error": "callback_failed",
|
|
1840
|
+
"message": str(exc),
|
|
1841
|
+
"prompt_id": execution.prompt_id,
|
|
1842
|
+
},
|
|
1843
|
+
)
|
|
1844
|
+
|
|
1845
|
+
|
|
1846
|
+
async def _post_worker_status_callback(callback: ProxyCallbackTarget | None, payload: Mapping[str, Any]) -> None:
|
|
1847
|
+
if callback is None:
|
|
1848
|
+
return
|
|
1849
|
+
await _post_worker_callback_request(callback, json_payload=dict(payload))
|
|
1850
|
+
|
|
1851
|
+
|
|
1852
|
+
async def _post_worker_completion_callback(
|
|
1853
|
+
callback: ProxyCallbackTarget,
|
|
1854
|
+
payload: Mapping[str, Any],
|
|
1855
|
+
uploads: list[WorkerCallbackUpload],
|
|
1856
|
+
) -> None:
|
|
1857
|
+
headers = _worker_callback_headers(callback)
|
|
1858
|
+
if not uploads:
|
|
1859
|
+
await _post_worker_callback_request(callback, json_payload=dict(payload))
|
|
1860
|
+
return
|
|
1861
|
+
|
|
1862
|
+
def build_form() -> aiohttp.FormData:
|
|
1863
|
+
form = aiohttp.FormData()
|
|
1864
|
+
form.add_field("payload", json.dumps(dict(payload)), content_type="application/json")
|
|
1865
|
+
for upload in uploads:
|
|
1866
|
+
form.add_field(
|
|
1867
|
+
upload.field_name,
|
|
1868
|
+
upload.body,
|
|
1869
|
+
filename=upload.filename,
|
|
1870
|
+
content_type=upload.content_type,
|
|
1871
|
+
)
|
|
1872
|
+
return form
|
|
1873
|
+
|
|
1874
|
+
await _post_worker_callback_request(callback, data_factory=build_form, headers=headers, timeout_seconds=120)
|
|
1875
|
+
|
|
1876
|
+
|
|
1877
|
+
async def _post_worker_callback_request(
|
|
1878
|
+
callback: ProxyCallbackTarget,
|
|
1879
|
+
*,
|
|
1880
|
+
json_payload: dict[str, Any] | None = None,
|
|
1881
|
+
data_factory: Callable[[], aiohttp.FormData] | None = None,
|
|
1882
|
+
headers: Mapping[str, str] | None = None,
|
|
1883
|
+
timeout_seconds: float = 30,
|
|
1884
|
+
) -> None:
|
|
1885
|
+
request_headers = dict(headers or _worker_callback_headers(callback))
|
|
1886
|
+
for attempt in range(20):
|
|
1887
|
+
async with aiohttp.ClientSession() as session:
|
|
1888
|
+
async with session.post(
|
|
1889
|
+
callback.url,
|
|
1890
|
+
json=json_payload,
|
|
1891
|
+
data=data_factory() if data_factory is not None else None,
|
|
1892
|
+
headers=request_headers,
|
|
1893
|
+
timeout=aiohttp.ClientTimeout(total=timeout_seconds),
|
|
1894
|
+
) as response:
|
|
1895
|
+
if response.status not in {404, 409} or attempt == 19:
|
|
1896
|
+
response.raise_for_status()
|
|
1897
|
+
return
|
|
1898
|
+
await asyncio.sleep(0.25)
|
|
1899
|
+
|
|
1900
|
+
|
|
1901
|
+
def _worker_callback_headers(callback: ProxyCallbackTarget) -> dict[str, str]:
|
|
1902
|
+
if not callback.token:
|
|
1903
|
+
return {}
|
|
1904
|
+
return {PROXY_AUTH_HEADER: f"Bearer {callback.token}"}
|
|
1905
|
+
|
|
1906
|
+
|
|
1907
|
+
async def _worker_callback_outputs_and_uploads(
|
|
1908
|
+
state: ServeState,
|
|
1909
|
+
outputs: list[dict[str, Any]],
|
|
1910
|
+
) -> tuple[list[dict[str, Any]], list[WorkerCallbackUpload]]:
|
|
1911
|
+
output_payloads = [dict(output) for output in outputs]
|
|
1912
|
+
uploads: list[WorkerCallbackUpload] = []
|
|
1913
|
+
for output_index, output in enumerate(output_payloads):
|
|
1914
|
+
artifacts = output.get("artifacts")
|
|
1915
|
+
if not isinstance(artifacts, list):
|
|
1916
|
+
continue
|
|
1917
|
+
output_type = str(output.get("type") or "output")
|
|
1918
|
+
for artifact_index, artifact in enumerate(artifacts):
|
|
1919
|
+
if not isinstance(artifact, dict):
|
|
1920
|
+
continue
|
|
1921
|
+
artifact = cast(dict[str, Any], artifact)
|
|
1922
|
+
filename = str(artifact.get("filename") or "")
|
|
1923
|
+
if not filename:
|
|
1924
|
+
continue
|
|
1925
|
+
response = await state.client.fetch_output(
|
|
1926
|
+
{
|
|
1927
|
+
"filename": filename,
|
|
1928
|
+
"subfolder": str(artifact.get("subfolder") or ""),
|
|
1929
|
+
"type": str(artifact.get("type") or "output"),
|
|
1930
|
+
},
|
|
1931
|
+
request_headers={},
|
|
1932
|
+
)
|
|
1933
|
+
field_name = f"artifact_{output_index}_{artifact_index}"
|
|
1934
|
+
safe_filename = _safe_artifact_filename(
|
|
1935
|
+
filename,
|
|
1936
|
+
fallback=f"{field_name}{_extension_for_content_type(response.content_type)}",
|
|
1937
|
+
)
|
|
1938
|
+
artifact["upload_field"] = field_name
|
|
1939
|
+
artifact["content_type"] = response.content_type
|
|
1940
|
+
artifact["kind"] = output_kind(output_type, safe_filename)
|
|
1941
|
+
uploads.append(
|
|
1942
|
+
WorkerCallbackUpload(
|
|
1943
|
+
field_name=field_name,
|
|
1944
|
+
filename=safe_filename,
|
|
1945
|
+
content_type=response.content_type,
|
|
1946
|
+
body=response.body,
|
|
1947
|
+
)
|
|
1948
|
+
)
|
|
1949
|
+
return output_payloads, uploads
|
|
1950
|
+
|
|
1951
|
+
|
|
1952
|
+
def _record_proxy_run_error(record: ProxyRuntimeRun, payload: dict[str, Any]) -> None:
|
|
1953
|
+
record.status = "error"
|
|
1954
|
+
record.error = str(payload.get("message") or payload.get("error") or "Proxy run failed.")
|
|
1955
|
+
record.raw_result = {"status": "error", **payload}
|
|
1956
|
+
|
|
1957
|
+
|
|
1958
|
+
def _register_proxy_artifacts(state: ServeState, outputs: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
|
1959
|
+
output_payloads = [dict(output) for output in outputs]
|
|
1960
|
+
for output in output_payloads:
|
|
1961
|
+
output_type = str(output.get("type") or "output")
|
|
1962
|
+
artifacts = output.get("artifacts")
|
|
1963
|
+
if not isinstance(artifacts, list):
|
|
1964
|
+
continue
|
|
1965
|
+
for artifact in artifacts:
|
|
1966
|
+
if not isinstance(artifact, dict):
|
|
1967
|
+
continue
|
|
1968
|
+
filename = str(artifact.get("filename") or "")
|
|
1969
|
+
if not filename:
|
|
1970
|
+
continue
|
|
1971
|
+
artifact_id = f"artifact_{uuid.uuid4().hex}"
|
|
1972
|
+
state.proxy_artifacts[artifact_id] = ProxyArtifactRef(
|
|
1973
|
+
params={
|
|
1974
|
+
"filename": filename,
|
|
1975
|
+
"subfolder": str(artifact.get("subfolder") or ""),
|
|
1976
|
+
"type": str(artifact.get("type") or "output"),
|
|
1977
|
+
}
|
|
1978
|
+
)
|
|
1979
|
+
artifact["proxy_artifact_id"] = artifact_id
|
|
1980
|
+
artifact["url"] = f"/proxy/artifacts/{artifact_id}"
|
|
1981
|
+
artifact["kind"] = output_kind(output_type, filename)
|
|
1982
|
+
return output_payloads
|
|
1983
|
+
|
|
1984
|
+
|
|
1985
|
+
def _proxy_run_payload(record: ProxyRuntimeRun) -> dict[str, Any]:
|
|
1986
|
+
return {"status": record.status, "prompt_id": record.prompt_id, **record.raw_result}
|
|
1987
|
+
|
|
1988
|
+
|
|
1989
|
+
def _proxy_auth_response(request: web.Request) -> web.Response | None:
|
|
1990
|
+
token = getattr(_state(request).config, "proxy_token", None)
|
|
1991
|
+
if not token:
|
|
1992
|
+
return None
|
|
1993
|
+
expected = f"Bearer {token}"
|
|
1994
|
+
received = request.headers.get(PROXY_AUTH_HEADER, "")
|
|
1995
|
+
if secrets.compare_digest(received, expected):
|
|
1996
|
+
return None
|
|
1997
|
+
return web.json_response({"error": "forbidden", "message": "Proxy token is invalid."}, status=403)
|
|
1998
|
+
|
|
1999
|
+
|
|
2000
|
+
def _serve_session(request: web.Request) -> ServeSession:
|
|
2001
|
+
state = _state(request)
|
|
2002
|
+
header_value = request.headers.get(SESSION_HEADER_NAME)
|
|
2003
|
+
cookie_value = request.cookies.get(SESSION_COOKIE_NAME)
|
|
2004
|
+
session_id = _session_id_from_value(cookie_value) or _session_id_from_value(header_value) or f"anon_{uuid.uuid4().hex}"
|
|
2005
|
+
scope_key = SHARED_GALLERY_SCOPE if state.config.gallery == "shared" else session_id
|
|
2006
|
+
return state.state_store.ensure_session(session_id, scope_key=scope_key)
|
|
2007
|
+
|
|
2008
|
+
|
|
2009
|
+
def _json_response_for_session(
|
|
2010
|
+
payload: Mapping[str, Any],
|
|
2011
|
+
session: ServeSession,
|
|
2012
|
+
*,
|
|
2013
|
+
status: int = 200,
|
|
2014
|
+
request: web.Request | None = None,
|
|
2015
|
+
) -> web.Response:
|
|
2016
|
+
if request is not None:
|
|
2017
|
+
payload = _public_runtime_payload(request, payload)
|
|
2018
|
+
response = web.json_response(payload, status=status)
|
|
2019
|
+
response.set_cookie(
|
|
2020
|
+
SESSION_COOKIE_NAME,
|
|
2021
|
+
session.session_id,
|
|
2022
|
+
httponly=True,
|
|
2023
|
+
samesite="Lax",
|
|
2024
|
+
max_age=60 * 60 * 24 * 365,
|
|
2025
|
+
)
|
|
2026
|
+
return response
|
|
2027
|
+
|
|
2028
|
+
|
|
2029
|
+
def _public_runtime_payload(request: web.Request, payload: Mapping[str, Any]) -> dict[str, Any]:
|
|
2030
|
+
api_base_path = request.app.get(STUDIO_API_BASE_PATH_KEY, "")
|
|
2031
|
+
return cast(dict[str, Any], _with_public_runtime_urls(payload, api_base_path))
|
|
2032
|
+
|
|
2033
|
+
|
|
2034
|
+
def _with_public_runtime_urls(value: Any, api_base_path: str, key: str | None = None) -> Any:
|
|
2035
|
+
if isinstance(value, Mapping):
|
|
2036
|
+
return {
|
|
2037
|
+
str(item_key): _with_public_runtime_urls(item_value, api_base_path, str(item_key))
|
|
2038
|
+
for item_key, item_value in value.items()
|
|
2039
|
+
}
|
|
2040
|
+
if isinstance(value, list):
|
|
2041
|
+
return [_with_public_runtime_urls(item, api_base_path) for item in value]
|
|
2042
|
+
if isinstance(value, tuple):
|
|
2043
|
+
return [_with_public_runtime_urls(item, api_base_path) for item in value]
|
|
2044
|
+
if key in {"url", "upload_url"} and isinstance(value, str):
|
|
2045
|
+
return _public_runtime_url(value, api_base_path)
|
|
2046
|
+
return value
|
|
2047
|
+
|
|
2048
|
+
|
|
2049
|
+
def _public_runtime_url(url: str, api_base_path: str) -> str:
|
|
2050
|
+
if not api_base_path or not url.startswith("/"):
|
|
2051
|
+
return url
|
|
2052
|
+
normalized_base = _normalize_public_prefix(api_base_path)
|
|
2053
|
+
if not normalized_base or url == normalized_base or url.startswith(f"{normalized_base}/"):
|
|
2054
|
+
return url
|
|
2055
|
+
if url.startswith(("/outputs/view", "/uploads/")):
|
|
2056
|
+
return f"{normalized_base}{url}"
|
|
2057
|
+
return url
|
|
2058
|
+
|
|
2059
|
+
|
|
2060
|
+
def _safe_token(value: str) -> str:
|
|
2061
|
+
return "".join(char for char in value if char.isalnum() or char in {"-", "_"})
|
|
2062
|
+
|
|
2063
|
+
|
|
2064
|
+
def _session_id_from_value(value: str | None) -> str | None:
|
|
2065
|
+
if value and _safe_token(value) == value:
|
|
2066
|
+
return value
|
|
2067
|
+
return None
|
|
2068
|
+
|
|
2069
|
+
|
|
2070
|
+
def _contracts_payload(state: ServeState) -> dict[str, Any]:
|
|
2071
|
+
manifest = state.manifest_snapshot()
|
|
2072
|
+
contracts: list[dict[str, Any]] = []
|
|
2073
|
+
for workflow_name, workflow in manifest.workflows.items():
|
|
2074
|
+
execution_contract = workflow.execution_contract
|
|
2075
|
+
if execution_contract is None:
|
|
2076
|
+
continue
|
|
2077
|
+
for contract_name, contract in execution_contract.contracts.items():
|
|
2078
|
+
contracts.append(_contract_payload(workflow_name, contract_name, contract))
|
|
2079
|
+
return {
|
|
2080
|
+
"environment": state.env.name,
|
|
2081
|
+
"contracts": contracts,
|
|
2082
|
+
}
|
|
2083
|
+
|
|
2084
|
+
|
|
2085
|
+
def _single_contract_payload(
|
|
2086
|
+
state: ServeState,
|
|
2087
|
+
workflow_name: str,
|
|
2088
|
+
contract_name: str,
|
|
2089
|
+
) -> dict[str, Any]:
|
|
2090
|
+
manifest = state.manifest_snapshot()
|
|
2091
|
+
workflow = manifest.workflows.get(workflow_name)
|
|
2092
|
+
if workflow is None or workflow.execution_contract is None:
|
|
2093
|
+
raise ValueError(f"Workflow '{workflow_name}' does not declare contracts.")
|
|
2094
|
+
contract = workflow.execution_contract.contracts.get(contract_name)
|
|
2095
|
+
if contract is None:
|
|
2096
|
+
raise ValueError(f"Workflow '{workflow_name}' does not declare contract '{contract_name}'.")
|
|
2097
|
+
return _contract_payload(workflow_name, contract_name, contract)
|
|
2098
|
+
|
|
2099
|
+
|
|
2100
|
+
def _contract_payload(
|
|
2101
|
+
workflow_name: str,
|
|
2102
|
+
contract_name: str,
|
|
2103
|
+
contract: NamedWorkflowContract,
|
|
2104
|
+
) -> dict[str, Any]:
|
|
2105
|
+
return {
|
|
2106
|
+
"workflow": workflow_name,
|
|
2107
|
+
"contract": contract_name,
|
|
2108
|
+
"display_name": contract.display_name,
|
|
2109
|
+
"description": contract.description,
|
|
2110
|
+
"inputs": [item.to_dict() for item in contract.inputs],
|
|
2111
|
+
"outputs": [item.to_dict() for item in contract.outputs],
|
|
2112
|
+
}
|
|
2113
|
+
|
|
2114
|
+
|
|
2115
|
+
async def _run_contract(
|
|
2116
|
+
state: ServeState,
|
|
2117
|
+
session: ServeSession,
|
|
2118
|
+
workflow_name: str,
|
|
2119
|
+
contract_name: str,
|
|
2120
|
+
body: dict[str, Any],
|
|
2121
|
+
) -> dict[str, Any]:
|
|
2122
|
+
if "inputs" in body:
|
|
2123
|
+
inputs = body["inputs"]
|
|
2124
|
+
else:
|
|
2125
|
+
control_keys = {"wait", "timeout_seconds", "poll_interval_seconds"}
|
|
2126
|
+
inputs = {key: value for key, value in body.items() if key not in control_keys}
|
|
2127
|
+
if not isinstance(inputs, dict):
|
|
2128
|
+
raise ValueError("'inputs' must be a JSON object.")
|
|
2129
|
+
wait = bool(body.get("wait", False))
|
|
2130
|
+
timeout_seconds = float(body.get("timeout_seconds", state.config.run_timeout_seconds))
|
|
2131
|
+
poll_interval_seconds = float(body.get("poll_interval_seconds", 1))
|
|
2132
|
+
|
|
2133
|
+
manifest = state.manifest_snapshot()
|
|
2134
|
+
prepared_contract = await _prepare_contract_run_inputs(state, workflow_name, contract_name, inputs)
|
|
2135
|
+
inputs = prepared_contract.inputs
|
|
2136
|
+
build_result = build_manifest_contract_prompt(
|
|
2137
|
+
manifest,
|
|
2138
|
+
state.env.cec_path,
|
|
2139
|
+
workflow_name,
|
|
2140
|
+
inputs,
|
|
2141
|
+
contract_name=contract_name,
|
|
2142
|
+
)
|
|
2143
|
+
if build_result.has_errors:
|
|
2144
|
+
payload: dict[str, Any] = {
|
|
2145
|
+
"status": "invalid_request",
|
|
2146
|
+
"issues": [asdict(issue) for issue in build_result.issues],
|
|
2147
|
+
"message": "Contract inputs could not be applied to the workflow prompt.",
|
|
2148
|
+
}
|
|
2149
|
+
error_record = _record_failed_run(state, session, workflow_name, contract_name, {"inputs": inputs}, payload)
|
|
2150
|
+
payload.update(error_record)
|
|
2151
|
+
return payload
|
|
2152
|
+
|
|
2153
|
+
run_id = f"run_{uuid.uuid4().hex}"
|
|
2154
|
+
callback_url = _worker_callback_url(state, run_id, wait=wait)
|
|
2155
|
+
callback_token = _callback_auth_token(state) if callback_url else None
|
|
2156
|
+
|
|
2157
|
+
async def record_submitted(prompt_id: str) -> None:
|
|
2158
|
+
response = {
|
|
2159
|
+
"status": "submitted",
|
|
2160
|
+
"run_id": run_id,
|
|
2161
|
+
"prompt_id": prompt_id,
|
|
2162
|
+
"issues": [asdict(issue) for issue in build_result.issues],
|
|
2163
|
+
}
|
|
2164
|
+
state.state_store.record_run(
|
|
2165
|
+
ServeRunRecord(
|
|
2166
|
+
run_id=run_id,
|
|
2167
|
+
session_id=session.session_id,
|
|
2168
|
+
scope_key=session.scope_key,
|
|
2169
|
+
workflow=workflow_name,
|
|
2170
|
+
contract=contract_name,
|
|
2171
|
+
status="submitted",
|
|
2172
|
+
prompt_id=prompt_id,
|
|
2173
|
+
inputs=_display_inputs(inputs),
|
|
2174
|
+
raw_result=dict(response),
|
|
2175
|
+
)
|
|
2176
|
+
)
|
|
2177
|
+
|
|
2178
|
+
execution = await state.executor.execute(
|
|
2179
|
+
RunExecutionRequest(
|
|
2180
|
+
prompt=build_result.prompt,
|
|
2181
|
+
outputs=build_result.outputs,
|
|
2182
|
+
wait=wait,
|
|
2183
|
+
timeout_seconds=timeout_seconds,
|
|
2184
|
+
poll_interval_seconds=poll_interval_seconds,
|
|
2185
|
+
cache_token=uuid.uuid4().hex[:10],
|
|
2186
|
+
on_submitted=record_submitted,
|
|
2187
|
+
staged_uploads=prepared_contract.staged_uploads,
|
|
2188
|
+
callback_run_id=run_id if callback_url else None,
|
|
2189
|
+
callback_url=callback_url,
|
|
2190
|
+
callback_token=callback_token,
|
|
2191
|
+
)
|
|
2192
|
+
)
|
|
2193
|
+
|
|
2194
|
+
response: dict[str, Any] = {
|
|
2195
|
+
"status": execution.status,
|
|
2196
|
+
"run_id": run_id,
|
|
2197
|
+
"prompt_id": execution.prompt_id,
|
|
2198
|
+
"issues": [asdict(issue) for issue in build_result.issues],
|
|
2199
|
+
}
|
|
2200
|
+
output_slots = _output_slots_for_run(
|
|
2201
|
+
run_id=run_id,
|
|
2202
|
+
session=session,
|
|
2203
|
+
workflow_name=workflow_name,
|
|
2204
|
+
contract_name=contract_name,
|
|
2205
|
+
outputs=build_result.outputs,
|
|
2206
|
+
prompt_id=execution.prompt_id,
|
|
2207
|
+
inputs=inputs,
|
|
2208
|
+
)
|
|
2209
|
+
if execution.status != "completed":
|
|
2210
|
+
if callback_url:
|
|
2211
|
+
existing_run = state.state_store.get_run_record(run_id)
|
|
2212
|
+
if existing_run is not None and existing_run.status != "submitted":
|
|
2213
|
+
snapshot = _callback_run_snapshot(state, session, run_id)
|
|
2214
|
+
snapshot["status"] = existing_run.status
|
|
2215
|
+
return snapshot
|
|
2216
|
+
pending_items = _pending_gallery_items_for_slots(
|
|
2217
|
+
output_slots,
|
|
2218
|
+
inputs=inputs,
|
|
2219
|
+
response=response,
|
|
2220
|
+
)
|
|
2221
|
+
state.state_store.record_output_slots(output_slots)
|
|
2222
|
+
state.state_store.record_gallery_items(pending_items)
|
|
2223
|
+
response["output_slots"] = [slot.to_public_dict() for slot in output_slots]
|
|
2224
|
+
response["gallery_items"] = [item.to_public_dict() for item in pending_items]
|
|
2225
|
+
if callback_url:
|
|
2226
|
+
return response
|
|
2227
|
+
task = asyncio.create_task(
|
|
2228
|
+
_complete_submitted_run(
|
|
2229
|
+
state,
|
|
2230
|
+
session,
|
|
2231
|
+
run_id=run_id,
|
|
2232
|
+
workflow_name=workflow_name,
|
|
2233
|
+
contract_name=contract_name,
|
|
2234
|
+
inputs=inputs,
|
|
2235
|
+
prompt_id=execution.prompt_id,
|
|
2236
|
+
outputs=build_result.outputs,
|
|
2237
|
+
timeout_seconds=timeout_seconds,
|
|
2238
|
+
poll_interval_seconds=poll_interval_seconds,
|
|
2239
|
+
issues=[asdict(issue) for issue in build_result.issues],
|
|
2240
|
+
output_slots=output_slots,
|
|
2241
|
+
created_at=output_slots[0].created_at if output_slots else utc_now(),
|
|
2242
|
+
)
|
|
2243
|
+
)
|
|
2244
|
+
_track_active_run_task(state, run_id, task)
|
|
2245
|
+
return response
|
|
2246
|
+
|
|
2247
|
+
response["outputs"] = execution.outputs
|
|
2248
|
+
state.state_store.record_run(
|
|
2249
|
+
ServeRunRecord(
|
|
2250
|
+
run_id=run_id,
|
|
2251
|
+
session_id=session.session_id,
|
|
2252
|
+
scope_key=session.scope_key,
|
|
2253
|
+
workflow=workflow_name,
|
|
2254
|
+
contract=contract_name,
|
|
2255
|
+
status="completed",
|
|
2256
|
+
prompt_id=execution.prompt_id,
|
|
2257
|
+
inputs=_display_inputs(inputs),
|
|
2258
|
+
raw_result=dict(response),
|
|
2259
|
+
)
|
|
2260
|
+
)
|
|
2261
|
+
resolved_slots = _resolved_output_slots_for_response(output_slots, response)
|
|
2262
|
+
state.state_store.record_output_slots(resolved_slots)
|
|
2263
|
+
gallery_items = _gallery_items_for_outputs(
|
|
2264
|
+
run_id=run_id,
|
|
2265
|
+
session=session,
|
|
2266
|
+
workflow_name=workflow_name,
|
|
2267
|
+
contract_name=contract_name,
|
|
2268
|
+
inputs=inputs,
|
|
2269
|
+
response=response,
|
|
2270
|
+
output_slots=resolved_slots,
|
|
2271
|
+
)
|
|
2272
|
+
gallery_items.extend(
|
|
2273
|
+
_empty_gallery_items_for_slots(
|
|
2274
|
+
resolved_slots,
|
|
2275
|
+
existing_items=gallery_items,
|
|
2276
|
+
inputs=inputs,
|
|
2277
|
+
response=response,
|
|
2278
|
+
)
|
|
2279
|
+
)
|
|
2280
|
+
state.state_store.record_gallery_items(gallery_items)
|
|
2281
|
+
response["output_slots"] = [slot.to_public_dict() for slot in resolved_slots]
|
|
2282
|
+
response["gallery_items"] = [item.to_public_dict() for item in gallery_items]
|
|
2283
|
+
return response
|
|
2284
|
+
|
|
2285
|
+
|
|
2286
|
+
async def _complete_submitted_run(
|
|
2287
|
+
state: ServeState,
|
|
2288
|
+
session: ServeSession,
|
|
2289
|
+
*,
|
|
2290
|
+
run_id: str,
|
|
2291
|
+
workflow_name: str,
|
|
2292
|
+
contract_name: str,
|
|
2293
|
+
inputs: dict[str, Any],
|
|
2294
|
+
prompt_id: str,
|
|
2295
|
+
outputs: tuple[Any, ...],
|
|
2296
|
+
timeout_seconds: float,
|
|
2297
|
+
poll_interval_seconds: float,
|
|
2298
|
+
issues: list[dict[str, Any]],
|
|
2299
|
+
output_slots: list[ServeRunOutputSlot],
|
|
2300
|
+
created_at: str,
|
|
2301
|
+
) -> None:
|
|
2302
|
+
running_slots = [
|
|
2303
|
+
_copy_output_slot(slot, status="running", prompt_id=prompt_id, raw_result={"status": "running", "run_id": run_id, "prompt_id": prompt_id})
|
|
2304
|
+
for slot in output_slots
|
|
2305
|
+
]
|
|
2306
|
+
state.state_store.record_output_slots(running_slots)
|
|
2307
|
+
state.state_store.record_run(
|
|
2308
|
+
ServeRunRecord(
|
|
2309
|
+
run_id=run_id,
|
|
2310
|
+
session_id=session.session_id,
|
|
2311
|
+
scope_key=session.scope_key,
|
|
2312
|
+
workflow=workflow_name,
|
|
2313
|
+
contract=contract_name,
|
|
2314
|
+
status="running",
|
|
2315
|
+
prompt_id=prompt_id,
|
|
2316
|
+
inputs=_display_inputs(inputs),
|
|
2317
|
+
raw_result={"status": "running", "run_id": run_id, "prompt_id": prompt_id, "issues": issues},
|
|
2318
|
+
created_at=created_at,
|
|
2319
|
+
)
|
|
2320
|
+
)
|
|
2321
|
+
try:
|
|
2322
|
+
execution = await state.executor.complete_submitted(
|
|
2323
|
+
prompt_id,
|
|
2324
|
+
outputs,
|
|
2325
|
+
timeout_seconds=timeout_seconds,
|
|
2326
|
+
poll_interval_seconds=poll_interval_seconds,
|
|
2327
|
+
)
|
|
2328
|
+
except ComfyGitServeTimeoutError as exc:
|
|
2329
|
+
payload = {"error": "timeout", "message": str(exc), "prompt_id": prompt_id}
|
|
2330
|
+
_record_failed_run(
|
|
2331
|
+
state,
|
|
2332
|
+
session,
|
|
2333
|
+
workflow_name,
|
|
2334
|
+
contract_name,
|
|
2335
|
+
{"inputs": inputs},
|
|
2336
|
+
payload,
|
|
2337
|
+
run_id=run_id,
|
|
2338
|
+
prompt_id=prompt_id,
|
|
2339
|
+
output_slots=output_slots,
|
|
2340
|
+
created_at=created_at,
|
|
2341
|
+
)
|
|
2342
|
+
return
|
|
2343
|
+
except ComfyUIRequestError as exc:
|
|
2344
|
+
payload = {
|
|
2345
|
+
"error": "comfyui_rejected_request",
|
|
2346
|
+
"message": str(exc),
|
|
2347
|
+
"comfy_status": exc.status,
|
|
2348
|
+
"comfy_url": exc.url,
|
|
2349
|
+
"comfyui": exc.payload,
|
|
2350
|
+
"prompt_id": prompt_id,
|
|
2351
|
+
}
|
|
2352
|
+
_record_failed_run(
|
|
2353
|
+
state,
|
|
2354
|
+
session,
|
|
2355
|
+
workflow_name,
|
|
2356
|
+
contract_name,
|
|
2357
|
+
{"inputs": inputs},
|
|
2358
|
+
payload,
|
|
2359
|
+
run_id=run_id,
|
|
2360
|
+
prompt_id=prompt_id,
|
|
2361
|
+
output_slots=output_slots,
|
|
2362
|
+
created_at=created_at,
|
|
2363
|
+
)
|
|
2364
|
+
return
|
|
2365
|
+
except ComfyUIExecutionError as exc:
|
|
2366
|
+
payload = {
|
|
2367
|
+
"error": "comfyui_execution_failed",
|
|
2368
|
+
"message": str(exc),
|
|
2369
|
+
"prompt_id": exc.prompt_id,
|
|
2370
|
+
"comfyui": exc.payload,
|
|
2371
|
+
}
|
|
2372
|
+
_record_failed_run(
|
|
2373
|
+
state,
|
|
2374
|
+
session,
|
|
2375
|
+
workflow_name,
|
|
2376
|
+
contract_name,
|
|
2377
|
+
{"inputs": inputs},
|
|
2378
|
+
payload,
|
|
2379
|
+
run_id=run_id,
|
|
2380
|
+
prompt_id=prompt_id,
|
|
2381
|
+
output_slots=output_slots,
|
|
2382
|
+
created_at=created_at,
|
|
2383
|
+
)
|
|
2384
|
+
return
|
|
2385
|
+
except (aiohttp.ClientError, asyncio.TimeoutError) as exc:
|
|
2386
|
+
payload = _executor_unavailable_payload(state, exc, prompt_id=prompt_id)
|
|
2387
|
+
_record_failed_run(
|
|
2388
|
+
state,
|
|
2389
|
+
session,
|
|
2390
|
+
workflow_name,
|
|
2391
|
+
contract_name,
|
|
2392
|
+
{"inputs": inputs},
|
|
2393
|
+
payload,
|
|
2394
|
+
run_id=run_id,
|
|
2395
|
+
prompt_id=prompt_id,
|
|
2396
|
+
output_slots=output_slots,
|
|
2397
|
+
created_at=created_at,
|
|
2398
|
+
)
|
|
2399
|
+
return
|
|
2400
|
+
except Exception as exc:
|
|
2401
|
+
payload = {"error": "internal_error", "message": str(exc), "prompt_id": prompt_id}
|
|
2402
|
+
_record_failed_run(
|
|
2403
|
+
state,
|
|
2404
|
+
session,
|
|
2405
|
+
workflow_name,
|
|
2406
|
+
contract_name,
|
|
2407
|
+
{"inputs": inputs},
|
|
2408
|
+
payload,
|
|
2409
|
+
run_id=run_id,
|
|
2410
|
+
prompt_id=prompt_id,
|
|
2411
|
+
output_slots=output_slots,
|
|
2412
|
+
created_at=created_at,
|
|
2413
|
+
)
|
|
2414
|
+
return
|
|
2415
|
+
|
|
2416
|
+
response: dict[str, Any] = {
|
|
2417
|
+
"status": "completed",
|
|
2418
|
+
"run_id": run_id,
|
|
2419
|
+
"prompt_id": execution.prompt_id,
|
|
2420
|
+
"issues": issues,
|
|
2421
|
+
"outputs": execution.outputs,
|
|
2422
|
+
}
|
|
2423
|
+
state.state_store.record_run(
|
|
2424
|
+
ServeRunRecord(
|
|
2425
|
+
run_id=run_id,
|
|
2426
|
+
session_id=session.session_id,
|
|
2427
|
+
scope_key=session.scope_key,
|
|
2428
|
+
workflow=workflow_name,
|
|
2429
|
+
contract=contract_name,
|
|
2430
|
+
status="completed",
|
|
2431
|
+
prompt_id=execution.prompt_id,
|
|
2432
|
+
inputs=_display_inputs(inputs),
|
|
2433
|
+
raw_result=dict(response),
|
|
2434
|
+
created_at=created_at,
|
|
2435
|
+
)
|
|
2436
|
+
)
|
|
2437
|
+
resolved_slots = _resolved_output_slots_for_response(output_slots, response)
|
|
2438
|
+
state.state_store.record_output_slots(resolved_slots)
|
|
2439
|
+
gallery_items = _gallery_items_for_outputs(
|
|
2440
|
+
run_id=run_id,
|
|
2441
|
+
session=session,
|
|
2442
|
+
workflow_name=workflow_name,
|
|
2443
|
+
contract_name=contract_name,
|
|
2444
|
+
inputs=inputs,
|
|
2445
|
+
response=response,
|
|
2446
|
+
output_slots=resolved_slots,
|
|
2447
|
+
created_at=created_at,
|
|
2448
|
+
)
|
|
2449
|
+
gallery_items.extend(
|
|
2450
|
+
_empty_gallery_items_for_slots(
|
|
2451
|
+
resolved_slots,
|
|
2452
|
+
existing_items=gallery_items,
|
|
2453
|
+
inputs=inputs,
|
|
2454
|
+
response=response,
|
|
2455
|
+
)
|
|
2456
|
+
)
|
|
2457
|
+
if not gallery_items:
|
|
2458
|
+
gallery_items = [
|
|
2459
|
+
_json_gallery_item_for_completed_run(
|
|
2460
|
+
run_id=run_id,
|
|
2461
|
+
session=session,
|
|
2462
|
+
workflow_name=workflow_name,
|
|
2463
|
+
contract_name=contract_name,
|
|
2464
|
+
inputs=inputs,
|
|
2465
|
+
response=response,
|
|
2466
|
+
item_id=_gallery_item_id_for_slot(resolved_slots[0]) if resolved_slots else f"gallery_{run_id}",
|
|
2467
|
+
slot_id=resolved_slots[0].slot_id if resolved_slots else None,
|
|
2468
|
+
created_at=created_at,
|
|
2469
|
+
)
|
|
2470
|
+
]
|
|
2471
|
+
state.state_store.record_gallery_items(gallery_items)
|
|
2472
|
+
|
|
2473
|
+
|
|
2474
|
+
def _record_completed_run_response(
|
|
2475
|
+
state: ServeState,
|
|
2476
|
+
session: ServeSession,
|
|
2477
|
+
*,
|
|
2478
|
+
workflow_name: str,
|
|
2479
|
+
contract_name: str,
|
|
2480
|
+
inputs: dict[str, Any],
|
|
2481
|
+
response: dict[str, Any],
|
|
2482
|
+
output_slots: list[ServeRunOutputSlot],
|
|
2483
|
+
created_at: str,
|
|
2484
|
+
) -> None:
|
|
2485
|
+
prompt_id = str(response.get("prompt_id") or "") or None
|
|
2486
|
+
state.state_store.record_run(
|
|
2487
|
+
ServeRunRecord(
|
|
2488
|
+
run_id=str(response.get("run_id") or ""),
|
|
2489
|
+
session_id=session.session_id,
|
|
2490
|
+
scope_key=session.scope_key,
|
|
2491
|
+
workflow=workflow_name,
|
|
2492
|
+
contract=contract_name,
|
|
2493
|
+
status="completed",
|
|
2494
|
+
prompt_id=prompt_id,
|
|
2495
|
+
inputs=_display_inputs(inputs),
|
|
2496
|
+
raw_result=dict(response),
|
|
2497
|
+
created_at=created_at,
|
|
2498
|
+
)
|
|
2499
|
+
)
|
|
2500
|
+
resolved_slots = _resolved_output_slots_for_response(output_slots, response)
|
|
2501
|
+
state.state_store.record_output_slots(resolved_slots)
|
|
2502
|
+
gallery_items = _gallery_items_for_outputs(
|
|
2503
|
+
run_id=str(response.get("run_id") or ""),
|
|
2504
|
+
session=session,
|
|
2505
|
+
workflow_name=workflow_name,
|
|
2506
|
+
contract_name=contract_name,
|
|
2507
|
+
inputs=inputs,
|
|
2508
|
+
response=response,
|
|
2509
|
+
output_slots=resolved_slots,
|
|
2510
|
+
created_at=created_at,
|
|
2511
|
+
)
|
|
2512
|
+
gallery_items.extend(
|
|
2513
|
+
_empty_gallery_items_for_slots(
|
|
2514
|
+
resolved_slots,
|
|
2515
|
+
existing_items=gallery_items,
|
|
2516
|
+
inputs=inputs,
|
|
2517
|
+
response=response,
|
|
2518
|
+
)
|
|
2519
|
+
)
|
|
2520
|
+
if not gallery_items:
|
|
2521
|
+
gallery_items = [
|
|
2522
|
+
_json_gallery_item_for_completed_run(
|
|
2523
|
+
run_id=str(response.get("run_id") or ""),
|
|
2524
|
+
session=session,
|
|
2525
|
+
workflow_name=workflow_name,
|
|
2526
|
+
contract_name=contract_name,
|
|
2527
|
+
inputs=inputs,
|
|
2528
|
+
response=response,
|
|
2529
|
+
item_id=_gallery_item_id_for_slot(resolved_slots[0]) if resolved_slots else f"gallery_{response.get('run_id')}",
|
|
2530
|
+
slot_id=resolved_slots[0].slot_id if resolved_slots else None,
|
|
2531
|
+
created_at=created_at,
|
|
2532
|
+
)
|
|
2533
|
+
]
|
|
2534
|
+
state.state_store.record_gallery_items(gallery_items)
|
|
2535
|
+
|
|
2536
|
+
|
|
2537
|
+
def _output_slots_for_run(
|
|
2538
|
+
*,
|
|
2539
|
+
run_id: str,
|
|
2540
|
+
session: ServeSession,
|
|
2541
|
+
workflow_name: str,
|
|
2542
|
+
contract_name: str,
|
|
2543
|
+
outputs: tuple[Any, ...],
|
|
2544
|
+
prompt_id: str | None,
|
|
2545
|
+
inputs: dict[str, Any],
|
|
2546
|
+
) -> list[ServeRunOutputSlot]:
|
|
2547
|
+
created_at = utc_now()
|
|
2548
|
+
slots: list[ServeRunOutputSlot] = []
|
|
2549
|
+
declared_outputs = outputs or (None,)
|
|
2550
|
+
for index, output in enumerate(declared_outputs):
|
|
2551
|
+
output_name = str(getattr(output, "name", "result") or "result")
|
|
2552
|
+
output_type = _slot_output_type(str(getattr(output, "type", "json") or "json"))
|
|
2553
|
+
width, height = _fallback_dimensions_for_type(output_type)
|
|
2554
|
+
slots.append(
|
|
2555
|
+
ServeRunOutputSlot(
|
|
2556
|
+
slot_id=_slot_id(run_id, index, output_name),
|
|
2557
|
+
run_id=run_id,
|
|
2558
|
+
session_id=session.session_id,
|
|
2559
|
+
scope_key=session.scope_key,
|
|
2560
|
+
workflow=workflow_name,
|
|
2561
|
+
contract=contract_name,
|
|
2562
|
+
output_name=output_name,
|
|
2563
|
+
output_type=output_type,
|
|
2564
|
+
status="pending",
|
|
2565
|
+
prompt_id=prompt_id,
|
|
2566
|
+
width=width,
|
|
2567
|
+
height=height,
|
|
2568
|
+
raw_result={"status": "pending", "run_id": run_id, "prompt_id": prompt_id, "inputs": _display_inputs(inputs)},
|
|
2569
|
+
created_at=created_at,
|
|
2570
|
+
updated_at=created_at,
|
|
2571
|
+
)
|
|
2572
|
+
)
|
|
2573
|
+
return slots
|
|
2574
|
+
|
|
2575
|
+
|
|
2576
|
+
def _pending_gallery_items_for_slots(
|
|
2577
|
+
slots: list[ServeRunOutputSlot],
|
|
2578
|
+
*,
|
|
2579
|
+
inputs: dict[str, Any],
|
|
2580
|
+
response: dict[str, Any],
|
|
2581
|
+
) -> list[ServeGalleryItem]:
|
|
2582
|
+
display_inputs = _display_inputs(inputs)
|
|
2583
|
+
return [
|
|
2584
|
+
ServeGalleryItem(
|
|
2585
|
+
item_id=_gallery_item_id_for_slot(slot),
|
|
2586
|
+
run_id=slot.run_id,
|
|
2587
|
+
session_id=slot.session_id,
|
|
2588
|
+
scope_key=slot.scope_key,
|
|
2589
|
+
workflow=slot.workflow,
|
|
2590
|
+
contract=slot.contract,
|
|
2591
|
+
status="pending",
|
|
2592
|
+
output_type=slot.output_type,
|
|
2593
|
+
slot_id=slot.slot_id,
|
|
2594
|
+
output_name=slot.output_name,
|
|
2595
|
+
prompt_id=slot.prompt_id,
|
|
2596
|
+
width=slot.width,
|
|
2597
|
+
height=slot.height,
|
|
2598
|
+
inputs=display_inputs,
|
|
2599
|
+
raw_result=dict(response),
|
|
2600
|
+
created_at=slot.created_at,
|
|
2601
|
+
updated_at=slot.created_at,
|
|
2602
|
+
)
|
|
2603
|
+
for slot in slots
|
|
2604
|
+
]
|
|
2605
|
+
|
|
2606
|
+
|
|
2607
|
+
def _resolved_output_slots_for_response(
|
|
2608
|
+
slots: list[ServeRunOutputSlot],
|
|
2609
|
+
response: dict[str, Any],
|
|
2610
|
+
) -> list[ServeRunOutputSlot]:
|
|
2611
|
+
response_outputs = [output for output in response.get("outputs") or [] if isinstance(output, Mapping)]
|
|
2612
|
+
resolved: list[ServeRunOutputSlot] = []
|
|
2613
|
+
for index, slot in enumerate(slots):
|
|
2614
|
+
output = response_outputs[index] if index < len(response_outputs) else None
|
|
2615
|
+
artifacts = output.get("artifacts") if isinstance(output, Mapping) else None
|
|
2616
|
+
artifact_list = artifacts if isinstance(artifacts, list) else []
|
|
2617
|
+
output_type = _slot_output_type(str(output.get("type") or slot.output_type)) if isinstance(output, Mapping) else slot.output_type
|
|
2618
|
+
width, height = slot.width, slot.height
|
|
2619
|
+
if artifact_list:
|
|
2620
|
+
first_artifact = artifact_list[0]
|
|
2621
|
+
if isinstance(first_artifact, Mapping):
|
|
2622
|
+
item_type = output_kind(output_type, str(first_artifact.get("filename") or ""))
|
|
2623
|
+
width, height = _gallery_dimensions_for_artifact(item_type, first_artifact)
|
|
2624
|
+
output_type = item_type
|
|
2625
|
+
resolved.append(
|
|
2626
|
+
_copy_output_slot(
|
|
2627
|
+
slot,
|
|
2628
|
+
status="done" if artifact_list else "empty",
|
|
2629
|
+
output_type=output_type,
|
|
2630
|
+
prompt_id=str(response.get("prompt_id") or "") or slot.prompt_id,
|
|
2631
|
+
width=width,
|
|
2632
|
+
height=height,
|
|
2633
|
+
raw_result=dict(response),
|
|
2634
|
+
)
|
|
2635
|
+
)
|
|
2636
|
+
return resolved
|
|
2637
|
+
|
|
2638
|
+
|
|
2639
|
+
def _copy_output_slot(
|
|
2640
|
+
slot: ServeRunOutputSlot,
|
|
2641
|
+
*,
|
|
2642
|
+
status: str,
|
|
2643
|
+
output_type: str | None = None,
|
|
2644
|
+
prompt_id: str | None = None,
|
|
2645
|
+
width: int | None = None,
|
|
2646
|
+
height: int | None = None,
|
|
2647
|
+
error: str | None = None,
|
|
2648
|
+
raw_result: dict[str, Any] | None = None,
|
|
2649
|
+
) -> ServeRunOutputSlot:
|
|
2650
|
+
return ServeRunOutputSlot(
|
|
2651
|
+
slot_id=slot.slot_id,
|
|
2652
|
+
run_id=slot.run_id,
|
|
2653
|
+
session_id=slot.session_id,
|
|
2654
|
+
scope_key=slot.scope_key,
|
|
2655
|
+
workflow=slot.workflow,
|
|
2656
|
+
contract=slot.contract,
|
|
2657
|
+
output_name=slot.output_name,
|
|
2658
|
+
output_type=output_type or slot.output_type,
|
|
2659
|
+
status=status,
|
|
2660
|
+
prompt_id=prompt_id or slot.prompt_id,
|
|
2661
|
+
width=width if width is not None else slot.width,
|
|
2662
|
+
height=height if height is not None else slot.height,
|
|
2663
|
+
error=error,
|
|
2664
|
+
raw_result=raw_result,
|
|
2665
|
+
created_at=slot.created_at,
|
|
2666
|
+
)
|
|
2667
|
+
|
|
2668
|
+
|
|
2669
|
+
def _slot_id(run_id: str, index: int, output_name: str) -> str:
|
|
2670
|
+
safe_name = _safe_token(output_name) or "output"
|
|
2671
|
+
return f"slot_{run_id}_{index}_{safe_name}"
|
|
2672
|
+
|
|
2673
|
+
|
|
2674
|
+
def _gallery_item_id_for_slot(slot: ServeRunOutputSlot) -> str:
|
|
2675
|
+
if slot.slot_id.startswith(f"slot_{slot.run_id}_0_"):
|
|
2676
|
+
return f"gallery_{slot.run_id}"
|
|
2677
|
+
return f"gallery_{slot.slot_id}"
|
|
2678
|
+
|
|
2679
|
+
|
|
2680
|
+
def _slot_output_type(output_type: str) -> str:
|
|
2681
|
+
normalized = output_type.lower()
|
|
2682
|
+
return normalized if normalized in {"image", "video", "audio", "json"} else "json"
|
|
2683
|
+
|
|
2684
|
+
|
|
2685
|
+
def _fallback_dimensions_for_type(output_type: str) -> tuple[int, int]:
|
|
2686
|
+
return (4, 1) if output_type == "audio" else (1, 1)
|
|
2687
|
+
|
|
2688
|
+
|
|
2689
|
+
def _gallery_dimensions_for_artifact(item_type: str, artifact: Mapping[str, Any]) -> tuple[int, int]:
|
|
2690
|
+
if item_type == "audio":
|
|
2691
|
+
return (4, 1)
|
|
2692
|
+
if item_type in {"image", "video"}:
|
|
2693
|
+
return artifact_dimensions(artifact)
|
|
2694
|
+
return (1, 1)
|
|
2695
|
+
|
|
2696
|
+
|
|
2697
|
+
def _json_gallery_item_for_completed_run(
|
|
2698
|
+
*,
|
|
2699
|
+
run_id: str,
|
|
2700
|
+
session: ServeSession,
|
|
2701
|
+
workflow_name: str,
|
|
2702
|
+
contract_name: str,
|
|
2703
|
+
inputs: dict[str, Any],
|
|
2704
|
+
response: dict[str, Any],
|
|
2705
|
+
item_id: str,
|
|
2706
|
+
slot_id: str | None,
|
|
2707
|
+
created_at: str,
|
|
2708
|
+
) -> ServeGalleryItem:
|
|
2709
|
+
return ServeGalleryItem(
|
|
2710
|
+
item_id=item_id,
|
|
2711
|
+
run_id=run_id,
|
|
2712
|
+
session_id=session.session_id,
|
|
2713
|
+
scope_key=session.scope_key,
|
|
2714
|
+
workflow=workflow_name,
|
|
2715
|
+
contract=contract_name,
|
|
2716
|
+
status="done",
|
|
2717
|
+
output_type="json",
|
|
2718
|
+
slot_id=slot_id,
|
|
2719
|
+
output_name="result",
|
|
2720
|
+
prompt_id=str(response.get("prompt_id") or "") or None,
|
|
2721
|
+
width=1,
|
|
2722
|
+
height=1,
|
|
2723
|
+
inputs=_display_inputs(inputs),
|
|
2724
|
+
raw_result=dict(response),
|
|
2725
|
+
created_at=created_at,
|
|
2726
|
+
updated_at=created_at,
|
|
2727
|
+
)
|
|
2728
|
+
|
|
2729
|
+
|
|
2730
|
+
def _record_failed_run(
|
|
2731
|
+
state: ServeState,
|
|
2732
|
+
session: ServeSession,
|
|
2733
|
+
workflow_name: str,
|
|
2734
|
+
contract_name: str,
|
|
2735
|
+
body: dict[str, Any],
|
|
2736
|
+
payload: dict[str, Any],
|
|
2737
|
+
*,
|
|
2738
|
+
run_id: str | None = None,
|
|
2739
|
+
prompt_id: str | None = None,
|
|
2740
|
+
gallery_item_id: str | None = None,
|
|
2741
|
+
output_slots: list[ServeRunOutputSlot] | None = None,
|
|
2742
|
+
created_at: str | None = None,
|
|
2743
|
+
) -> dict[str, Any]:
|
|
2744
|
+
run_id = run_id or f"run_{uuid.uuid4().hex}"
|
|
2745
|
+
created_at = created_at or utc_now()
|
|
2746
|
+
inputs = _extract_run_inputs(body)
|
|
2747
|
+
message = str(payload.get("message") or payload.get("error") or "Generation failed")
|
|
2748
|
+
raw_result = dict(payload)
|
|
2749
|
+
state.state_store.record_run(
|
|
2750
|
+
ServeRunRecord(
|
|
2751
|
+
run_id=run_id,
|
|
2752
|
+
session_id=session.session_id,
|
|
2753
|
+
scope_key=session.scope_key,
|
|
2754
|
+
workflow=workflow_name,
|
|
2755
|
+
contract=contract_name,
|
|
2756
|
+
status="error",
|
|
2757
|
+
inputs=_display_inputs(inputs),
|
|
2758
|
+
prompt_id=prompt_id,
|
|
2759
|
+
raw_result=raw_result,
|
|
2760
|
+
error=message,
|
|
2761
|
+
created_at=created_at,
|
|
2762
|
+
)
|
|
2763
|
+
)
|
|
2764
|
+
slots = output_slots or []
|
|
2765
|
+
errored_slots: list[ServeRunOutputSlot] = []
|
|
2766
|
+
if slots:
|
|
2767
|
+
errored_slots = [
|
|
2768
|
+
_copy_output_slot(slot, status="error", prompt_id=prompt_id, error=message, raw_result=raw_result)
|
|
2769
|
+
for slot in slots
|
|
2770
|
+
]
|
|
2771
|
+
state.state_store.record_output_slots(errored_slots)
|
|
2772
|
+
gallery_items = _error_gallery_items_for_slots(
|
|
2773
|
+
errored_slots,
|
|
2774
|
+
inputs=inputs,
|
|
2775
|
+
raw_result=raw_result,
|
|
2776
|
+
message=message,
|
|
2777
|
+
)
|
|
2778
|
+
else:
|
|
2779
|
+
gallery_items = [
|
|
2780
|
+
ServeGalleryItem(
|
|
2781
|
+
item_id=gallery_item_id or f"gallery_{uuid.uuid4().hex}",
|
|
2782
|
+
run_id=run_id,
|
|
2783
|
+
session_id=session.session_id,
|
|
2784
|
+
scope_key=session.scope_key,
|
|
2785
|
+
workflow=workflow_name,
|
|
2786
|
+
contract=contract_name,
|
|
2787
|
+
status="error",
|
|
2788
|
+
output_type="image",
|
|
2789
|
+
inputs=_display_inputs(inputs),
|
|
2790
|
+
prompt_id=prompt_id,
|
|
2791
|
+
width=1,
|
|
2792
|
+
height=1,
|
|
2793
|
+
raw_result=raw_result,
|
|
2794
|
+
error=message,
|
|
2795
|
+
created_at=created_at,
|
|
2796
|
+
updated_at=created_at,
|
|
2797
|
+
)
|
|
2798
|
+
]
|
|
2799
|
+
state.state_store.record_gallery_items(gallery_items)
|
|
2800
|
+
return {
|
|
2801
|
+
"run_id": run_id,
|
|
2802
|
+
"output_slots": [slot.to_public_dict() for slot in errored_slots],
|
|
2803
|
+
"gallery_items": [item.to_public_dict() for item in gallery_items],
|
|
2804
|
+
}
|
|
2805
|
+
|
|
2806
|
+
|
|
2807
|
+
def _error_gallery_items_for_slots(
|
|
2808
|
+
slots: list[ServeRunOutputSlot],
|
|
2809
|
+
*,
|
|
2810
|
+
inputs: dict[str, Any],
|
|
2811
|
+
raw_result: dict[str, Any],
|
|
2812
|
+
message: str,
|
|
2813
|
+
) -> list[ServeGalleryItem]:
|
|
2814
|
+
display_inputs = _display_inputs(inputs)
|
|
2815
|
+
return [
|
|
2816
|
+
ServeGalleryItem(
|
|
2817
|
+
item_id=_gallery_item_id_for_slot(slot),
|
|
2818
|
+
run_id=slot.run_id,
|
|
2819
|
+
session_id=slot.session_id,
|
|
2820
|
+
scope_key=slot.scope_key,
|
|
2821
|
+
workflow=slot.workflow,
|
|
2822
|
+
contract=slot.contract,
|
|
2823
|
+
status="error",
|
|
2824
|
+
output_type=slot.output_type if slot.output_type in {"image", "video", "audio", "json"} else "image",
|
|
2825
|
+
slot_id=slot.slot_id,
|
|
2826
|
+
output_name=slot.output_name,
|
|
2827
|
+
inputs=display_inputs,
|
|
2828
|
+
prompt_id=slot.prompt_id,
|
|
2829
|
+
width=slot.width,
|
|
2830
|
+
height=slot.height,
|
|
2831
|
+
raw_result=raw_result,
|
|
2832
|
+
error=message,
|
|
2833
|
+
created_at=slot.created_at,
|
|
2834
|
+
updated_at=slot.created_at,
|
|
2835
|
+
)
|
|
2836
|
+
for slot in slots
|
|
2837
|
+
]
|
|
2838
|
+
|
|
2839
|
+
|
|
2840
|
+
def _empty_gallery_items_for_slots(
|
|
2841
|
+
slots: list[ServeRunOutputSlot],
|
|
2842
|
+
*,
|
|
2843
|
+
existing_items: list[ServeGalleryItem],
|
|
2844
|
+
inputs: dict[str, Any],
|
|
2845
|
+
response: dict[str, Any],
|
|
2846
|
+
) -> list[ServeGalleryItem]:
|
|
2847
|
+
existing_slot_ids = {item.slot_id for item in existing_items if item.slot_id}
|
|
2848
|
+
display_inputs = _display_inputs(inputs)
|
|
2849
|
+
return [
|
|
2850
|
+
ServeGalleryItem(
|
|
2851
|
+
item_id=_gallery_item_id_for_slot(slot),
|
|
2852
|
+
run_id=slot.run_id,
|
|
2853
|
+
session_id=slot.session_id,
|
|
2854
|
+
scope_key=slot.scope_key,
|
|
2855
|
+
workflow=slot.workflow,
|
|
2856
|
+
contract=slot.contract,
|
|
2857
|
+
status="done",
|
|
2858
|
+
output_type="json",
|
|
2859
|
+
slot_id=slot.slot_id,
|
|
2860
|
+
output_name=slot.output_name,
|
|
2861
|
+
inputs=display_inputs,
|
|
2862
|
+
prompt_id=slot.prompt_id,
|
|
2863
|
+
width=1,
|
|
2864
|
+
height=1,
|
|
2865
|
+
raw_result=dict(response),
|
|
2866
|
+
created_at=slot.created_at,
|
|
2867
|
+
updated_at=slot.created_at,
|
|
2868
|
+
)
|
|
2869
|
+
for slot in slots
|
|
2870
|
+
if slot.status == "empty" and slot.slot_id not in existing_slot_ids
|
|
2871
|
+
]
|
|
2872
|
+
|
|
2873
|
+
|
|
2874
|
+
def _gallery_items_for_outputs(
|
|
2875
|
+
*,
|
|
2876
|
+
run_id: str,
|
|
2877
|
+
session: ServeSession,
|
|
2878
|
+
workflow_name: str,
|
|
2879
|
+
contract_name: str,
|
|
2880
|
+
inputs: dict[str, Any],
|
|
2881
|
+
response: dict[str, Any],
|
|
2882
|
+
output_slots: list[ServeRunOutputSlot] | None = None,
|
|
2883
|
+
created_at: str | None = None,
|
|
2884
|
+
) -> list[ServeGalleryItem]:
|
|
2885
|
+
items: list[ServeGalleryItem] = []
|
|
2886
|
+
display_inputs = _display_inputs(inputs)
|
|
2887
|
+
prompt_id = str(response.get("prompt_id") or "")
|
|
2888
|
+
created_at = created_at or utc_now()
|
|
2889
|
+
raw_result = dict(response)
|
|
2890
|
+
slots = output_slots or []
|
|
2891
|
+
for output_index, output in enumerate(response.get("outputs") or []):
|
|
2892
|
+
if not isinstance(output, Mapping):
|
|
2893
|
+
continue
|
|
2894
|
+
slot = slots[output_index] if output_index < len(slots) else None
|
|
2895
|
+
output_name = str(output.get("name") or "output")
|
|
2896
|
+
output_type = str(output.get("type") or "json").lower()
|
|
2897
|
+
artifacts = output.get("artifacts")
|
|
2898
|
+
if not isinstance(artifacts, list):
|
|
2899
|
+
continue
|
|
2900
|
+
for artifact_index, artifact in enumerate(artifacts):
|
|
2901
|
+
if not isinstance(artifact, Mapping):
|
|
2902
|
+
continue
|
|
2903
|
+
artifact_payload = dict(artifact)
|
|
2904
|
+
filename = artifact_payload.get("filename")
|
|
2905
|
+
item_type = output_kind(output_type, str(filename or ""))
|
|
2906
|
+
width, height = _gallery_dimensions_for_artifact(item_type, artifact_payload)
|
|
2907
|
+
item_id = _gallery_item_id_for_slot(slot) if slot and artifact_index == 0 else f"gallery_{uuid.uuid4().hex}"
|
|
2908
|
+
items.append(
|
|
2909
|
+
ServeGalleryItem(
|
|
2910
|
+
item_id=item_id,
|
|
2911
|
+
run_id=run_id,
|
|
2912
|
+
session_id=session.session_id,
|
|
2913
|
+
scope_key=session.scope_key,
|
|
2914
|
+
workflow=workflow_name,
|
|
2915
|
+
contract=contract_name,
|
|
2916
|
+
status="done",
|
|
2917
|
+
output_type=item_type,
|
|
2918
|
+
slot_id=slot.slot_id if slot else None,
|
|
2919
|
+
output_name=output_name,
|
|
2920
|
+
prompt_id=prompt_id or None,
|
|
2921
|
+
filename=str(filename) if filename else None,
|
|
2922
|
+
url=str(artifact_payload.get("url")) if artifact_payload.get("url") else None,
|
|
2923
|
+
width=width if item_type in {"image", "video", "audio"} else 1,
|
|
2924
|
+
height=height if item_type in {"image", "video", "audio"} else 1,
|
|
2925
|
+
inputs=display_inputs,
|
|
2926
|
+
artifact=artifact_payload,
|
|
2927
|
+
raw_result=raw_result,
|
|
2928
|
+
created_at=created_at,
|
|
2929
|
+
updated_at=created_at,
|
|
2930
|
+
)
|
|
2931
|
+
)
|
|
2932
|
+
return items
|
|
2933
|
+
|
|
2934
|
+
|
|
2935
|
+
def _extract_run_inputs(body: Mapping[str, Any]) -> dict[str, Any]:
|
|
2936
|
+
if isinstance(body.get("inputs"), dict):
|
|
2937
|
+
return dict(body["inputs"])
|
|
2938
|
+
control_keys = {"wait", "timeout_seconds", "poll_interval_seconds"}
|
|
2939
|
+
return {key: value for key, value in body.items() if key not in control_keys}
|
|
2940
|
+
|
|
2941
|
+
|
|
2942
|
+
def _display_inputs(inputs: Mapping[str, Any]) -> dict[str, Any]:
|
|
2943
|
+
return {str(key): _display_value(value) for key, value in inputs.items()}
|
|
2944
|
+
|
|
2945
|
+
|
|
2946
|
+
def _display_value(value: Any) -> Any:
|
|
2947
|
+
if isinstance(value, Mapping):
|
|
2948
|
+
if value.get("kind") == "file_ref":
|
|
2949
|
+
return {
|
|
2950
|
+
key: value.get(key)
|
|
2951
|
+
for key in ("kind", "ref", "filename", "mime_type", "size")
|
|
2952
|
+
if key in value
|
|
2953
|
+
}
|
|
2954
|
+
return {str(key): _display_value(child) for key, child in value.items()}
|
|
2955
|
+
if isinstance(value, list):
|
|
2956
|
+
return [_display_value(child) for child in value]
|
|
2957
|
+
if isinstance(value, str) and value.startswith("data:"):
|
|
2958
|
+
return f"{value[:48]}... [inline data omitted]"
|
|
2959
|
+
return value
|
|
2960
|
+
|
|
2961
|
+
|
|
2962
|
+
async def _prepare_contract_inputs(
|
|
2963
|
+
state: ServeState,
|
|
2964
|
+
workflow_name: str,
|
|
2965
|
+
contract_name: str,
|
|
2966
|
+
inputs: dict[str, Any],
|
|
2967
|
+
) -> dict[str, Any]:
|
|
2968
|
+
return (await _prepare_contract_run_inputs(state, workflow_name, contract_name, inputs)).inputs
|
|
2969
|
+
|
|
2970
|
+
|
|
2971
|
+
async def _prepare_contract_run_inputs(
|
|
2972
|
+
state: ServeState,
|
|
2973
|
+
workflow_name: str,
|
|
2974
|
+
contract_name: str,
|
|
2975
|
+
inputs: dict[str, Any],
|
|
2976
|
+
) -> PreparedContractInputs:
|
|
2977
|
+
"""Resolve uploaded media refs before building the ComfyUI prompt."""
|
|
2978
|
+
|
|
2979
|
+
manifest = state.manifest_snapshot()
|
|
2980
|
+
workflow = manifest.workflows.get(workflow_name)
|
|
2981
|
+
execution_contract = getattr(workflow, "execution_contract", None) if workflow else None
|
|
2982
|
+
contract = execution_contract.contracts.get(contract_name) if execution_contract else None
|
|
2983
|
+
if contract is None:
|
|
2984
|
+
return PreparedContractInputs(dict(inputs))
|
|
2985
|
+
|
|
2986
|
+
prepared = dict(inputs)
|
|
2987
|
+
staged_uploads: list[StagedUpload] = []
|
|
2988
|
+
for contract_input in contract.inputs:
|
|
2989
|
+
if str(contract_input.type).lower() not in FILE_UPLOAD_CONTRACT_INPUT_TYPES:
|
|
2990
|
+
continue
|
|
2991
|
+
if contract_input.name not in prepared:
|
|
2992
|
+
continue
|
|
2993
|
+
resolved, record = _resolve_upload_binding(
|
|
2994
|
+
state,
|
|
2995
|
+
prepared[contract_input.name],
|
|
2996
|
+
input_name=contract_input.name,
|
|
2997
|
+
)
|
|
2998
|
+
prepared[contract_input.name] = resolved
|
|
2999
|
+
if record is not None:
|
|
3000
|
+
staged_upload = _staged_upload_from_record(contract_input.name, record)
|
|
3001
|
+
if staged_upload is not None:
|
|
3002
|
+
staged_uploads.append(staged_upload)
|
|
3003
|
+
return PreparedContractInputs(prepared, tuple(staged_uploads))
|
|
3004
|
+
|
|
3005
|
+
|
|
3006
|
+
def _staged_upload_from_record(input_name: str, record: Any) -> StagedUpload | None:
|
|
3007
|
+
path = getattr(record, "path", None)
|
|
3008
|
+
filename = getattr(record, "filename", None)
|
|
3009
|
+
content_type = getattr(record, "content_type", None)
|
|
3010
|
+
comfyui_filename = getattr(record, "comfyui_filename", None)
|
|
3011
|
+
if path is None or filename is None or content_type is None or comfyui_filename is None:
|
|
3012
|
+
return None
|
|
3013
|
+
return StagedUpload(
|
|
3014
|
+
input_name=input_name,
|
|
3015
|
+
path=Path(path),
|
|
3016
|
+
filename=str(filename),
|
|
3017
|
+
content_type=str(content_type),
|
|
3018
|
+
size=getattr(record, "size", None),
|
|
3019
|
+
comfyui_filename=str(comfyui_filename),
|
|
3020
|
+
)
|
|
3021
|
+
|
|
3022
|
+
|
|
3023
|
+
def _prepare_upload_slot(state: ServeState, body: Mapping[str, Any]) -> UploadRecord:
|
|
3024
|
+
requested_content_type = _canonical_content_type(body.get("mime_type") or body.get("content_type"))
|
|
3025
|
+
filename = _safe_upload_filename(body.get("filename"), requested_content_type)
|
|
3026
|
+
content_type = requested_content_type or _content_type_for_filename(filename)
|
|
3027
|
+
size = _optional_positive_int(body.get("size"))
|
|
3028
|
+
max_bytes = _max_request_bytes(state)
|
|
3029
|
+
if size is not None and size > max_bytes:
|
|
3030
|
+
raise ValueError(f"Upload is too large. This cg serve instance accepts files up to {max_bytes} bytes.")
|
|
3031
|
+
|
|
3032
|
+
upload_id = f"upload_{uuid.uuid4().hex}"
|
|
3033
|
+
stored_filename = f"{upload_id}_{filename}"
|
|
3034
|
+
input_dir = _comfyui_input_dir(state.env)
|
|
3035
|
+
record = UploadRecord(
|
|
3036
|
+
upload_id=upload_id,
|
|
3037
|
+
token=secrets.token_urlsafe(UPLOAD_TOKEN_BYTES),
|
|
3038
|
+
filename=filename,
|
|
3039
|
+
content_type=content_type,
|
|
3040
|
+
size=size,
|
|
3041
|
+
path=input_dir / stored_filename,
|
|
3042
|
+
comfyui_filename=stored_filename,
|
|
3043
|
+
)
|
|
3044
|
+
state.uploads[upload_id] = record
|
|
3045
|
+
return record
|
|
3046
|
+
|
|
3047
|
+
|
|
3048
|
+
def _resolve_upload_ref(state: ServeState, value: Any, *, input_name: str) -> Any:
|
|
3049
|
+
return _resolve_upload_binding(state, value, input_name=input_name)[0]
|
|
3050
|
+
|
|
3051
|
+
|
|
3052
|
+
def _resolve_upload_binding(state: ServeState, value: Any, *, input_name: str) -> tuple[Any, UploadRecord | None]:
|
|
3053
|
+
if isinstance(value, str):
|
|
3054
|
+
if value.startswith("data:"):
|
|
3055
|
+
raise ValueError(
|
|
3056
|
+
f"Input '{input_name}' uses an inline data URL. Upload the file first and submit a file_ref."
|
|
3057
|
+
)
|
|
3058
|
+
return value, None
|
|
3059
|
+
|
|
3060
|
+
if not isinstance(value, Mapping):
|
|
3061
|
+
return value, None
|
|
3062
|
+
|
|
3063
|
+
if value.get("kind") == "file_ref":
|
|
3064
|
+
upload_id = value.get("ref") or value.get("upload_id")
|
|
3065
|
+
if not isinstance(upload_id, str) or not upload_id:
|
|
3066
|
+
raise ValueError(f"Input '{input_name}' file_ref is missing a ref.")
|
|
3067
|
+
record = state.uploads.get(upload_id)
|
|
3068
|
+
if record is None:
|
|
3069
|
+
raise ValueError(f"Input '{input_name}' references an unknown upload.")
|
|
3070
|
+
if record.status != "ready":
|
|
3071
|
+
raise ValueError(f"Input '{input_name}' references an upload that is not ready.")
|
|
3072
|
+
return record.comfyui_filename, record
|
|
3073
|
+
|
|
3074
|
+
if any(key in value for key in ("data_url", "base64", "data")):
|
|
3075
|
+
raise ValueError(
|
|
3076
|
+
f"Input '{input_name}' uses inline file bytes. Upload the file first and submit a file_ref."
|
|
3077
|
+
)
|
|
3078
|
+
|
|
3079
|
+
return value, None
|
|
3080
|
+
|
|
3081
|
+
|
|
3082
|
+
def _safe_upload_filename(value: Any, content_type: Any = None) -> str:
|
|
3083
|
+
raw = str(value or "").strip().replace("\\", "/").split("/")[-1]
|
|
3084
|
+
safe = "".join(char for char in raw if char.isalnum() or char in {"-", "_", "."}).strip(" .")
|
|
3085
|
+
if safe and "." in safe and any(char.isalnum() for char in safe):
|
|
3086
|
+
return safe
|
|
3087
|
+
return _generated_upload_filename(str(content_type or ""), stem=safe or None)
|
|
3088
|
+
|
|
3089
|
+
|
|
3090
|
+
def _generated_upload_filename(content_type: str, stem: str | None = None) -> str:
|
|
3091
|
+
extension = _extension_for_content_type(content_type)
|
|
3092
|
+
return f"{stem or f'comfygit-upload-{uuid.uuid4().hex}'}{extension}"
|
|
3093
|
+
|
|
3094
|
+
|
|
3095
|
+
def _content_type_for_filename(filename: str) -> str:
|
|
3096
|
+
extension = Path(filename).suffix.lower()
|
|
3097
|
+
return UPLOAD_MIME_TYPE_BY_EXTENSION.get(extension, DEFAULT_UPLOAD_CONTENT_TYPE)
|
|
3098
|
+
|
|
3099
|
+
|
|
3100
|
+
def _extension_for_content_type(content_type: str) -> str:
|
|
3101
|
+
return UPLOAD_EXTENSION_BY_MIME_TYPE.get(
|
|
3102
|
+
_canonical_content_type(content_type),
|
|
3103
|
+
DEFAULT_UPLOAD_EXTENSION,
|
|
3104
|
+
)
|
|
3105
|
+
|
|
3106
|
+
|
|
3107
|
+
def _canonical_content_type(value: Any) -> str:
|
|
3108
|
+
content_type = str(value or "").split(";", 1)[0].strip().lower()
|
|
3109
|
+
return UPLOAD_MIME_TYPE_ALIASES.get(content_type, content_type)
|
|
3110
|
+
|
|
3111
|
+
|
|
3112
|
+
def _optional_positive_int(value: Any) -> int | None:
|
|
3113
|
+
if value is None:
|
|
3114
|
+
return None
|
|
3115
|
+
try:
|
|
3116
|
+
parsed = int(value)
|
|
3117
|
+
except (TypeError, ValueError):
|
|
3118
|
+
return None
|
|
3119
|
+
return parsed if parsed >= 0 else None
|
|
3120
|
+
|
|
3121
|
+
|
|
3122
|
+
def _optional_positive_int_query(
|
|
3123
|
+
request: web.Request,
|
|
3124
|
+
name: str,
|
|
3125
|
+
*,
|
|
3126
|
+
max_value: int | None = None,
|
|
3127
|
+
) -> int | None:
|
|
3128
|
+
value = request.query.get(name)
|
|
3129
|
+
if value in {None, ""}:
|
|
3130
|
+
return None
|
|
3131
|
+
try:
|
|
3132
|
+
parsed = int(str(value))
|
|
3133
|
+
except ValueError as exc:
|
|
3134
|
+
raise ValueError(f"'{name}' must be a positive integer.") from exc
|
|
3135
|
+
if parsed < 1:
|
|
3136
|
+
raise ValueError(f"'{name}' must be a positive integer.")
|
|
3137
|
+
if max_value is not None and parsed > max_value:
|
|
3138
|
+
raise ValueError(f"'{name}' must be less than or equal to {max_value}.")
|
|
3139
|
+
return parsed
|
|
3140
|
+
|
|
3141
|
+
|
|
3142
|
+
def _comfyui_input_dir(env: Environment) -> Path:
|
|
3143
|
+
comfyui_path = Path(getattr(env, "comfyui_path", Path(getattr(env, "path", ".")) / "ComfyUI"))
|
|
3144
|
+
return comfyui_path / "input"
|
|
3145
|
+
|
|
3146
|
+
|
|
3147
|
+
def _upload_too_large_response(max_bytes: int) -> web.Response:
|
|
3148
|
+
max_mib = max_bytes // (1024 * 1024)
|
|
3149
|
+
return web.json_response(
|
|
3150
|
+
{
|
|
3151
|
+
"error": "request_too_large",
|
|
3152
|
+
"message": f"Upload is too large. This cg serve instance accepts uploads up to {max_mib} MiB.",
|
|
3153
|
+
},
|
|
3154
|
+
status=413,
|
|
3155
|
+
)
|