Files

338 lines
14 KiB
Python

#!/usr/bin/env python3
"""
OpenAI-compatible API wrapper for ComfyUI.
Supports: ideogram4, flux2
Endpoint: POST /v1/images/generations
Extended params: seed, steps, cfg, quality, negative_prompt, style
Timeout: 1800s (30 min)
"""
import os
import re
import json
import base64
import time
import asyncio
import httpx
from typing import Optional
from fastapi import FastAPI, Request, HTTPException, Header
from fastapi.responses import JSONResponse
from pydantic import BaseModel, Field
COMFY_HOST = os.environ.get("COMFYUI_HOST", "localhost")
COMFY_PORT = os.environ.get("COMFYUI_PORT", "8188")
COMFY_URL = f"http://{COMFY_HOST}:{COMFY_PORT}"
API_KEY = os.environ.get("API_KEY", "")
ALLOWED_SIZES = {
"256x256", "512x512", "512x768", "768x512", "768x1024", "1024x768",
"1024x1024", "1024x1536", "1536x1024", "832x1216", "1216x832",
"896x1152", "1152x896", "640x1536", "1536x640"
}
DEFAULT_SIZE = "1024x1024"
QUALITY_STEPS = {
"low": 8,
"medium": 20,
"high": 30,
"ultra": 50,
}
IDEOGRAM_ALLOWED = {"1024x1024", "1024x1536", "1536x1024", "832x1216", "1216x832", "896x1152", "1152x896"}
IDEOGRAM_DEFAULT = "1024x1024"
FLUX_ALLOWED = {"256x256", "512x512", "512x768", "768x512", "768x1024", "1024x768",
"1024x1024", "1024x1536", "1536x1024", "832x1216", "1216x832"}
# Only two models active on this system
ALLOWED_MODELS = {"ideogram4", "flux2"}
DEFAULT_MODEL = "ideogram4"
def _parse_size(size: str) -> tuple[int, int]:
m = re.match(r"(\d+)x(\d+)", size)
if not m:
raise ValueError(f"Invalid size: {size}")
return int(m.group(1)), int(m.group(2))
def _validate_model_size(model: str, size: str) -> str:
if model == "ideogram4":
if size not in IDEOGRAM_ALLOWED:
return IDEOGRAM_DEFAULT
return size
if model.startswith("flux"):
if size not in FLUX_ALLOWED:
return DEFAULT_SIZE
return size
return DEFAULT_SIZE
def _build_ideogram4_workflow(prompt: str, width: int, height: int, seed: int,
steps: int = 20, cfg: float = 1.0,
negative_prompt: str = "") -> dict:
"""
Ideogram 4 requires asymmetric CFG wiring:
- Main UNet goes through CFGOverride before DualModelGuider
- Separate unconditional UNet loaded via second UNETLoader
- ConditioningZeroOut for the negative (not a text prompt)
- DualModelGuider takes model_negative (unconditional UNet), not model_1
"""
return {
"1": {"inputs": {"vae_name": "flux2-vae.safetensors"}, "class_type": "VAELoader"},
"2": {"inputs": {"unet_name": "ideogram4_fp8_scaled.safetensors", "weight_dtype": "default"},
"class_type": "UNETLoader"},
"3": {"inputs": {"clip_name": "qwen3vl_8b_fp8_scaled.safetensors", "type": "ideogram4", "device": "default"},
"class_type": "CLIPLoader"},
"4": {"inputs": {"text": prompt, "clip": ["3", 0]}, "class_type": "CLIPTextEncode"},
"5": {"inputs": {"width": width, "height": height, "batch_size": 1}, "class_type": "EmptyFlux2LatentImage"},
"6": {"inputs": {"noise_seed": seed, "batch_count": 1}, "class_type": "RandomNoise"},
"7": {"inputs": {"sampler_name": "euler"}, "class_type": "KSamplerSelect"},
"8": {"inputs": {"steps": steps, "width": width, "height": height, "mu": 0.5, "std": 1.75},
"class_type": "Ideogram4Scheduler"},
# Asymmetric CFG: second unconditional UNet
"15": {"inputs": {"unet_name": "ideogram4_unconditional_fp8_scaled.safetensors", "weight_dtype": "default"},
"class_type": "UNETLoader"},
# Zero out conditioning for negative path
"16": {"inputs": {"conditioning": ["4", 0]}, "class_type": "ConditioningZeroOut"},
# CFG override on main model (70%-100% of steps)
"17": {"inputs": {"cfg": 3.0, "start_percent": 0.7, "end_percent": 1.0, "model": ["2", 0]},
"class_type": "CFGOverride"},
# DualModelGuider: main model (CFG overridden), unconditional model, positive, zeroed negative
"9": {"inputs": {"cfg": 7.0, "model": ["17", 0], "positive": ["4", 0],
"model_negative": ["15", 0], "negative": ["16", 0]},
"class_type": "DualModelGuider"},
"10": {"inputs": {"noise": ["6", 0], "guider": ["9", 0], "sampler": ["7", 0],
"sigmas": ["8", 0], "latent_image": ["5", 0]}, "class_type": "SamplerCustomAdvanced"},
"11": {"inputs": {"samples": ["10", 0], "vae": ["1", 0]}, "class_type": "VAEDecode"},
"12": {"inputs": {"filename_prefix": "ideogram4_api", "images": ["11", 0]}, "class_type": "SaveImage"},
}
def _build_flux2_workflow(prompt: str, width: int, height: int, seed: int,
steps: int = 20, cfg: float = 3.5,
negative_prompt: str = "") -> dict:
"""FLUX.2-klein-9B workflow using UnetLoaderGGUF + CLIPLoaderGGUF (Qwen3VL)."""
return {
"1": {"inputs": {"vae_name": "flux2-vae.safetensors"}, "class_type": "VAELoader"},
"2": {"inputs": {"unet_name": "flux-2-klein-9b-Q4_K_S.gguf"},
"class_type": "UnetLoaderGGUF"},
"3": {"inputs": {"model": ["2", 0], "max_shift": 1.15, "base_shift": 0.5, "width": width, "height": height},
"class_type": "ModelSamplingFlux"},
# FLUX.2 uses Qwen3VL via CLIPLoaderGGUF, NOT DualCLIPLoader (CLIP-L + T5XXL)
"4": {"inputs": {"clip_name": "Qwen3VL-8B-Instruct-Q4_K_M.gguf", "type": "flux2"},
"class_type": "CLIPLoaderGGUF"},
"5": {"inputs": {"text": prompt, "clip": ["4", 0]}, "class_type": "CLIPTextEncode"},
"6": {"inputs": {"conditioning": ["5", 0], "guidance": cfg}, "class_type": "FluxGuidance"},
"7": {"inputs": {"width": width, "height": height, "batch_size": 1}, "class_type": "EmptyFlux2LatentImage"},
"8": {"inputs": {"noise_seed": seed}, "class_type": "RandomNoise"},
"9": {"inputs": {"sampler_name": "euler"}, "class_type": "KSamplerSelect"},
"10": {"inputs": {"model": ["3", 0], "steps": steps, "denoise": 1.0, "scheduler": "simple",
"sampler": ["9", 0]}, "class_type": "BasicScheduler"},
"11": {"inputs": {"model": ["3", 0], "conditioning": ["6", 0]}, "class_type": "BasicGuider"},
"12": {"inputs": {"noise": ["8", 0], "guider": ["11", 0], "sampler": ["9", 0],
"sigmas": ["10", 0], "latent_image": ["7", 0]}, "class_type": "SamplerCustomAdvanced"},
"13": {"inputs": {"samples": ["12", 0], "vae": ["1", 0]}, "class_type": "VAEDecode"},
"14": {"inputs": {"filename_prefix": "flux2_api", "images": ["13", 0]}, "class_type": "SaveImage"},
}
def _build_workflow(model: str, prompt: str, width: int, height: int,
seed: int, steps: int, cfg: float,
negative_prompt: str = "", style: str = "") -> dict:
if model == "ideogram4":
return _build_ideogram4_workflow(prompt, width, height, seed, steps, cfg, negative_prompt)
elif model == "flux2":
return _build_flux2_workflow(prompt, width, height, seed, steps, cfg, negative_prompt)
else:
raise ValueError(f"Unknown model: {model}")
async def _comfy_post(path: str, data: dict, timeout: float = 2400) -> dict:
async with httpx.AsyncClient() as client:
r = await client.post(f"{COMFY_URL}{path}", json=data, timeout=timeout)
if r.status_code >= 400:
try:
detail = r.json()
except Exception:
detail = r.text
raise HTTPException(status_code=r.status_code, detail=detail)
return r.json()
async def _comfy_get(path: str, timeout: float = 30) -> dict:
async with httpx.AsyncClient() as client:
r = await client.get(f"{COMFY_URL}{path}", timeout=timeout)
r.raise_for_status()
return r.json()
async def _generate_image(model: str, prompt: str, width: int, height: int, seed: int,
steps: int, cfg: float, negative_prompt: str = "", style: str = "") -> bytes:
workflow = _build_workflow(model, prompt, width, height, seed, steps, cfg, negative_prompt, style)
resp = await _comfy_post("/prompt", {"prompt": workflow})
prompt_id = resp.get("prompt_id")
if not prompt_id:
raise HTTPException(status_code=500, detail="No prompt_id from ComfyUI")
# Poll for completion - up to 30 min (180 iterations x 10s)
for i in range(180):
await asyncio.sleep(10)
try:
hist = await _comfy_get(f"/history/{prompt_id}", timeout=10)
except Exception:
continue
if prompt_id not in hist:
continue
d = hist[prompt_id]
status = d.get("status", {})
status_str = status.get("status_str", "?")
if status_str == "error":
errors = []
for node_id, msgs in status.get("messages", {}).items():
for level, msg, detail, _ in msgs:
if level == "error":
errors.append(f"[{node_id}] {msg}: {detail}")
detail = d.get("outputs", {}).get("error", "") or "; ".join(errors) or "ComfyUI execution error"
raise HTTPException(status_code=500, detail=detail)
if status_str == "success":
outputs = d.get("outputs", {})
for node_id, node_out in outputs.items():
for img in node_out.get("images", []):
fname = img.get("filename")
subfolder = img.get("subfolder", "")
ftype = img.get("type", "output")
url = f"{COMFY_URL}/view?filename={fname}&subfolder={subfolder}&type={ftype}"
async with httpx.AsyncClient() as client:
img_resp = await client.get(url, timeout=60)
img_resp.raise_for_status()
return img_resp.content
raise HTTPException(status_code=500, detail="No image in ComfyUI outputs")
raise HTTPException(status_code=504, detail="ComfyUI generation timeout")
app = FastAPI(title="ComfyUI Image Gen OpenAI API", version="2.1.0")
class ImageGenerationRequest(BaseModel):
model: str = Field(default="ideogram4", description="Model: ideogram4, flux2")
prompt: str = Field(..., description="Text prompt")
n: int = Field(default=1, ge=1, le=4, description="Number of images")
size: str = Field(default="1024x1024")
response_format: str = Field(default="url", pattern="^(url|b64_json)$")
seed: Optional[int] = Field(default=None, description="Random seed (int). Random if omitted.")
steps: Optional[int] = Field(default=None, ge=1, le=100, description="Denoising steps. Overrides quality.")
cfg: Optional[float] = Field(default=None, ge=0.0, le=100.0, description="CFG / guidance scale")
quality: Optional[str] = Field(default="medium", description="low|medium|high|ultra -> maps to steps")
negative_prompt: Optional[str] = Field(default="", description="Negative prompt")
style: Optional[str] = Field(default="", description="Style modifier (model-dependent)")
user: Optional[str] = Field(default=None, description="OpenAI user field (logged, not used)")
class ImageData(BaseModel):
url: Optional[str] = None
b64_json: Optional[str] = None
revised_prompt: Optional[str] = None
class ImageGenerationResponse(BaseModel):
created: int
data: list[ImageData]
@app.middleware("http")
async def auth_middleware(request: Request, call_next):
if API_KEY and request.url.path.startswith("/v1"):
auth = request.headers.get("authorization", "")
if not auth.startswith("Bearer ") or auth[7:] != API_KEY:
return JSONResponse(status_code=401, content={"error": "Unauthorized"})
return await call_next(request)
@app.post("/v1/images/generations", response_model=ImageGenerationResponse)
async def create_image_generation(req: ImageGenerationRequest):
if req.model not in ALLOWED_MODELS:
raise HTTPException(
status_code=400,
detail=f"Invalid model: {req.model}. Allowed: {ALLOWED_MODELS}"
)
validated_size = _validate_model_size(req.model, req.size)
if validated_size not in ALLOWED_SIZES:
raise HTTPException(status_code=400, detail=f"Invalid size: {validated_size}. Allowed: {ALLOWED_SIZES}")
width, height = _parse_size(validated_size)
created = int(time.time())
data = []
if req.steps is not None:
steps = req.steps
elif req.quality in QUALITY_STEPS:
steps = QUALITY_STEPS[req.quality]
else:
steps = QUALITY_STEPS["medium"]
# Ideogram4: enforce minimum ~12 steps to avoid blank/gray output
if req.model == "ideogram4" and steps < 12:
steps = 12
if req.model == "ideogram4":
cfg = req.cfg if req.cfg is not None else 1.0
elif req.model.startswith("flux"):
cfg = req.cfg if req.cfg is not None else 3.5
else:
cfg = req.cfg if req.cfg is not None else 7.5
for i in range(req.n):
seed = req.seed + i if req.seed is not None else int(time.time() * 1000) + i
image_bytes = await _generate_image(
req.model, req.prompt, width, height, seed,
steps, cfg, req.negative_prompt or "", req.style or ""
)
if req.response_format == "b64_json":
b64 = base64.b64encode(image_bytes).decode("utf-8")
data.append(ImageData(b64_json=b64, revised_prompt=req.prompt))
else:
b64 = base64.b64encode(image_bytes).decode("utf-8")
data_url = f"data:image/png;base64,{b64}"
data.append(ImageData(url=data_url, revised_prompt=req.prompt))
return ImageGenerationResponse(created=created, data=data)
@app.get("/v1/models")
async def list_models():
"""List available image models - dynamically from ALLOWED_MODELS."""
return {
"object": "list",
"data": [
{"id": m, "object": "model", "owned_by": "comfyui", "permission": []}
for m in sorted(ALLOWED_MODELS)
]
}
@app.get("/health")
async def health():
try:
await _comfy_get("/system_stats", timeout=5)
comfy_status = "ok"
except Exception as e:
comfy_status = f"error: {e}"
return {
"status": "ok",
"comfyui": COMFY_URL,
"comfyui_status": comfy_status,
"models": list(ALLOWED_MODELS)
}
if __name__ == "__main__":
import uvicorn
port = int(os.environ.get("PORT", "8000"))
uvicorn.run(app, host="0.0.0.0", port=port)