nai-aclab 1.0.0

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.
package/web_app.py ADDED
@@ -0,0 +1,511 @@
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ import mimetypes
5
+ import random
6
+ import shutil
7
+ import sys
8
+ import threading
9
+ import time
10
+ import urllib.parse
11
+ from dataclasses import asdict, fields
12
+ from copy import deepcopy
13
+ from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
14
+ from pathlib import Path
15
+
16
+ from app import (
17
+ APP_DIR,
18
+ OUTPUT_DIR,
19
+ ApiSettings,
20
+ AppState,
21
+ BattingScene,
22
+ Category,
23
+ CharacterPreset,
24
+ GenerationSettings,
25
+ NovelAIClient,
26
+ PromptPreset,
27
+ float_range,
28
+ load_state,
29
+ now_id,
30
+ parse_artist_tags,
31
+ safe_path_name,
32
+ save_state,
33
+ weight_tag,
34
+ )
35
+
36
+
37
+ WEB_DIR = APP_DIR / "web"
38
+ JOBS: dict[str, dict] = {}
39
+ STATE_LOCK = threading.Lock()
40
+
41
+
42
+ def dataclass_from_dict(cls, data: dict):
43
+ valid = {f.name for f in fields(cls)}
44
+ return cls(**{k: v for k, v in (data or {}).items() if k in valid})
45
+
46
+
47
+ def state_from_dict(data: dict) -> AppState:
48
+ base_data = data.get("base_presets", [])
49
+ quality_override = data.get("quality_override_prompt", "")
50
+ if not quality_override.strip():
51
+ selected_base_name = (data.get("generation", {}) or {}).get("base_preset", "")
52
+ selected_base = next((item for item in base_data if item.get("name") == selected_base_name), None)
53
+ fallback_base = (
54
+ selected_base
55
+ if selected_base and selected_base.get("quality_override_prompt")
56
+ else next((item for item in base_data if item.get("quality_override_prompt")), None)
57
+ )
58
+ quality_override = (fallback_base or {}).get("quality_override_prompt", "")
59
+ return AppState(
60
+ categories=[dataclass_from_dict(Category, item) for item in data.get("categories", [])],
61
+ base_presets=[dataclass_from_dict(PromptPreset, item) for item in base_data],
62
+ character_presets=[dataclass_from_dict(CharacterPreset, item) for item in data.get("character_presets", [])],
63
+ quality_override_prompt=quality_override,
64
+ negative_prompt=data.get("negative_prompt", ""),
65
+ uc_prompt=data.get("uc_prompt", data.get("negative_prompt", "")),
66
+ api=dataclass_from_dict(ApiSettings, data.get("api", {})),
67
+ generation=dataclass_from_dict(GenerationSettings, data.get("generation", {})),
68
+ batting_scenes=[dataclass_from_dict(BattingScene, item) for item in data.get("batting_scenes", [])],
69
+ history=data.get("history", []),
70
+ )
71
+
72
+
73
+ def media_url(path: str) -> str:
74
+ try:
75
+ rel = Path(path).resolve().relative_to(OUTPUT_DIR.resolve())
76
+ except (ValueError, OSError):
77
+ return ""
78
+ return "/media/" + urllib.parse.quote(str(rel).replace("\\", "/"))
79
+
80
+
81
+ def state_payload(state: AppState) -> dict:
82
+ data = asdict(state)
83
+ for category in data.get("categories", []):
84
+ category["recognized_tags"] = parse_artist_tags(category.get("tags", []))
85
+ for history in data.get("history", []):
86
+ for item in history.get("items", []):
87
+ item["image_url"] = media_url(item.get("path", ""))
88
+ item["request_url"] = media_url(item.get("request_path", ""))
89
+ return data
90
+
91
+
92
+ def item_payload(item: dict) -> dict:
93
+ data = dict(item)
94
+ data["image_url"] = media_url(data.get("path", ""))
95
+ data["request_url"] = media_url(data.get("request_path", ""))
96
+ return data
97
+
98
+
99
+ def selected_base(state: AppState) -> PromptPreset | None:
100
+ name = state.generation.base_preset
101
+ return next((item for item in state.base_presets if item.name == name), None) or (
102
+ state.base_presets[0] if state.base_presets else None
103
+ )
104
+
105
+
106
+ def selected_character(state: AppState) -> CharacterPreset | None:
107
+ name = state.generation.character_preset
108
+ return next((item for item in state.character_presets if item.name == name), None) or (
109
+ state.character_presets[0] if state.character_presets else None
110
+ )
111
+
112
+
113
+ def random_artist_tags(state: AppState) -> list[dict]:
114
+ result = []
115
+ for category in state.categories:
116
+ tags = parse_artist_tags(category.tags)
117
+ if not tags:
118
+ continue
119
+ weights = float_range(category.min_weight, category.max_weight, category.granule)
120
+ pick_count = len(tags) if category.picks <= 0 else min(category.picks, len(tags))
121
+ for tag in random.sample(tags, pick_count):
122
+ weight = random.choice(weights) if weights else category.min_weight
123
+ result.append(
124
+ {
125
+ "category": category.name,
126
+ "tag": tag,
127
+ "weight": weight,
128
+ "prompt": weight_tag(tag, weight),
129
+ }
130
+ )
131
+ return result
132
+
133
+
134
+ def fixed_artist_tags(state: AppState) -> list[dict]:
135
+ result = []
136
+ for item in getattr(state.generation, "fixed_artists", []) or []:
137
+ tag = str(item.get("tag", "")).strip()
138
+ if not tag:
139
+ continue
140
+ try:
141
+ weight = float(item.get("weight", 1.0))
142
+ except (TypeError, ValueError):
143
+ weight = 1.0
144
+ result.append(
145
+ {
146
+ "category": item.get("category", "fixed"),
147
+ "tag": tag,
148
+ "weight": weight,
149
+ "prompt": item.get("prompt") or weight_tag(tag, weight),
150
+ }
151
+ )
152
+ return result
153
+
154
+
155
+ def build_prompt(state: AppState) -> tuple[str, str, str, list[dict]]:
156
+ base = selected_base(state)
157
+ character = selected_character(state)
158
+ artists = fixed_artist_tags(state) or random_artist_tags(state)
159
+ base_chunks = []
160
+ if base and base.prompt.strip():
161
+ base_chunks.append(base.prompt.strip())
162
+ if artists:
163
+ base_chunks.append(", ".join(item["prompt"] for item in artists))
164
+ quality_prompt = ""
165
+ if state.quality_override_prompt.strip():
166
+ quality_prompt = state.quality_override_prompt.strip()
167
+ elif base and base.quality_prompt.strip():
168
+ quality_prompt = base.quality_prompt.strip()
169
+ if quality_prompt:
170
+ base_chunks.append(quality_prompt)
171
+
172
+ prompt_parts = []
173
+ base_prompt = ", ".join(base_chunks)
174
+ if base_prompt:
175
+ prompt_parts.append(base_prompt)
176
+ if character:
177
+ prompt_parts.extend(prompt.strip() for prompt in character.prompts if prompt.strip())
178
+
179
+ negative = [state.negative_prompt.strip()]
180
+ if character:
181
+ negative.extend(item.strip() for item in character.negatives if item.strip())
182
+ negative_prompt = ", ".join(item for item in negative if item)
183
+ uc_prompt = state.uc_prompt.strip() or negative_prompt
184
+ return " | ".join(prompt_parts), negative_prompt, uc_prompt, artists
185
+
186
+
187
+ def save_incoming_state(data: dict) -> AppState:
188
+ state = state_from_dict(data)
189
+ with STATE_LOCK:
190
+ current = load_state()
191
+ state.history = current.history
192
+ save_state(state)
193
+ return state
194
+
195
+
196
+ def delete_history_entries(ids: list[str], delete_files: bool) -> AppState:
197
+ ids_set = set(ids)
198
+ with STATE_LOCK:
199
+ state = load_state()
200
+ removed = [item for item in state.history if item.get("id") in ids_set]
201
+ state.history = [item for item in state.history if item.get("id") not in ids_set]
202
+ save_state(state)
203
+ if delete_files:
204
+ output_root = OUTPUT_DIR.resolve()
205
+ for history in removed:
206
+ raw_dir = history.get("output_dir", "")
207
+ if not raw_dir:
208
+ continue
209
+ try:
210
+ target = Path(raw_dir).resolve()
211
+ except OSError:
212
+ continue
213
+ if target == output_root or output_root not in target.parents:
214
+ continue
215
+ if target.exists():
216
+ shutil.rmtree(target, ignore_errors=True)
217
+ return load_state()
218
+
219
+
220
+ def clear_history(delete_files: bool) -> AppState:
221
+ with STATE_LOCK:
222
+ state = load_state()
223
+ ids = [item.get("id", "") for item in state.history]
224
+ return delete_history_entries(ids, delete_files)
225
+
226
+
227
+ def run_generation(job_id: str, state: AppState) -> None:
228
+ base = selected_base(state)
229
+ character = selected_character(state)
230
+ count = max(1, int(state.generation.count or 1))
231
+ run_id = now_id()
232
+ out_dir = OUTPUT_DIR / f"{run_id}_{safe_path_name(base.name if base else 'base')}_{safe_path_name(character.name if character else 'character')}"
233
+ out_dir.mkdir(parents=True, exist_ok=True)
234
+ client = NovelAIClient(state.api)
235
+ items = []
236
+
237
+ JOBS[job_id].update({"status": "running", "progress": 0, "total": count, "output_dir": str(out_dir), "items": []})
238
+ cancelled = False
239
+ for idx in range(count):
240
+ if JOBS.get(job_id, {}).get("cancel_requested"):
241
+ cancelled = True
242
+ JOBS[job_id]["log"].append("중지 요청으로 남은 생성을 건너뜁니다.")
243
+ break
244
+ prompt, negative, uc_prompt, artists = build_prompt(state)
245
+ path = out_dir / f"image_{idx + 1:03}.png"
246
+ metadata = {
247
+ "path": str(path),
248
+ "request_path": str(path.with_name(path.stem + "_request.json")),
249
+ "prompt": prompt,
250
+ "negative_prompt": negative,
251
+ "uc_prompt": uc_prompt,
252
+ "artists": artists,
253
+ "created_at": time.strftime("%Y-%m-%d %H:%M:%S"),
254
+ }
255
+ try:
256
+ client.generate(prompt, negative, uc_prompt, path)
257
+ JOBS[job_id]["log"].append(f"{idx + 1}/{count} 완료: {path.name}")
258
+ except Exception as exc:
259
+ metadata["error"] = str(exc)
260
+ JOBS[job_id]["log"].append(f"{idx + 1}/{count} 실패: {exc}")
261
+ items.append(metadata)
262
+ JOBS[job_id].update({"progress": idx + 1, "items": [item_payload(item) for item in items]})
263
+
264
+ history = {
265
+ "id": out_dir.name,
266
+ "base_preset": base.name if base else "",
267
+ "character_preset": character.name if character else "",
268
+ "created_at": time.strftime("%Y-%m-%d %H:%M:%S"),
269
+ "output_dir": str(out_dir),
270
+ "items": items,
271
+ }
272
+ update = {
273
+ "status": "cancelled" if cancelled or JOBS.get(job_id, {}).get("cancel_requested") else "done",
274
+ "log": JOBS[job_id]["log"] + ["생성 작업이 중지되었습니다." if cancelled else "생성 작업이 끝났습니다."],
275
+ }
276
+ if items:
277
+ with STATE_LOCK:
278
+ latest = load_state()
279
+ latest.history.insert(0, history)
280
+ save_state(latest)
281
+ update["history"] = history
282
+ JOBS[job_id].update(update)
283
+
284
+
285
+ def scene_state(state: AppState, scene: BattingScene) -> AppState:
286
+ request_state = deepcopy(state)
287
+ request_state.generation.base_preset = scene.base_preset
288
+ request_state.generation.character_preset = scene.character_preset
289
+ request_state.generation.count = scene_count(scene)
290
+ return request_state
291
+
292
+
293
+ def scene_count(scene: BattingScene) -> int:
294
+ try:
295
+ return max(1, int(scene.count or 2))
296
+ except (TypeError, ValueError):
297
+ return 2
298
+
299
+
300
+ def run_batting_test(job_id: str, state: AppState) -> None:
301
+ scenes = [scene for scene in state.batting_scenes if scene.base_preset and scene.character_preset]
302
+ if not scenes:
303
+ JOBS[job_id].update({"status": "done", "progress": 0, "total": 0, "items": [], "log": ["타율 테스트에 사용할 씬이 없습니다."]})
304
+ return
305
+
306
+ total = sum(scene_count(scene) for scene in scenes)
307
+ run_id = now_id()
308
+ out_dir = OUTPUT_DIR / f"{run_id}_타율테스트_{len(scenes)}씬"
309
+ out_dir.mkdir(parents=True, exist_ok=True)
310
+ client = NovelAIClient(state.api)
311
+ items = []
312
+ progress = 0
313
+ JOBS[job_id].update(
314
+ {
315
+ "status": "running",
316
+ "progress": 0,
317
+ "total": total,
318
+ "output_dir": str(out_dir),
319
+ "items": [],
320
+ "log": [f"타율 테스트 시작: {len(scenes)}개 씬, 총 {total}장"],
321
+ }
322
+ )
323
+
324
+ cancelled = False
325
+ for scene_index, scene in enumerate(scenes, start=1):
326
+ if JOBS.get(job_id, {}).get("cancel_requested"):
327
+ cancelled = True
328
+ JOBS[job_id]["log"].append("중지 요청으로 남은 씬을 건너뜁니다.")
329
+ break
330
+ current = scene_state(state, scene)
331
+ base = selected_base(current)
332
+ character = selected_character(current)
333
+ count = scene_count(scene)
334
+ scene_name = scene.name.strip() or f"Scene {scene_index}"
335
+ scene_dir = out_dir / f"{scene_index:02}_{safe_path_name(scene_name)}"
336
+ scene_dir.mkdir(parents=True, exist_ok=True)
337
+ JOBS[job_id]["log"].append(
338
+ f"[{scene_index}/{len(scenes)}] {scene_name}: {base.name if base else ''} + {character.name if character else ''}"
339
+ )
340
+
341
+ for image_index in range(count):
342
+ if JOBS.get(job_id, {}).get("cancel_requested"):
343
+ cancelled = True
344
+ JOBS[job_id]["log"].append(f"{scene_name}: 중지 요청으로 남은 이미지를 건너뜁니다.")
345
+ break
346
+ prompt, negative, uc_prompt, artists = build_prompt(current)
347
+ path = scene_dir / f"image_{image_index + 1:03}.png"
348
+ metadata = {
349
+ "path": str(path),
350
+ "request_path": str(path.with_name(path.stem + "_request.json")),
351
+ "prompt": prompt,
352
+ "negative_prompt": negative,
353
+ "uc_prompt": uc_prompt,
354
+ "artists": artists,
355
+ "scene_name": scene_name,
356
+ "scene_index": scene_index,
357
+ "source_base_preset": base.name if base else "",
358
+ "source_character_preset": character.name if character else "",
359
+ "created_at": time.strftime("%Y-%m-%d %H:%M:%S"),
360
+ }
361
+ try:
362
+ client.generate(prompt, negative, uc_prompt, path)
363
+ JOBS[job_id]["log"].append(f"{scene_name} {image_index + 1}/{count} 완료: {path.name}")
364
+ except Exception as exc:
365
+ metadata["error"] = str(exc)
366
+ JOBS[job_id]["log"].append(f"{scene_name} {image_index + 1}/{count} 실패: {exc}")
367
+ items.append(metadata)
368
+ progress += 1
369
+ JOBS[job_id].update({"progress": progress, "items": [item_payload(item) for item in items]})
370
+ if cancelled:
371
+ break
372
+
373
+ history = {
374
+ "id": out_dir.name,
375
+ "type": "batting_test",
376
+ "base_preset": "타율 테스트",
377
+ "character_preset": f"{len(scenes)}개 씬",
378
+ "created_at": time.strftime("%Y-%m-%d %H:%M:%S"),
379
+ "output_dir": str(out_dir),
380
+ "scenes": [asdict(scene) for scene in scenes],
381
+ "items": items,
382
+ }
383
+ update = {
384
+ "status": "cancelled" if cancelled or JOBS.get(job_id, {}).get("cancel_requested") else "done",
385
+ "log": JOBS[job_id]["log"] + ["타율 테스트가 중지되었습니다." if cancelled else "타율 테스트가 끝났습니다."],
386
+ }
387
+ if items:
388
+ with STATE_LOCK:
389
+ latest = load_state()
390
+ latest.history.insert(0, history)
391
+ save_state(latest)
392
+ update["history"] = history
393
+ JOBS[job_id].update(update)
394
+
395
+
396
+ class Handler(BaseHTTPRequestHandler):
397
+ server_version = "NAIArtistWeb/0.1"
398
+
399
+ def log_message(self, format: str, *args) -> None:
400
+ return
401
+
402
+ def read_json(self) -> dict:
403
+ length = int(self.headers.get("Content-Length", "0"))
404
+ if length <= 0:
405
+ return {}
406
+ return json.loads(self.rfile.read(length).decode("utf-8"))
407
+
408
+ def send_json(self, data: dict, status: int = 200) -> None:
409
+ body = json.dumps(data, ensure_ascii=False).encode("utf-8")
410
+ self.send_response(status)
411
+ self.send_header("Content-Type", "application/json; charset=utf-8")
412
+ self.send_header("Content-Length", str(len(body)))
413
+ self.end_headers()
414
+ self.wfile.write(body)
415
+
416
+ def send_file(self, path: Path) -> None:
417
+ if not path.exists() or not path.is_file():
418
+ self.send_error(404)
419
+ return
420
+ body = path.read_bytes()
421
+ self.send_response(200)
422
+ self.send_header("Content-Type", mimetypes.guess_type(path.name)[0] or "application/octet-stream")
423
+ self.send_header("Content-Length", str(len(body)))
424
+ self.end_headers()
425
+ self.wfile.write(body)
426
+
427
+ def do_GET(self) -> None:
428
+ parsed = urllib.parse.urlparse(self.path)
429
+ route = parsed.path
430
+ if route == "/":
431
+ self.send_file(WEB_DIR / "index.html")
432
+ elif route.startswith("/web/"):
433
+ self.send_file((WEB_DIR / route.removeprefix("/web/")).resolve())
434
+ elif route == "/api/state":
435
+ self.send_json({"state": state_payload(load_state())})
436
+ elif route == "/api/preview":
437
+ state = load_state()
438
+ prompt, negative, uc_prompt, artists = build_prompt(state)
439
+ self.send_json({"prompt": prompt, "negative": negative, "uc": uc_prompt, "artists": artists})
440
+ elif route == "/api/job":
441
+ query = urllib.parse.parse_qs(parsed.query)
442
+ job_id = query.get("id", [""])[0]
443
+ self.send_json({"job": JOBS.get(job_id, {"status": "missing"})})
444
+ elif route.startswith("/media/"):
445
+ rel = urllib.parse.unquote(route.removeprefix("/media/"))
446
+ path = (OUTPUT_DIR / rel).resolve()
447
+ if not str(path).startswith(str(OUTPUT_DIR.resolve())):
448
+ self.send_error(403)
449
+ return
450
+ self.send_file(path)
451
+ else:
452
+ self.send_error(404)
453
+
454
+ def do_POST(self) -> None:
455
+ parsed = urllib.parse.urlparse(self.path)
456
+ route = parsed.path
457
+ data = self.read_json()
458
+ if route == "/api/state":
459
+ state = save_incoming_state(data.get("state", data))
460
+ self.send_json({"state": state_payload(state)})
461
+ elif route == "/api/preview":
462
+ state = save_incoming_state(data.get("state", data))
463
+ prompt, negative, uc_prompt, artists = build_prompt(state)
464
+ self.send_json({"prompt": prompt, "negative": negative, "uc": uc_prompt, "artists": artists})
465
+ elif route == "/api/generate":
466
+ state = save_incoming_state(data.get("state", data))
467
+ job_id = f"job_{int(time.time() * 1000)}"
468
+ JOBS[job_id] = {"id": job_id, "status": "queued", "progress": 0, "total": state.generation.count, "log": []}
469
+ thread = threading.Thread(target=run_generation, args=(job_id, state), daemon=True)
470
+ thread.start()
471
+ self.send_json({"job_id": job_id})
472
+ elif route == "/api/batting/generate":
473
+ state = save_incoming_state(data.get("state", data))
474
+ total = sum(scene_count(scene) for scene in state.batting_scenes)
475
+ job_id = f"job_{int(time.time() * 1000)}"
476
+ JOBS[job_id] = {"id": job_id, "status": "queued", "progress": 0, "total": total, "log": []}
477
+ thread = threading.Thread(target=run_batting_test, args=(job_id, state), daemon=True)
478
+ thread.start()
479
+ self.send_json({"job_id": job_id})
480
+ elif route == "/api/job/cancel":
481
+ job_id = str(data.get("job_id", ""))
482
+ job = JOBS.get(job_id)
483
+ if not job:
484
+ self.send_json({"job": {"status": "missing"}}, status=404)
485
+ return
486
+ if job.get("status") not in ("done", "cancelled", "missing"):
487
+ job["cancel_requested"] = True
488
+ job["status"] = "cancelling"
489
+ if "중지 요청을 받았습니다. 현재 처리 중인 이미지가 끝나면 멈춥니다." not in job.setdefault("log", []):
490
+ job["log"].append("중지 요청을 받았습니다. 현재 처리 중인 이미지가 끝나면 멈춥니다.")
491
+ self.send_json({"job": job})
492
+ elif route == "/api/history/delete":
493
+ state = delete_history_entries(data.get("ids", []), bool(data.get("delete_files", False)))
494
+ self.send_json({"state": state_payload(state)})
495
+ elif route == "/api/history/clear":
496
+ state = clear_history(bool(data.get("delete_files", False)))
497
+ self.send_json({"state": state_payload(state)})
498
+ else:
499
+ self.send_error(404)
500
+
501
+
502
+ def main() -> None:
503
+ port = int(sys.argv[1]) if len(sys.argv) > 1 else 8765
504
+ server = ThreadingHTTPServer(("127.0.0.1", port), Handler)
505
+ print("NAI Artist Combination Web UI")
506
+ print(f"http://127.0.0.1:{port}")
507
+ server.serve_forever()
508
+
509
+
510
+ if __name__ == "__main__":
511
+ main()