271 lines
11 KiB
Python
271 lines
11 KiB
Python
#!/usr/bin/env python3
|
|
"""Router-compatible Qwen-Image-2.1 worker backed by pinned ComfyUI."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import os
|
|
import random
|
|
import shutil
|
|
import signal
|
|
import subprocess
|
|
import threading
|
|
import time
|
|
import urllib.request
|
|
import uuid
|
|
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
|
from pathlib import Path
|
|
|
|
|
|
HOST = os.environ.get("WORKER_HOST", "0.0.0.0")
|
|
PORT = int(os.environ.get("WORKER_PORT", "8086"))
|
|
TOKEN = os.environ.get("WORKER_TOKEN", "").strip()
|
|
COMFY = "http://127.0.0.1:8188"
|
|
COMFY_DIR = Path("/opt/ComfyUI")
|
|
INPUT_DIR = COMFY_DIR / "input"
|
|
OUTPUT_DIR = Path(os.environ.get("IMAGE_DIR", "/data/images")).resolve()
|
|
MODEL = "Qwen-Image-2.1-int8"
|
|
GENERATION_LOCK = threading.Lock()
|
|
READY = threading.Event()
|
|
ACTIVE = False
|
|
COMFY_PROCESS: subprocess.Popen | None = None
|
|
|
|
if len(TOKEN) < 32:
|
|
raise RuntimeError("WORKER_TOKEN is missing or too short")
|
|
|
|
|
|
def request_json(path: str, payload: dict | None = None,
|
|
timeout: float = 30) -> dict:
|
|
body = None if payload is None else json.dumps(payload).encode()
|
|
headers = {"Content-Type": "application/json"} if body is not None else {}
|
|
method = "POST" if body is not None else "GET"
|
|
request = urllib.request.Request(COMFY + path, data=body,
|
|
headers=headers, method=method)
|
|
with urllib.request.urlopen(request, timeout=timeout) as response:
|
|
return json.load(response)
|
|
|
|
|
|
def wait_for_comfy() -> None:
|
|
for _ in range(300):
|
|
if COMFY_PROCESS is not None and COMFY_PROCESS.poll() is not None:
|
|
return
|
|
try:
|
|
request_json("/system_stats", timeout=2)
|
|
READY.set()
|
|
return
|
|
except Exception:
|
|
time.sleep(1)
|
|
|
|
|
|
def start_comfy() -> None:
|
|
global COMFY_PROCESS
|
|
INPUT_DIR.mkdir(parents=True, exist_ok=True)
|
|
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
|
|
command = [
|
|
"python", "main.py", "--listen", "127.0.0.1", "--port", "8188",
|
|
"--lowvram", "--preview-method", "none",
|
|
"--input-directory", str(INPUT_DIR),
|
|
"--output-directory", str(OUTPUT_DIR),
|
|
]
|
|
COMFY_PROCESS = subprocess.Popen(command, cwd=COMFY_DIR)
|
|
threading.Thread(target=wait_for_comfy, daemon=True).start()
|
|
|
|
|
|
def stop(*_: object) -> None:
|
|
if COMFY_PROCESS is not None and COMFY_PROCESS.poll() is None:
|
|
COMFY_PROCESS.terminate()
|
|
try:
|
|
COMFY_PROCESS.wait(timeout=15)
|
|
except subprocess.TimeoutExpired:
|
|
COMFY_PROCESS.kill()
|
|
raise SystemExit(0)
|
|
|
|
|
|
def validate_request(data: dict) -> tuple[str, str, int, int, int, float, int]:
|
|
prompt = data.get("prompt")
|
|
filename = data.get("filename")
|
|
if not isinstance(prompt, str) or not prompt.strip() or len(prompt) > 8000:
|
|
raise ValueError("invalid prompt")
|
|
if (not isinstance(filename, str) or Path(filename).name != filename
|
|
or not filename.endswith(".png")):
|
|
raise ValueError("invalid filename")
|
|
width = int(data.get("width", 1024))
|
|
height = int(data.get("height", 1024))
|
|
if width % 32 or height % 32 or not 512 <= width <= 2048 or not 512 <= height <= 2048:
|
|
raise ValueError("width and height must be multiples of 32 between 512 and 2048")
|
|
steps = int(data.get("steps", 25))
|
|
guidance = float(data.get("guidance", 1.0))
|
|
if steps != 25 or guidance != 1.0:
|
|
raise ValueError("Qwen-Image-2.1 requires steps=25 and guidance=1.0")
|
|
seed = data.get("seed")
|
|
seed = random.randrange(2**32) if seed is None else int(seed)
|
|
if not 0 <= seed <= 2**32 - 1:
|
|
raise ValueError("invalid seed")
|
|
return prompt.strip(), filename, width, height, steps, guidance, seed
|
|
|
|
|
|
def make_workflow(prompt: str, seed: int, steps: int, width: int, height: int,
|
|
references: list[str], prefix: str) -> dict:
|
|
encode_inputs: dict = {
|
|
"clip": ["2", 0], "prompt": prompt, "negative_prompt": "",
|
|
"resolution": max(width, height),
|
|
}
|
|
workflow: dict = {
|
|
"1": {"class_type": "UNETLoader", "inputs": {
|
|
"unet_name": "qwen_image_2.1_int8_convrot.safetensors",
|
|
"weight_dtype": "default"}},
|
|
"2": {"class_type": "CLIPLoader", "inputs": {
|
|
"clip_name": "qwen3vl_8b_int8_convrot.safetensors",
|
|
"type": "qwen_image", "device": "default"}},
|
|
"3": {"class_type": "VAELoader", "inputs": {
|
|
"vae_name": "qwen_image_2.1_vae_bf16.safetensors"}},
|
|
"4": {"class_type": "TextEncodeQwenImage21", "inputs": encode_inputs},
|
|
"6": {"class_type": "KSampler", "inputs": {
|
|
"model": ["1", 0], "positive": ["4", 0], "negative": ["4", 1],
|
|
"latent_image": ["4", 2] if references else ["5", 0],
|
|
"seed": seed, "steps": steps, "cfg": 1.0,
|
|
"sampler_name": "euler", "scheduler": "simple", "denoise": 1.0}},
|
|
"7": {"class_type": "VAEDecode", "inputs": {
|
|
"samples": ["6", 0], "vae": ["3", 0]}},
|
|
"8": {"class_type": "SaveImage", "inputs": {
|
|
"filename_prefix": prefix, "images": ["7", 0]}},
|
|
}
|
|
if references:
|
|
encode_inputs["vae"] = ["3", 0]
|
|
for index, name in enumerate(references, 1):
|
|
node = str(8 + index)
|
|
workflow[node] = {"class_type": "LoadImage", "inputs": {"image": name}}
|
|
# ComfyUI V3 dynamic inputs use the parent path as a prefix. The
|
|
# normalizer turns images.image_N back into the execute() dict.
|
|
encode_inputs[f"images.image_{index}"] = [node, 0]
|
|
else:
|
|
workflow["5"] = {"class_type": "EmptyLatentImage", "inputs": {
|
|
"width": width, "height": height, "batch_size": 1}}
|
|
return workflow
|
|
|
|
|
|
def wait_result(prompt_id: str, timeout: int = 1800) -> dict:
|
|
deadline = time.monotonic() + timeout
|
|
while time.monotonic() < deadline:
|
|
history = request_json(f"/history/{prompt_id}", timeout=10)
|
|
if prompt_id in history:
|
|
result = history[prompt_id]
|
|
status = result.get("status", {})
|
|
if status.get("status_str") == "error" or not status.get("completed", True):
|
|
raise RuntimeError("ComfyUI generation failed: " + json.dumps(status))
|
|
return result
|
|
time.sleep(1)
|
|
raise TimeoutError("Qwen image generation timed out")
|
|
|
|
|
|
def prepare_references(data: dict, job: str) -> list[str]:
|
|
source_files = data.get("source_files") or []
|
|
if not isinstance(source_files, list) or len(source_files) > 4:
|
|
raise ValueError("invalid source image list")
|
|
copied: list[str] = []
|
|
for index, name in enumerate(source_files, 1):
|
|
if not isinstance(name, str) or Path(name).name != name:
|
|
raise ValueError("invalid source image filename")
|
|
source = (OUTPUT_DIR / name).resolve()
|
|
if source.parent != OUTPUT_DIR or not source.is_file():
|
|
raise ValueError("source image not found")
|
|
with source.open("rb") as stream:
|
|
header = stream.read(16)
|
|
if header.startswith(b"\x89PNG\r\n\x1a\n"):
|
|
suffix = ".png"
|
|
elif header.startswith(b"\xff\xd8\xff"):
|
|
suffix = ".jpg"
|
|
elif header.startswith((b"RIFF",)) and header[8:12] == b"WEBP":
|
|
suffix = ".webp"
|
|
else:
|
|
raise ValueError("unsupported source image format")
|
|
target_name = f"{job}-ref-{index}{suffix}"
|
|
shutil.copyfile(source, INPUT_DIR / target_name)
|
|
copied.append(target_name)
|
|
return copied
|
|
|
|
|
|
def generate(data: dict) -> dict:
|
|
global ACTIVE
|
|
with GENERATION_LOCK:
|
|
ACTIVE = True
|
|
started = time.monotonic()
|
|
copied: list[str] = []
|
|
try:
|
|
prompt, filename, width, height, steps, _, seed = validate_request(data)
|
|
job = "router-" + uuid.uuid4().hex
|
|
copied = prepare_references(data, job)
|
|
queued = request_json("/prompt", {"prompt": make_workflow(
|
|
prompt, seed, steps, width, height, copied, job),
|
|
"client_id": "athena-image-router"})
|
|
result = wait_result(queued["prompt_id"])
|
|
images = [image for output in result.get("outputs", {}).values()
|
|
for image in output.get("images", [])]
|
|
if not images:
|
|
raise RuntimeError("generation completed without an image")
|
|
image = images[0]
|
|
generated = (OUTPUT_DIR / image.get("subfolder", "") /
|
|
image["filename"]).resolve()
|
|
if OUTPUT_DIR not in generated.parents or not generated.is_file():
|
|
raise RuntimeError("ComfyUI returned an invalid output path")
|
|
target = (OUTPUT_DIR / filename).resolve()
|
|
if target.parent != OUTPUT_DIR:
|
|
raise RuntimeError("invalid output path")
|
|
generated.replace(target)
|
|
return {"status": "ok", "filename": filename, "seed": seed,
|
|
"seconds": round(time.monotonic() - started, 3),
|
|
"model": MODEL}
|
|
finally:
|
|
for name in copied:
|
|
try:
|
|
(INPUT_DIR / name).unlink()
|
|
except FileNotFoundError:
|
|
pass
|
|
ACTIVE = False
|
|
|
|
|
|
class Handler(BaseHTTPRequestHandler):
|
|
def log_message(self, fmt: str, *args: object) -> None:
|
|
print(f"[qwen-image-2.1] {self.client_address[0]} {fmt % args}", flush=True)
|
|
|
|
def reply(self, status: int, payload: dict) -> None:
|
|
body = json.dumps(payload, separators=(",", ":")).encode()
|
|
self.send_response(status)
|
|
self.send_header("Content-Type", "application/json")
|
|
self.send_header("Content-Length", str(len(body)))
|
|
self.end_headers()
|
|
self.wfile.write(body)
|
|
|
|
def do_GET(self) -> None: # noqa: N802
|
|
if self.path != "/health":
|
|
self.reply(404, {"error": "not found"})
|
|
elif not READY.is_set():
|
|
self.reply(503, {"status": "starting", "model": MODEL})
|
|
else:
|
|
self.reply(200, {"status": "ok", "model_loaded": ACTIVE,
|
|
"model": MODEL})
|
|
|
|
def do_POST(self) -> None: # noqa: N802
|
|
if self.headers.get("Authorization", "") != f"Bearer {TOKEN}":
|
|
self.reply(401, {"error": "unauthorized"})
|
|
return
|
|
if self.path != "/generate":
|
|
self.reply(404, {"error": "not found"})
|
|
return
|
|
try:
|
|
length = int(self.headers.get("Content-Length", "0"))
|
|
if length < 2 or length > 32768:
|
|
raise ValueError("invalid request size")
|
|
self.reply(200, generate(json.loads(self.rfile.read(length))))
|
|
except Exception as exc:
|
|
print(f"[qwen-image-2.1] generation failed: {type(exc).__name__}: "
|
|
f"{str(exc)[:1000]}", flush=True)
|
|
self.reply(500, {"status": "error", "message": str(exc)})
|
|
|
|
|
|
if __name__ == "__main__":
|
|
signal.signal(signal.SIGTERM, stop)
|
|
signal.signal(signal.SIGINT, stop)
|
|
start_comfy()
|
|
ThreadingHTTPServer((HOST, PORT), Handler).serve_forever()
|