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
comfygit_studio/state.py
ADDED
|
@@ -0,0 +1,1019 @@
|
|
|
1
|
+
"""Runtime state adapters for ComfyGit Studio.
|
|
2
|
+
|
|
3
|
+
This module is intentionally part of the Studio runtime package, not core. The
|
|
4
|
+
data recorded here describes Studio sessions and runtime output history, not
|
|
5
|
+
portable environment truth.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import base64
|
|
11
|
+
import binascii
|
|
12
|
+
import json
|
|
13
|
+
import sqlite3
|
|
14
|
+
from dataclasses import dataclass
|
|
15
|
+
from datetime import UTC, datetime
|
|
16
|
+
from pathlib import Path
|
|
17
|
+
from typing import Any
|
|
18
|
+
|
|
19
|
+
SERVE_STATE_SCHEMA_VERSION = 2
|
|
20
|
+
GALLERY_CURSOR_VERSION = 1
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def utc_now() -> str:
|
|
24
|
+
return datetime.now(UTC).isoformat(timespec="milliseconds").replace("+00:00", "Z")
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
@dataclass(frozen=True)
|
|
28
|
+
class ServeSession:
|
|
29
|
+
session_id: str
|
|
30
|
+
scope_key: str
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
@dataclass(frozen=True)
|
|
34
|
+
class ServeRunRecord:
|
|
35
|
+
run_id: str
|
|
36
|
+
session_id: str
|
|
37
|
+
scope_key: str
|
|
38
|
+
workflow: str
|
|
39
|
+
contract: str
|
|
40
|
+
status: str
|
|
41
|
+
inputs: dict[str, Any]
|
|
42
|
+
prompt_id: str | None = None
|
|
43
|
+
raw_result: dict[str, Any] | None = None
|
|
44
|
+
error: str | None = None
|
|
45
|
+
created_at: str = ""
|
|
46
|
+
updated_at: str = ""
|
|
47
|
+
|
|
48
|
+
def to_public_dict(self) -> dict[str, Any]:
|
|
49
|
+
payload: dict[str, Any] = {
|
|
50
|
+
"run_id": self.run_id,
|
|
51
|
+
"session_id": self.session_id,
|
|
52
|
+
"workflow": self.workflow,
|
|
53
|
+
"contract": self.contract,
|
|
54
|
+
"status": self.status,
|
|
55
|
+
"inputs": self.inputs,
|
|
56
|
+
"createdAt": self.created_at,
|
|
57
|
+
"updatedAt": self.updated_at,
|
|
58
|
+
}
|
|
59
|
+
optional = {
|
|
60
|
+
"prompt_id": self.prompt_id,
|
|
61
|
+
"raw_result": self.raw_result,
|
|
62
|
+
"error": self.error,
|
|
63
|
+
}
|
|
64
|
+
for key, value in optional.items():
|
|
65
|
+
if value is not None:
|
|
66
|
+
payload[key] = value
|
|
67
|
+
return payload
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
@dataclass(frozen=True)
|
|
71
|
+
class ServeRunOutputSlot:
|
|
72
|
+
slot_id: str
|
|
73
|
+
run_id: str
|
|
74
|
+
session_id: str
|
|
75
|
+
scope_key: str
|
|
76
|
+
workflow: str
|
|
77
|
+
contract: str
|
|
78
|
+
output_name: str
|
|
79
|
+
output_type: str
|
|
80
|
+
status: str
|
|
81
|
+
prompt_id: str | None = None
|
|
82
|
+
width: int | None = None
|
|
83
|
+
height: int | None = None
|
|
84
|
+
error: str | None = None
|
|
85
|
+
raw_result: dict[str, Any] | None = None
|
|
86
|
+
created_at: str = ""
|
|
87
|
+
updated_at: str = ""
|
|
88
|
+
|
|
89
|
+
def to_public_dict(self) -> dict[str, Any]:
|
|
90
|
+
payload: dict[str, Any] = {
|
|
91
|
+
"slot_id": self.slot_id,
|
|
92
|
+
"run_id": self.run_id,
|
|
93
|
+
"contract": f"{self.workflow} / {self.contract}",
|
|
94
|
+
"contractWorkflow": self.workflow,
|
|
95
|
+
"contractName": self.contract,
|
|
96
|
+
"outputName": self.output_name,
|
|
97
|
+
"type": self.output_type,
|
|
98
|
+
"status": self.status,
|
|
99
|
+
"createdAt": self.created_at,
|
|
100
|
+
"updatedAt": self.updated_at,
|
|
101
|
+
}
|
|
102
|
+
optional = {
|
|
103
|
+
"promptId": self.prompt_id,
|
|
104
|
+
"width": self.width,
|
|
105
|
+
"height": self.height,
|
|
106
|
+
"error": self.error,
|
|
107
|
+
"rawResult": self.raw_result,
|
|
108
|
+
}
|
|
109
|
+
for key, value in optional.items():
|
|
110
|
+
if value is not None:
|
|
111
|
+
payload[key] = value
|
|
112
|
+
return payload
|
|
113
|
+
|
|
114
|
+
|
|
115
|
+
@dataclass(frozen=True)
|
|
116
|
+
class ServeGalleryItem:
|
|
117
|
+
item_id: str
|
|
118
|
+
run_id: str
|
|
119
|
+
session_id: str
|
|
120
|
+
scope_key: str
|
|
121
|
+
workflow: str
|
|
122
|
+
contract: str
|
|
123
|
+
status: str
|
|
124
|
+
output_type: str
|
|
125
|
+
inputs: dict[str, Any]
|
|
126
|
+
slot_id: str | None = None
|
|
127
|
+
output_name: str | None = None
|
|
128
|
+
prompt_id: str | None = None
|
|
129
|
+
filename: str | None = None
|
|
130
|
+
url: str | None = None
|
|
131
|
+
width: int | None = None
|
|
132
|
+
height: int | None = None
|
|
133
|
+
artifact: dict[str, Any] | None = None
|
|
134
|
+
raw_result: dict[str, Any] | None = None
|
|
135
|
+
error: str | None = None
|
|
136
|
+
created_at: str = ""
|
|
137
|
+
updated_at: str = ""
|
|
138
|
+
|
|
139
|
+
def to_public_dict(self) -> dict[str, Any]:
|
|
140
|
+
payload: dict[str, Any] = {
|
|
141
|
+
"id": self.item_id,
|
|
142
|
+
"run_id": self.run_id,
|
|
143
|
+
"contract": f"{self.workflow} / {self.contract}",
|
|
144
|
+
"contractWorkflow": self.workflow,
|
|
145
|
+
"contractName": self.contract,
|
|
146
|
+
"status": self.status,
|
|
147
|
+
"type": self.output_type,
|
|
148
|
+
"inputs": self.inputs,
|
|
149
|
+
"createdAt": self.created_at,
|
|
150
|
+
}
|
|
151
|
+
optional = {
|
|
152
|
+
"slotId": self.slot_id,
|
|
153
|
+
"promptId": self.prompt_id,
|
|
154
|
+
"outputName": self.output_name,
|
|
155
|
+
"filename": self.filename,
|
|
156
|
+
"url": self.url,
|
|
157
|
+
"width": self.width,
|
|
158
|
+
"height": self.height,
|
|
159
|
+
"artifact": self.artifact,
|
|
160
|
+
"rawResult": self.raw_result,
|
|
161
|
+
"error": self.error,
|
|
162
|
+
}
|
|
163
|
+
for key, value in optional.items():
|
|
164
|
+
if value is not None:
|
|
165
|
+
payload[key] = value
|
|
166
|
+
return payload
|
|
167
|
+
|
|
168
|
+
|
|
169
|
+
@dataclass(frozen=True)
|
|
170
|
+
class GalleryPage:
|
|
171
|
+
items: list[dict[str, Any]]
|
|
172
|
+
next_cursor: str | None
|
|
173
|
+
has_more: bool
|
|
174
|
+
limit: int | None = None
|
|
175
|
+
|
|
176
|
+
|
|
177
|
+
class ServeStateStore:
|
|
178
|
+
"""Interface for serve-owned session/run/gallery runtime state."""
|
|
179
|
+
|
|
180
|
+
persistent = False
|
|
181
|
+
|
|
182
|
+
def close(self) -> None:
|
|
183
|
+
return None
|
|
184
|
+
|
|
185
|
+
def ensure_session(self, session_id: str, *, scope_key: str) -> ServeSession:
|
|
186
|
+
raise NotImplementedError
|
|
187
|
+
|
|
188
|
+
def record_run(self, run: ServeRunRecord) -> None:
|
|
189
|
+
raise NotImplementedError
|
|
190
|
+
|
|
191
|
+
def get_run(self, scope_key: str, run_id: str) -> dict[str, Any] | None:
|
|
192
|
+
raise NotImplementedError
|
|
193
|
+
|
|
194
|
+
def get_run_record(self, run_id: str) -> ServeRunRecord | None:
|
|
195
|
+
raise NotImplementedError
|
|
196
|
+
|
|
197
|
+
def list_runs(self, scope_key: str, statuses: set[str] | None = None) -> list[dict[str, Any]]:
|
|
198
|
+
raise NotImplementedError
|
|
199
|
+
|
|
200
|
+
def list_active_runs(self, statuses: set[str]) -> list[ServeRunRecord]:
|
|
201
|
+
raise NotImplementedError
|
|
202
|
+
|
|
203
|
+
def cancel_run(self, scope_key: str, run_id: str, *, raw_result: dict[str, Any], error: str) -> bool:
|
|
204
|
+
raise NotImplementedError
|
|
205
|
+
|
|
206
|
+
def record_output_slots(self, slots: list[ServeRunOutputSlot]) -> None:
|
|
207
|
+
raise NotImplementedError
|
|
208
|
+
|
|
209
|
+
def list_output_slots(self, scope_key: str, run_id: str) -> list[dict[str, Any]]:
|
|
210
|
+
raise NotImplementedError
|
|
211
|
+
|
|
212
|
+
def record_gallery_items(self, items: list[ServeGalleryItem]) -> None:
|
|
213
|
+
raise NotImplementedError
|
|
214
|
+
|
|
215
|
+
def list_gallery_items(self, scope_key: str) -> list[dict[str, Any]]:
|
|
216
|
+
raise NotImplementedError
|
|
217
|
+
|
|
218
|
+
def list_gallery_page(self, scope_key: str, *, limit: int | None, cursor: str | None) -> GalleryPage:
|
|
219
|
+
raise NotImplementedError
|
|
220
|
+
|
|
221
|
+
def list_gallery_items_for_run(self, scope_key: str, run_id: str) -> list[dict[str, Any]]:
|
|
222
|
+
raise NotImplementedError
|
|
223
|
+
|
|
224
|
+
def delete_gallery_item(self, scope_key: str, item_id: str) -> bool:
|
|
225
|
+
raise NotImplementedError
|
|
226
|
+
|
|
227
|
+
|
|
228
|
+
class EphemeralServeStateStore(ServeStateStore):
|
|
229
|
+
"""In-memory serve state discarded when the process exits."""
|
|
230
|
+
|
|
231
|
+
def __init__(self) -> None:
|
|
232
|
+
self.sessions: dict[str, ServeSession] = {}
|
|
233
|
+
self.runs: dict[str, ServeRunRecord] = {}
|
|
234
|
+
self.output_slots: dict[str, ServeRunOutputSlot] = {}
|
|
235
|
+
self.gallery_items: dict[str, ServeGalleryItem] = {}
|
|
236
|
+
|
|
237
|
+
def ensure_session(self, session_id: str, *, scope_key: str) -> ServeSession:
|
|
238
|
+
session = self.sessions.get(session_id) or ServeSession(session_id=session_id, scope_key=scope_key)
|
|
239
|
+
self.sessions[session_id] = session
|
|
240
|
+
return session
|
|
241
|
+
|
|
242
|
+
def record_run(self, run: ServeRunRecord) -> None:
|
|
243
|
+
self.runs[run.run_id] = _with_run_timestamps(run)
|
|
244
|
+
|
|
245
|
+
def get_run(self, scope_key: str, run_id: str) -> dict[str, Any] | None:
|
|
246
|
+
run = self.runs.get(run_id)
|
|
247
|
+
if run is None or run.scope_key != scope_key:
|
|
248
|
+
return None
|
|
249
|
+
return run.to_public_dict()
|
|
250
|
+
|
|
251
|
+
def get_run_record(self, run_id: str) -> ServeRunRecord | None:
|
|
252
|
+
return self.runs.get(run_id)
|
|
253
|
+
|
|
254
|
+
def list_runs(self, scope_key: str, statuses: set[str] | None = None) -> list[dict[str, Any]]:
|
|
255
|
+
runs = [run for run in self.runs.values() if run.scope_key == scope_key]
|
|
256
|
+
if statuses is not None:
|
|
257
|
+
runs = [run for run in runs if run.status in statuses]
|
|
258
|
+
runs.sort(key=lambda run: run.created_at, reverse=True)
|
|
259
|
+
return [run.to_public_dict() for run in runs]
|
|
260
|
+
|
|
261
|
+
def list_active_runs(self, statuses: set[str]) -> list[ServeRunRecord]:
|
|
262
|
+
runs = [run for run in self.runs.values() if run.status in statuses]
|
|
263
|
+
runs.sort(key=lambda run: run.created_at, reverse=True)
|
|
264
|
+
return runs
|
|
265
|
+
|
|
266
|
+
def cancel_run(self, scope_key: str, run_id: str, *, raw_result: dict[str, Any], error: str) -> bool:
|
|
267
|
+
run = self.runs.get(run_id)
|
|
268
|
+
if run is None or run.scope_key != scope_key or run.status not in {"submitted", "running"}:
|
|
269
|
+
return False
|
|
270
|
+
self.runs[run_id] = _with_run_timestamps(
|
|
271
|
+
ServeRunRecord(
|
|
272
|
+
run_id=run.run_id,
|
|
273
|
+
session_id=run.session_id,
|
|
274
|
+
scope_key=run.scope_key,
|
|
275
|
+
workflow=run.workflow,
|
|
276
|
+
contract=run.contract,
|
|
277
|
+
status="cancelled",
|
|
278
|
+
inputs=run.inputs,
|
|
279
|
+
prompt_id=run.prompt_id,
|
|
280
|
+
raw_result=raw_result,
|
|
281
|
+
error=error,
|
|
282
|
+
created_at=run.created_at,
|
|
283
|
+
)
|
|
284
|
+
)
|
|
285
|
+
for slot_id, slot in list(self.output_slots.items()):
|
|
286
|
+
if slot.scope_key == scope_key and slot.run_id == run_id and slot.status in {"pending", "running"}:
|
|
287
|
+
self.output_slots[slot_id] = _with_output_slot_timestamps(
|
|
288
|
+
ServeRunOutputSlot(
|
|
289
|
+
slot_id=slot.slot_id,
|
|
290
|
+
run_id=slot.run_id,
|
|
291
|
+
session_id=slot.session_id,
|
|
292
|
+
scope_key=slot.scope_key,
|
|
293
|
+
workflow=slot.workflow,
|
|
294
|
+
contract=slot.contract,
|
|
295
|
+
output_name=slot.output_name,
|
|
296
|
+
output_type=slot.output_type,
|
|
297
|
+
status="cancelled",
|
|
298
|
+
prompt_id=slot.prompt_id,
|
|
299
|
+
width=slot.width,
|
|
300
|
+
height=slot.height,
|
|
301
|
+
error=error,
|
|
302
|
+
raw_result=raw_result,
|
|
303
|
+
created_at=slot.created_at,
|
|
304
|
+
)
|
|
305
|
+
)
|
|
306
|
+
for item_id, item in list(self.gallery_items.items()):
|
|
307
|
+
if item.scope_key == scope_key and item.run_id == run_id and item.status == "pending":
|
|
308
|
+
del self.gallery_items[item_id]
|
|
309
|
+
return True
|
|
310
|
+
|
|
311
|
+
def record_output_slots(self, slots: list[ServeRunOutputSlot]) -> None:
|
|
312
|
+
for slot in slots:
|
|
313
|
+
stamped = _with_output_slot_timestamps(slot)
|
|
314
|
+
self.output_slots[stamped.slot_id] = stamped
|
|
315
|
+
|
|
316
|
+
def list_output_slots(self, scope_key: str, run_id: str) -> list[dict[str, Any]]:
|
|
317
|
+
slots = [
|
|
318
|
+
slot
|
|
319
|
+
for slot in self.output_slots.values()
|
|
320
|
+
if slot.scope_key == scope_key and slot.run_id == run_id
|
|
321
|
+
]
|
|
322
|
+
slots.sort(key=lambda slot: slot.created_at)
|
|
323
|
+
return [slot.to_public_dict() for slot in slots]
|
|
324
|
+
|
|
325
|
+
def record_gallery_items(self, items: list[ServeGalleryItem]) -> None:
|
|
326
|
+
for item in items:
|
|
327
|
+
stamped = _with_gallery_timestamps(item)
|
|
328
|
+
self.gallery_items[stamped.item_id] = stamped
|
|
329
|
+
|
|
330
|
+
def list_gallery_items(self, scope_key: str) -> list[dict[str, Any]]:
|
|
331
|
+
items = [item for item in self.gallery_items.values() if item.scope_key == scope_key]
|
|
332
|
+
items.sort(key=lambda item: item.created_at, reverse=True)
|
|
333
|
+
return [item.to_public_dict() for item in items]
|
|
334
|
+
|
|
335
|
+
def list_gallery_page(self, scope_key: str, *, limit: int | None, cursor: str | None) -> GalleryPage:
|
|
336
|
+
decoded_cursor = _decode_gallery_cursor(cursor)
|
|
337
|
+
items = [item for item in self.gallery_items.values() if item.scope_key == scope_key]
|
|
338
|
+
items.sort(key=lambda item: (item.created_at, item.item_id), reverse=True)
|
|
339
|
+
if decoded_cursor is not None:
|
|
340
|
+
items = [
|
|
341
|
+
item
|
|
342
|
+
for item in items
|
|
343
|
+
if _gallery_item_after_cursor(
|
|
344
|
+
item,
|
|
345
|
+
created_at=decoded_cursor[0],
|
|
346
|
+
item_id=decoded_cursor[1],
|
|
347
|
+
)
|
|
348
|
+
]
|
|
349
|
+
return _gallery_page_from_items(items, limit=limit)
|
|
350
|
+
|
|
351
|
+
def list_gallery_items_for_run(self, scope_key: str, run_id: str) -> list[dict[str, Any]]:
|
|
352
|
+
items = [
|
|
353
|
+
item
|
|
354
|
+
for item in self.gallery_items.values()
|
|
355
|
+
if item.scope_key == scope_key and item.run_id == run_id
|
|
356
|
+
]
|
|
357
|
+
items.sort(key=lambda item: item.created_at, reverse=True)
|
|
358
|
+
return [item.to_public_dict() for item in items]
|
|
359
|
+
|
|
360
|
+
def delete_gallery_item(self, scope_key: str, item_id: str) -> bool:
|
|
361
|
+
item = self.gallery_items.get(item_id)
|
|
362
|
+
if item is None or item.scope_key != scope_key:
|
|
363
|
+
return False
|
|
364
|
+
del self.gallery_items[item_id]
|
|
365
|
+
return True
|
|
366
|
+
|
|
367
|
+
|
|
368
|
+
class SQLiteServeStateStore(ServeStateStore):
|
|
369
|
+
"""SQLite-backed serve state for local persistent Studio sessions."""
|
|
370
|
+
|
|
371
|
+
persistent = True
|
|
372
|
+
|
|
373
|
+
def __init__(self, database_path: Path) -> None:
|
|
374
|
+
self.database_path = database_path
|
|
375
|
+
self.database_path.parent.mkdir(parents=True, exist_ok=True)
|
|
376
|
+
self.connection = sqlite3.connect(self.database_path)
|
|
377
|
+
self.connection.row_factory = sqlite3.Row
|
|
378
|
+
self._init_schema()
|
|
379
|
+
|
|
380
|
+
def close(self) -> None:
|
|
381
|
+
self.connection.close()
|
|
382
|
+
|
|
383
|
+
def ensure_session(self, session_id: str, *, scope_key: str) -> ServeSession:
|
|
384
|
+
now = utc_now()
|
|
385
|
+
with self.connection:
|
|
386
|
+
self.connection.execute(
|
|
387
|
+
"""
|
|
388
|
+
INSERT INTO sessions (session_id, scope_key, created_at, updated_at)
|
|
389
|
+
VALUES (?, ?, ?, ?)
|
|
390
|
+
ON CONFLICT(session_id) DO UPDATE SET
|
|
391
|
+
scope_key = excluded.scope_key,
|
|
392
|
+
updated_at = excluded.updated_at
|
|
393
|
+
""",
|
|
394
|
+
(session_id, scope_key, now, now),
|
|
395
|
+
)
|
|
396
|
+
return ServeSession(session_id=session_id, scope_key=scope_key)
|
|
397
|
+
|
|
398
|
+
def record_run(self, run: ServeRunRecord) -> None:
|
|
399
|
+
run = _with_run_timestamps(run)
|
|
400
|
+
with self.connection:
|
|
401
|
+
self.connection.execute(
|
|
402
|
+
"""
|
|
403
|
+
INSERT INTO runs (
|
|
404
|
+
run_id, session_id, scope_key, workflow, contract, status,
|
|
405
|
+
prompt_id, inputs_json, raw_result_json, error, created_at, updated_at
|
|
406
|
+
)
|
|
407
|
+
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
|
408
|
+
ON CONFLICT(run_id) DO UPDATE SET
|
|
409
|
+
status = excluded.status,
|
|
410
|
+
prompt_id = excluded.prompt_id,
|
|
411
|
+
raw_result_json = excluded.raw_result_json,
|
|
412
|
+
inputs_json = excluded.inputs_json,
|
|
413
|
+
error = excluded.error,
|
|
414
|
+
updated_at = excluded.updated_at
|
|
415
|
+
""",
|
|
416
|
+
(
|
|
417
|
+
run.run_id,
|
|
418
|
+
run.session_id,
|
|
419
|
+
run.scope_key,
|
|
420
|
+
run.workflow,
|
|
421
|
+
run.contract,
|
|
422
|
+
run.status,
|
|
423
|
+
run.prompt_id,
|
|
424
|
+
_json(run.inputs),
|
|
425
|
+
_json(run.raw_result),
|
|
426
|
+
run.error,
|
|
427
|
+
run.created_at,
|
|
428
|
+
run.updated_at,
|
|
429
|
+
),
|
|
430
|
+
)
|
|
431
|
+
|
|
432
|
+
def get_run(self, scope_key: str, run_id: str) -> dict[str, Any] | None:
|
|
433
|
+
row = self.connection.execute(
|
|
434
|
+
"""
|
|
435
|
+
SELECT *
|
|
436
|
+
FROM runs
|
|
437
|
+
WHERE scope_key = ? AND run_id = ?
|
|
438
|
+
""",
|
|
439
|
+
(scope_key, run_id),
|
|
440
|
+
).fetchone()
|
|
441
|
+
return _run_from_row(row).to_public_dict() if row else None
|
|
442
|
+
|
|
443
|
+
def get_run_record(self, run_id: str) -> ServeRunRecord | None:
|
|
444
|
+
row = self.connection.execute(
|
|
445
|
+
"""
|
|
446
|
+
SELECT *
|
|
447
|
+
FROM runs
|
|
448
|
+
WHERE run_id = ?
|
|
449
|
+
""",
|
|
450
|
+
(run_id,),
|
|
451
|
+
).fetchone()
|
|
452
|
+
return _run_from_row(row) if row else None
|
|
453
|
+
|
|
454
|
+
def list_runs(self, scope_key: str, statuses: set[str] | None = None) -> list[dict[str, Any]]:
|
|
455
|
+
params: list[Any] = [scope_key]
|
|
456
|
+
status_clause = ""
|
|
457
|
+
if statuses:
|
|
458
|
+
placeholders = ", ".join("?" for _ in statuses)
|
|
459
|
+
status_clause = f" AND status IN ({placeholders})"
|
|
460
|
+
params.extend(sorted(statuses))
|
|
461
|
+
rows = self.connection.execute(
|
|
462
|
+
f"""
|
|
463
|
+
SELECT *
|
|
464
|
+
FROM runs
|
|
465
|
+
WHERE scope_key = ?{status_clause}
|
|
466
|
+
ORDER BY created_at DESC, run_id DESC
|
|
467
|
+
""",
|
|
468
|
+
params,
|
|
469
|
+
).fetchall()
|
|
470
|
+
return [_run_from_row(row).to_public_dict() for row in rows]
|
|
471
|
+
|
|
472
|
+
def list_active_runs(self, statuses: set[str]) -> list[ServeRunRecord]:
|
|
473
|
+
if not statuses:
|
|
474
|
+
return []
|
|
475
|
+
placeholders = ", ".join("?" for _ in statuses)
|
|
476
|
+
rows = self.connection.execute(
|
|
477
|
+
f"""
|
|
478
|
+
SELECT *
|
|
479
|
+
FROM runs
|
|
480
|
+
WHERE status IN ({placeholders})
|
|
481
|
+
ORDER BY created_at DESC, run_id DESC
|
|
482
|
+
""",
|
|
483
|
+
sorted(statuses),
|
|
484
|
+
).fetchall()
|
|
485
|
+
return [_run_from_row(row) for row in rows]
|
|
486
|
+
|
|
487
|
+
def cancel_run(self, scope_key: str, run_id: str, *, raw_result: dict[str, Any], error: str) -> bool:
|
|
488
|
+
now = utc_now()
|
|
489
|
+
with self.connection:
|
|
490
|
+
cursor = self.connection.execute(
|
|
491
|
+
"""
|
|
492
|
+
UPDATE runs
|
|
493
|
+
SET status = 'cancelled',
|
|
494
|
+
raw_result_json = ?,
|
|
495
|
+
error = ?,
|
|
496
|
+
updated_at = ?
|
|
497
|
+
WHERE scope_key = ?
|
|
498
|
+
AND run_id = ?
|
|
499
|
+
AND status IN ('submitted', 'running')
|
|
500
|
+
""",
|
|
501
|
+
(_json(raw_result), error, now, scope_key, run_id),
|
|
502
|
+
)
|
|
503
|
+
if cursor.rowcount == 0:
|
|
504
|
+
return False
|
|
505
|
+
self.connection.execute(
|
|
506
|
+
"""
|
|
507
|
+
UPDATE output_slots
|
|
508
|
+
SET status = 'cancelled',
|
|
509
|
+
raw_result_json = ?,
|
|
510
|
+
error = ?,
|
|
511
|
+
updated_at = ?
|
|
512
|
+
WHERE scope_key = ?
|
|
513
|
+
AND run_id = ?
|
|
514
|
+
AND status IN ('pending', 'running')
|
|
515
|
+
""",
|
|
516
|
+
(_json(raw_result), error, now, scope_key, run_id),
|
|
517
|
+
)
|
|
518
|
+
self.connection.execute(
|
|
519
|
+
"""
|
|
520
|
+
DELETE FROM gallery_items
|
|
521
|
+
WHERE scope_key = ?
|
|
522
|
+
AND run_id = ?
|
|
523
|
+
AND status = 'pending'
|
|
524
|
+
""",
|
|
525
|
+
(scope_key, run_id),
|
|
526
|
+
)
|
|
527
|
+
return True
|
|
528
|
+
|
|
529
|
+
def record_output_slots(self, slots: list[ServeRunOutputSlot]) -> None:
|
|
530
|
+
with self.connection:
|
|
531
|
+
for slot in slots:
|
|
532
|
+
slot = _with_output_slot_timestamps(slot)
|
|
533
|
+
self.connection.execute(
|
|
534
|
+
"""
|
|
535
|
+
INSERT INTO output_slots (
|
|
536
|
+
slot_id, run_id, session_id, scope_key, workflow, contract,
|
|
537
|
+
output_name, output_type, status, prompt_id, width, height,
|
|
538
|
+
error, raw_result_json, created_at, updated_at
|
|
539
|
+
)
|
|
540
|
+
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
|
541
|
+
ON CONFLICT(slot_id) DO UPDATE SET
|
|
542
|
+
status = excluded.status,
|
|
543
|
+
prompt_id = excluded.prompt_id,
|
|
544
|
+
width = excluded.width,
|
|
545
|
+
height = excluded.height,
|
|
546
|
+
error = excluded.error,
|
|
547
|
+
raw_result_json = excluded.raw_result_json,
|
|
548
|
+
updated_at = excluded.updated_at
|
|
549
|
+
""",
|
|
550
|
+
(
|
|
551
|
+
slot.slot_id,
|
|
552
|
+
slot.run_id,
|
|
553
|
+
slot.session_id,
|
|
554
|
+
slot.scope_key,
|
|
555
|
+
slot.workflow,
|
|
556
|
+
slot.contract,
|
|
557
|
+
slot.output_name,
|
|
558
|
+
slot.output_type,
|
|
559
|
+
slot.status,
|
|
560
|
+
slot.prompt_id,
|
|
561
|
+
slot.width,
|
|
562
|
+
slot.height,
|
|
563
|
+
slot.error,
|
|
564
|
+
_json(slot.raw_result),
|
|
565
|
+
slot.created_at,
|
|
566
|
+
slot.updated_at,
|
|
567
|
+
),
|
|
568
|
+
)
|
|
569
|
+
|
|
570
|
+
def list_output_slots(self, scope_key: str, run_id: str) -> list[dict[str, Any]]:
|
|
571
|
+
rows = self.connection.execute(
|
|
572
|
+
"""
|
|
573
|
+
SELECT *
|
|
574
|
+
FROM output_slots
|
|
575
|
+
WHERE scope_key = ? AND run_id = ?
|
|
576
|
+
ORDER BY created_at ASC, slot_id ASC
|
|
577
|
+
""",
|
|
578
|
+
(scope_key, run_id),
|
|
579
|
+
).fetchall()
|
|
580
|
+
return [_output_slot_from_row(row).to_public_dict() for row in rows]
|
|
581
|
+
|
|
582
|
+
def record_gallery_items(self, items: list[ServeGalleryItem]) -> None:
|
|
583
|
+
with self.connection:
|
|
584
|
+
for item in items:
|
|
585
|
+
item = _with_gallery_timestamps(item)
|
|
586
|
+
self.connection.execute(
|
|
587
|
+
"""
|
|
588
|
+
INSERT INTO gallery_items (
|
|
589
|
+
item_id, run_id, session_id, scope_key, workflow, contract,
|
|
590
|
+
status, output_type, slot_id, output_name, prompt_id, filename, url,
|
|
591
|
+
width, height, inputs_json, artifact_json, raw_result_json,
|
|
592
|
+
error, created_at, updated_at
|
|
593
|
+
)
|
|
594
|
+
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
|
595
|
+
ON CONFLICT(item_id) DO UPDATE SET
|
|
596
|
+
status = excluded.status,
|
|
597
|
+
output_type = excluded.output_type,
|
|
598
|
+
slot_id = excluded.slot_id,
|
|
599
|
+
output_name = excluded.output_name,
|
|
600
|
+
prompt_id = excluded.prompt_id,
|
|
601
|
+
filename = excluded.filename,
|
|
602
|
+
url = excluded.url,
|
|
603
|
+
width = excluded.width,
|
|
604
|
+
height = excluded.height,
|
|
605
|
+
inputs_json = excluded.inputs_json,
|
|
606
|
+
artifact_json = excluded.artifact_json,
|
|
607
|
+
raw_result_json = excluded.raw_result_json,
|
|
608
|
+
error = excluded.error,
|
|
609
|
+
updated_at = excluded.updated_at
|
|
610
|
+
""",
|
|
611
|
+
(
|
|
612
|
+
item.item_id,
|
|
613
|
+
item.run_id,
|
|
614
|
+
item.session_id,
|
|
615
|
+
item.scope_key,
|
|
616
|
+
item.workflow,
|
|
617
|
+
item.contract,
|
|
618
|
+
item.status,
|
|
619
|
+
item.output_type,
|
|
620
|
+
item.slot_id,
|
|
621
|
+
item.output_name,
|
|
622
|
+
item.prompt_id,
|
|
623
|
+
item.filename,
|
|
624
|
+
item.url,
|
|
625
|
+
item.width,
|
|
626
|
+
item.height,
|
|
627
|
+
_json(item.inputs),
|
|
628
|
+
_json(item.artifact),
|
|
629
|
+
_json(item.raw_result),
|
|
630
|
+
item.error,
|
|
631
|
+
item.created_at,
|
|
632
|
+
item.updated_at,
|
|
633
|
+
),
|
|
634
|
+
)
|
|
635
|
+
|
|
636
|
+
def list_gallery_items(self, scope_key: str) -> list[dict[str, Any]]:
|
|
637
|
+
rows = self.connection.execute(
|
|
638
|
+
"""
|
|
639
|
+
SELECT *
|
|
640
|
+
FROM gallery_items
|
|
641
|
+
WHERE scope_key = ?
|
|
642
|
+
ORDER BY created_at DESC, item_id DESC
|
|
643
|
+
""",
|
|
644
|
+
(scope_key,),
|
|
645
|
+
).fetchall()
|
|
646
|
+
return [_gallery_item_from_row(row).to_public_dict() for row in rows]
|
|
647
|
+
|
|
648
|
+
def list_gallery_page(self, scope_key: str, *, limit: int | None, cursor: str | None) -> GalleryPage:
|
|
649
|
+
decoded_cursor = _decode_gallery_cursor(cursor)
|
|
650
|
+
params: list[Any] = [scope_key]
|
|
651
|
+
cursor_clause = ""
|
|
652
|
+
if decoded_cursor is not None:
|
|
653
|
+
cursor_clause = """
|
|
654
|
+
AND (
|
|
655
|
+
created_at < ?
|
|
656
|
+
OR (created_at = ? AND item_id < ?)
|
|
657
|
+
)
|
|
658
|
+
"""
|
|
659
|
+
params.extend([decoded_cursor[0], decoded_cursor[0], decoded_cursor[1]])
|
|
660
|
+
limit_clause = ""
|
|
661
|
+
if limit is not None:
|
|
662
|
+
limit_clause = "LIMIT ?"
|
|
663
|
+
params.append(limit + 1)
|
|
664
|
+
rows = self.connection.execute(
|
|
665
|
+
f"""
|
|
666
|
+
SELECT *
|
|
667
|
+
FROM gallery_items
|
|
668
|
+
WHERE scope_key = ?
|
|
669
|
+
{cursor_clause}
|
|
670
|
+
ORDER BY created_at DESC, item_id DESC
|
|
671
|
+
{limit_clause}
|
|
672
|
+
""",
|
|
673
|
+
params,
|
|
674
|
+
).fetchall()
|
|
675
|
+
return _gallery_page_from_items([_gallery_item_from_row(row) for row in rows], limit=limit)
|
|
676
|
+
|
|
677
|
+
def list_gallery_items_for_run(self, scope_key: str, run_id: str) -> list[dict[str, Any]]:
|
|
678
|
+
rows = self.connection.execute(
|
|
679
|
+
"""
|
|
680
|
+
SELECT *
|
|
681
|
+
FROM gallery_items
|
|
682
|
+
WHERE scope_key = ? AND run_id = ?
|
|
683
|
+
ORDER BY created_at DESC, item_id DESC
|
|
684
|
+
""",
|
|
685
|
+
(scope_key, run_id),
|
|
686
|
+
).fetchall()
|
|
687
|
+
return [_gallery_item_from_row(row).to_public_dict() for row in rows]
|
|
688
|
+
|
|
689
|
+
def delete_gallery_item(self, scope_key: str, item_id: str) -> bool:
|
|
690
|
+
with self.connection:
|
|
691
|
+
cursor = self.connection.execute(
|
|
692
|
+
"DELETE FROM gallery_items WHERE scope_key = ? AND item_id = ?",
|
|
693
|
+
(scope_key, item_id),
|
|
694
|
+
)
|
|
695
|
+
return cursor.rowcount > 0
|
|
696
|
+
|
|
697
|
+
def _init_schema(self) -> None:
|
|
698
|
+
with self.connection:
|
|
699
|
+
self.connection.execute("PRAGMA journal_mode=WAL")
|
|
700
|
+
self.connection.execute(
|
|
701
|
+
"""
|
|
702
|
+
CREATE TABLE IF NOT EXISTS serve_state_meta (
|
|
703
|
+
key TEXT PRIMARY KEY,
|
|
704
|
+
value TEXT NOT NULL
|
|
705
|
+
)
|
|
706
|
+
"""
|
|
707
|
+
)
|
|
708
|
+
self.connection.execute(
|
|
709
|
+
"""
|
|
710
|
+
INSERT INTO serve_state_meta (key, value)
|
|
711
|
+
VALUES ('schema_version', ?)
|
|
712
|
+
ON CONFLICT(key) DO UPDATE SET value = excluded.value
|
|
713
|
+
""",
|
|
714
|
+
(str(SERVE_STATE_SCHEMA_VERSION),),
|
|
715
|
+
)
|
|
716
|
+
self.connection.execute(
|
|
717
|
+
"""
|
|
718
|
+
CREATE TABLE IF NOT EXISTS sessions (
|
|
719
|
+
session_id TEXT PRIMARY KEY,
|
|
720
|
+
scope_key TEXT NOT NULL,
|
|
721
|
+
created_at TEXT NOT NULL,
|
|
722
|
+
updated_at TEXT NOT NULL
|
|
723
|
+
)
|
|
724
|
+
"""
|
|
725
|
+
)
|
|
726
|
+
self.connection.execute(
|
|
727
|
+
"""
|
|
728
|
+
CREATE TABLE IF NOT EXISTS runs (
|
|
729
|
+
run_id TEXT PRIMARY KEY,
|
|
730
|
+
session_id TEXT NOT NULL,
|
|
731
|
+
scope_key TEXT NOT NULL,
|
|
732
|
+
workflow TEXT NOT NULL,
|
|
733
|
+
contract TEXT NOT NULL,
|
|
734
|
+
status TEXT NOT NULL,
|
|
735
|
+
prompt_id TEXT,
|
|
736
|
+
inputs_json TEXT NOT NULL,
|
|
737
|
+
raw_result_json TEXT,
|
|
738
|
+
error TEXT,
|
|
739
|
+
created_at TEXT NOT NULL,
|
|
740
|
+
updated_at TEXT NOT NULL
|
|
741
|
+
)
|
|
742
|
+
"""
|
|
743
|
+
)
|
|
744
|
+
self.connection.execute(
|
|
745
|
+
"""
|
|
746
|
+
CREATE TABLE IF NOT EXISTS gallery_items (
|
|
747
|
+
item_id TEXT PRIMARY KEY,
|
|
748
|
+
run_id TEXT NOT NULL,
|
|
749
|
+
session_id TEXT NOT NULL,
|
|
750
|
+
scope_key TEXT NOT NULL,
|
|
751
|
+
workflow TEXT NOT NULL,
|
|
752
|
+
contract TEXT NOT NULL,
|
|
753
|
+
status TEXT NOT NULL,
|
|
754
|
+
output_type TEXT NOT NULL,
|
|
755
|
+
slot_id TEXT,
|
|
756
|
+
output_name TEXT,
|
|
757
|
+
prompt_id TEXT,
|
|
758
|
+
filename TEXT,
|
|
759
|
+
url TEXT,
|
|
760
|
+
width INTEGER,
|
|
761
|
+
height INTEGER,
|
|
762
|
+
inputs_json TEXT NOT NULL,
|
|
763
|
+
artifact_json TEXT,
|
|
764
|
+
raw_result_json TEXT,
|
|
765
|
+
error TEXT,
|
|
766
|
+
created_at TEXT NOT NULL,
|
|
767
|
+
updated_at TEXT NOT NULL
|
|
768
|
+
)
|
|
769
|
+
"""
|
|
770
|
+
)
|
|
771
|
+
self.connection.execute(
|
|
772
|
+
"""
|
|
773
|
+
CREATE TABLE IF NOT EXISTS output_slots (
|
|
774
|
+
slot_id TEXT PRIMARY KEY,
|
|
775
|
+
run_id TEXT NOT NULL,
|
|
776
|
+
session_id TEXT NOT NULL,
|
|
777
|
+
scope_key TEXT NOT NULL,
|
|
778
|
+
workflow TEXT NOT NULL,
|
|
779
|
+
contract TEXT NOT NULL,
|
|
780
|
+
output_name TEXT NOT NULL,
|
|
781
|
+
output_type TEXT NOT NULL,
|
|
782
|
+
status TEXT NOT NULL,
|
|
783
|
+
prompt_id TEXT,
|
|
784
|
+
width INTEGER,
|
|
785
|
+
height INTEGER,
|
|
786
|
+
error TEXT,
|
|
787
|
+
raw_result_json TEXT,
|
|
788
|
+
created_at TEXT NOT NULL,
|
|
789
|
+
updated_at TEXT NOT NULL
|
|
790
|
+
)
|
|
791
|
+
"""
|
|
792
|
+
)
|
|
793
|
+
self._ensure_column("gallery_items", "slot_id", "TEXT")
|
|
794
|
+
self.connection.execute(
|
|
795
|
+
"""
|
|
796
|
+
CREATE INDEX IF NOT EXISTS idx_gallery_items_scope_created_item
|
|
797
|
+
ON gallery_items(scope_key, created_at DESC, item_id DESC)
|
|
798
|
+
"""
|
|
799
|
+
)
|
|
800
|
+
self.connection.execute(
|
|
801
|
+
"CREATE INDEX IF NOT EXISTS idx_gallery_items_scope_run ON gallery_items(scope_key, run_id)"
|
|
802
|
+
)
|
|
803
|
+
self.connection.execute(
|
|
804
|
+
"CREATE INDEX IF NOT EXISTS idx_runs_scope_status_created ON runs(scope_key, status, created_at DESC)"
|
|
805
|
+
)
|
|
806
|
+
self.connection.execute(
|
|
807
|
+
"CREATE INDEX IF NOT EXISTS idx_output_slots_scope_run ON output_slots(scope_key, run_id)"
|
|
808
|
+
)
|
|
809
|
+
|
|
810
|
+
def _ensure_column(self, table_name: str, column_name: str, declaration: str) -> None:
|
|
811
|
+
rows = self.connection.execute(f"PRAGMA table_info({table_name})").fetchall()
|
|
812
|
+
if any(row["name"] == column_name for row in rows):
|
|
813
|
+
return
|
|
814
|
+
self.connection.execute(f"ALTER TABLE {table_name} ADD COLUMN {column_name} {declaration}")
|
|
815
|
+
|
|
816
|
+
|
|
817
|
+
def _with_run_timestamps(run: ServeRunRecord) -> ServeRunRecord:
|
|
818
|
+
now = utc_now()
|
|
819
|
+
created_at = run.created_at or now
|
|
820
|
+
return ServeRunRecord(
|
|
821
|
+
run_id=run.run_id,
|
|
822
|
+
session_id=run.session_id,
|
|
823
|
+
scope_key=run.scope_key,
|
|
824
|
+
workflow=run.workflow,
|
|
825
|
+
contract=run.contract,
|
|
826
|
+
status=run.status,
|
|
827
|
+
inputs=run.inputs,
|
|
828
|
+
prompt_id=run.prompt_id,
|
|
829
|
+
raw_result=run.raw_result,
|
|
830
|
+
error=run.error,
|
|
831
|
+
created_at=created_at,
|
|
832
|
+
updated_at=run.updated_at or now,
|
|
833
|
+
)
|
|
834
|
+
|
|
835
|
+
|
|
836
|
+
def _with_output_slot_timestamps(slot: ServeRunOutputSlot) -> ServeRunOutputSlot:
|
|
837
|
+
now = utc_now()
|
|
838
|
+
created_at = slot.created_at or now
|
|
839
|
+
return ServeRunOutputSlot(
|
|
840
|
+
slot_id=slot.slot_id,
|
|
841
|
+
run_id=slot.run_id,
|
|
842
|
+
session_id=slot.session_id,
|
|
843
|
+
scope_key=slot.scope_key,
|
|
844
|
+
workflow=slot.workflow,
|
|
845
|
+
contract=slot.contract,
|
|
846
|
+
output_name=slot.output_name,
|
|
847
|
+
output_type=slot.output_type,
|
|
848
|
+
status=slot.status,
|
|
849
|
+
prompt_id=slot.prompt_id,
|
|
850
|
+
width=slot.width,
|
|
851
|
+
height=slot.height,
|
|
852
|
+
error=slot.error,
|
|
853
|
+
raw_result=slot.raw_result,
|
|
854
|
+
created_at=created_at,
|
|
855
|
+
updated_at=slot.updated_at or now,
|
|
856
|
+
)
|
|
857
|
+
|
|
858
|
+
|
|
859
|
+
def _with_gallery_timestamps(item: ServeGalleryItem) -> ServeGalleryItem:
|
|
860
|
+
now = utc_now()
|
|
861
|
+
created_at = item.created_at or now
|
|
862
|
+
return ServeGalleryItem(
|
|
863
|
+
item_id=item.item_id,
|
|
864
|
+
run_id=item.run_id,
|
|
865
|
+
session_id=item.session_id,
|
|
866
|
+
scope_key=item.scope_key,
|
|
867
|
+
workflow=item.workflow,
|
|
868
|
+
contract=item.contract,
|
|
869
|
+
status=item.status,
|
|
870
|
+
output_type=item.output_type,
|
|
871
|
+
inputs=item.inputs,
|
|
872
|
+
slot_id=item.slot_id,
|
|
873
|
+
output_name=item.output_name,
|
|
874
|
+
prompt_id=item.prompt_id,
|
|
875
|
+
filename=item.filename,
|
|
876
|
+
url=item.url,
|
|
877
|
+
width=item.width,
|
|
878
|
+
height=item.height,
|
|
879
|
+
artifact=item.artifact,
|
|
880
|
+
raw_result=item.raw_result,
|
|
881
|
+
error=item.error,
|
|
882
|
+
created_at=created_at,
|
|
883
|
+
updated_at=item.updated_at or now,
|
|
884
|
+
)
|
|
885
|
+
|
|
886
|
+
|
|
887
|
+
def _output_slot_from_row(row: sqlite3.Row) -> ServeRunOutputSlot:
|
|
888
|
+
return ServeRunOutputSlot(
|
|
889
|
+
slot_id=str(row["slot_id"]),
|
|
890
|
+
run_id=str(row["run_id"]),
|
|
891
|
+
session_id=str(row["session_id"]),
|
|
892
|
+
scope_key=str(row["scope_key"]),
|
|
893
|
+
workflow=str(row["workflow"]),
|
|
894
|
+
contract=str(row["contract"]),
|
|
895
|
+
output_name=str(row["output_name"]),
|
|
896
|
+
output_type=str(row["output_type"]),
|
|
897
|
+
status=str(row["status"]),
|
|
898
|
+
prompt_id=row["prompt_id"],
|
|
899
|
+
width=row["width"],
|
|
900
|
+
height=row["height"],
|
|
901
|
+
error=row["error"],
|
|
902
|
+
raw_result=_loads(row["raw_result_json"], None),
|
|
903
|
+
created_at=str(row["created_at"]),
|
|
904
|
+
updated_at=str(row["updated_at"]),
|
|
905
|
+
)
|
|
906
|
+
|
|
907
|
+
|
|
908
|
+
def _gallery_item_from_row(row: sqlite3.Row) -> ServeGalleryItem:
|
|
909
|
+
return ServeGalleryItem(
|
|
910
|
+
item_id=str(row["item_id"]),
|
|
911
|
+
run_id=str(row["run_id"]),
|
|
912
|
+
session_id=str(row["session_id"]),
|
|
913
|
+
scope_key=str(row["scope_key"]),
|
|
914
|
+
workflow=str(row["workflow"]),
|
|
915
|
+
contract=str(row["contract"]),
|
|
916
|
+
status=str(row["status"]),
|
|
917
|
+
output_type=str(row["output_type"]),
|
|
918
|
+
slot_id=_row_value(row, "slot_id"),
|
|
919
|
+
output_name=row["output_name"],
|
|
920
|
+
prompt_id=row["prompt_id"],
|
|
921
|
+
filename=row["filename"],
|
|
922
|
+
url=row["url"],
|
|
923
|
+
width=row["width"],
|
|
924
|
+
height=row["height"],
|
|
925
|
+
inputs=_loads(row["inputs_json"], {}),
|
|
926
|
+
artifact=_loads(row["artifact_json"], None),
|
|
927
|
+
raw_result=_loads(row["raw_result_json"], None),
|
|
928
|
+
error=row["error"],
|
|
929
|
+
created_at=str(row["created_at"]),
|
|
930
|
+
updated_at=str(row["updated_at"]),
|
|
931
|
+
)
|
|
932
|
+
|
|
933
|
+
|
|
934
|
+
def _gallery_page_from_items(items: list[ServeGalleryItem], *, limit: int | None) -> GalleryPage:
|
|
935
|
+
if limit is None:
|
|
936
|
+
return GalleryPage(
|
|
937
|
+
items=[item.to_public_dict() for item in items],
|
|
938
|
+
next_cursor=None,
|
|
939
|
+
has_more=False,
|
|
940
|
+
limit=None,
|
|
941
|
+
)
|
|
942
|
+
|
|
943
|
+
has_more = len(items) > limit
|
|
944
|
+
page_items = items[:limit]
|
|
945
|
+
next_cursor = _encode_gallery_cursor(page_items[-1]) if has_more and page_items else None
|
|
946
|
+
return GalleryPage(
|
|
947
|
+
items=[item.to_public_dict() for item in page_items],
|
|
948
|
+
next_cursor=next_cursor,
|
|
949
|
+
has_more=has_more,
|
|
950
|
+
limit=limit,
|
|
951
|
+
)
|
|
952
|
+
|
|
953
|
+
|
|
954
|
+
def _gallery_item_after_cursor(item: ServeGalleryItem, *, created_at: str, item_id: str) -> bool:
|
|
955
|
+
return item.created_at < created_at or (item.created_at == created_at and item.item_id < item_id)
|
|
956
|
+
|
|
957
|
+
|
|
958
|
+
def _encode_gallery_cursor(item: ServeGalleryItem) -> str:
|
|
959
|
+
payload = {
|
|
960
|
+
"v": GALLERY_CURSOR_VERSION,
|
|
961
|
+
"created_at": item.created_at,
|
|
962
|
+
"item_id": item.item_id,
|
|
963
|
+
}
|
|
964
|
+
raw = json.dumps(payload, sort_keys=True, separators=(",", ":")).encode("utf-8")
|
|
965
|
+
return base64.urlsafe_b64encode(raw).decode("ascii").rstrip("=")
|
|
966
|
+
|
|
967
|
+
|
|
968
|
+
def _decode_gallery_cursor(cursor: str | None) -> tuple[str, str] | None:
|
|
969
|
+
if cursor is None or cursor == "":
|
|
970
|
+
return None
|
|
971
|
+
try:
|
|
972
|
+
padding = "=" * (-len(cursor) % 4)
|
|
973
|
+
raw = base64.urlsafe_b64decode(f"{cursor}{padding}".encode("ascii"))
|
|
974
|
+
payload = json.loads(raw.decode("utf-8"))
|
|
975
|
+
except (binascii.Error, json.JSONDecodeError, UnicodeDecodeError, ValueError) as exc:
|
|
976
|
+
raise ValueError("Gallery cursor is invalid.") from exc
|
|
977
|
+
if not isinstance(payload, dict) or payload.get("v") != GALLERY_CURSOR_VERSION:
|
|
978
|
+
raise ValueError("Gallery cursor is invalid.")
|
|
979
|
+
created_at = payload.get("created_at")
|
|
980
|
+
item_id = payload.get("item_id")
|
|
981
|
+
if not isinstance(created_at, str) or not isinstance(item_id, str):
|
|
982
|
+
raise ValueError("Gallery cursor is invalid.")
|
|
983
|
+
return created_at, item_id
|
|
984
|
+
|
|
985
|
+
|
|
986
|
+
def _run_from_row(row: sqlite3.Row) -> ServeRunRecord:
|
|
987
|
+
return ServeRunRecord(
|
|
988
|
+
run_id=str(row["run_id"]),
|
|
989
|
+
session_id=str(row["session_id"]),
|
|
990
|
+
scope_key=str(row["scope_key"]),
|
|
991
|
+
workflow=str(row["workflow"]),
|
|
992
|
+
contract=str(row["contract"]),
|
|
993
|
+
status=str(row["status"]),
|
|
994
|
+
prompt_id=row["prompt_id"],
|
|
995
|
+
inputs=_loads(row["inputs_json"], {}),
|
|
996
|
+
raw_result=_loads(row["raw_result_json"], None),
|
|
997
|
+
error=row["error"],
|
|
998
|
+
created_at=str(row["created_at"]),
|
|
999
|
+
updated_at=str(row["updated_at"]),
|
|
1000
|
+
)
|
|
1001
|
+
|
|
1002
|
+
|
|
1003
|
+
def _json(value: Any) -> str | None:
|
|
1004
|
+
if value is None:
|
|
1005
|
+
return None
|
|
1006
|
+
return json.dumps(value, sort_keys=True, separators=(",", ":"))
|
|
1007
|
+
|
|
1008
|
+
|
|
1009
|
+
def _loads(value: str | None, fallback: Any) -> Any:
|
|
1010
|
+
if not value:
|
|
1011
|
+
return fallback
|
|
1012
|
+
try:
|
|
1013
|
+
return json.loads(value)
|
|
1014
|
+
except json.JSONDecodeError:
|
|
1015
|
+
return fallback
|
|
1016
|
+
|
|
1017
|
+
|
|
1018
|
+
def _row_value(row: sqlite3.Row, key: str) -> Any:
|
|
1019
|
+
return row[key] if key in row.keys() else None
|