134 lines
5.7 KiB
Python
134 lines
5.7 KiB
Python
"""SwarmUI diffusion backend — talk to a running SwarmUI server over HTTP.
|
|
|
|
SwarmUI (a REST wrapper over ComfyUI) exposes ``/API/GetNewSession`` to obtain a session
|
|
id, then ``/API/GenerateText2Image`` to run a generation. For inpaint we send the frame
|
|
and the mask as base64 PNG plus the prompt/params, and read the produced image back.
|
|
|
|
Implementation notes:
|
|
- Uses only stdlib ``urllib`` — **no new dependency**; the diffusion model runs in
|
|
SwarmUI's own process (so our app never imports torch on this path).
|
|
- Exact API field names drift between SwarmUI versions, so the request body is built in
|
|
one place (:meth:`_build_payload`) for easy tuning; errors surface the URL + a hint.
|
|
- The response may carry image data inline (``data:`` URI) or as a server-relative path
|
|
— :meth:`_fetch_image_bytes` handles both.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import base64
|
|
import json
|
|
import urllib.error
|
|
import urllib.request
|
|
|
|
import cv2
|
|
import numpy as np
|
|
|
|
from .base import Cancelled
|
|
from .diffusion import DiffusionBackend, InpaintParams
|
|
|
|
|
|
class SwarmUIBackend(DiffusionBackend):
|
|
def __init__(self, url: str | None, timeout: float = 600.0) -> None:
|
|
self._url = (url or "http://localhost:7801").rstrip("/")
|
|
self._timeout = timeout
|
|
self._session: str | None = None
|
|
|
|
@property
|
|
def name(self) -> str:
|
|
return "SwarmUI"
|
|
|
|
# ------------------------------------------------------------------ HTTP
|
|
def _post(self, route: str, payload: dict) -> dict:
|
|
data = json.dumps(payload).encode("utf-8")
|
|
req = urllib.request.Request(
|
|
self._url + route, data=data, headers={"Content-Type": "application/json"}
|
|
)
|
|
try:
|
|
with urllib.request.urlopen(req, timeout=self._timeout) as resp:
|
|
return json.loads(resp.read().decode("utf-8"))
|
|
except urllib.error.URLError as e:
|
|
raise RuntimeError(
|
|
f"Не удалось связаться со SwarmUI ({self._url}{route}): {e}.\n"
|
|
"Проверьте, что сервер SwarmUI запущен и адрес верный "
|
|
"(Файл → Движок восстановления…)."
|
|
) from e
|
|
|
|
def ping(self) -> str:
|
|
"""Open a fresh session to verify the server is reachable; return the session id.
|
|
|
|
Used by the settings dialog's "Проверить соединение" — forces a new
|
|
``GetNewSession`` (ignores any cached id) so repeated checks really re-test, and
|
|
raises a clear RuntimeError (URL + hint) if the server is down/unreachable.
|
|
"""
|
|
self._session = None
|
|
return self._session_id()
|
|
|
|
def _session_id(self) -> str:
|
|
if self._session is None:
|
|
r = self._post("/API/GetNewSession", {})
|
|
self._session = r.get("session_id") or r.get("sessionId")
|
|
if not self._session:
|
|
raise RuntimeError(f"SwarmUI не вернул session_id: {r}")
|
|
return self._session
|
|
|
|
@staticmethod
|
|
def _b64_png(img: np.ndarray) -> str:
|
|
ok, buf = cv2.imencode(".png", img)
|
|
if not ok:
|
|
raise RuntimeError("Не удалось закодировать изображение в PNG для SwarmUI")
|
|
return base64.b64encode(buf.tobytes()).decode("ascii")
|
|
|
|
def _build_payload(
|
|
self, session: str, image_b64: str, mask_b64: str, params: InpaintParams, h: int, w: int
|
|
) -> dict:
|
|
"""Map our params onto SwarmUI's GenerateText2Image body (centralised for tuning)."""
|
|
payload = {
|
|
"session_id": session,
|
|
"images": 1,
|
|
"prompt": params.prompt,
|
|
"negativeprompt": params.negative,
|
|
"width": w,
|
|
"height": h,
|
|
"steps": int(params.steps),
|
|
"cfgscale": float(params.cfg),
|
|
"seed": int(params.seed),
|
|
"initimage": image_b64,
|
|
"maskimage": mask_b64, # white = regenerate
|
|
"initimagecreativity": float(params.denoise), # 0..1 inpaint denoise
|
|
"maskblur": int(params.mask_blur),
|
|
}
|
|
if params.model:
|
|
payload["model"] = params.model
|
|
return payload
|
|
|
|
# --------------------------------------------------------------- backend
|
|
def inpaint(self, image_bgr, mask, params, should_cancel=None):
|
|
if should_cancel is not None and should_cancel():
|
|
raise Cancelled("Восстановление отменено")
|
|
session = self._session_id()
|
|
h, w = image_bgr.shape[:2]
|
|
payload = self._build_payload(
|
|
session, self._b64_png(image_bgr), self._b64_png(mask), params, h, w
|
|
)
|
|
if should_cancel is not None and should_cancel():
|
|
raise Cancelled("Восстановление отменено")
|
|
resp = self._post("/API/GenerateText2Image", payload)
|
|
return self._decode_result(resp)
|
|
|
|
def _decode_result(self, resp: dict) -> np.ndarray:
|
|
images = resp.get("images") or []
|
|
if not images:
|
|
raise RuntimeError(f"SwarmUI не вернул изображений (ответ: {resp})")
|
|
raw = self._fetch_image_bytes(images[0])
|
|
arr = cv2.imdecode(np.frombuffer(raw, np.uint8), cv2.IMREAD_COLOR)
|
|
if arr is None:
|
|
raise RuntimeError("Не удалось декодировать результат SwarmUI")
|
|
return arr
|
|
|
|
def _fetch_image_bytes(self, ref: str) -> bytes:
|
|
if ref.startswith("data:"): # inline base64 data URI
|
|
return base64.b64decode(ref.split(",", 1)[1])
|
|
url = ref if ref.startswith("http") else f"{self._url}/{ref.lstrip('/')}"
|
|
with urllib.request.urlopen(url, timeout=self._timeout) as r:
|
|
return r.read()
|