Files
HVideoTool/hvideotool/core/restore/swarmui.py
T

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()