175 lines
6.8 KiB
Python
175 lines
6.8 KiB
Python
from __future__ import annotations
|
|
|
|
import argparse
|
|
import base64
|
|
import json
|
|
import struct
|
|
import sys
|
|
import time
|
|
import urllib.error
|
|
import urllib.parse
|
|
import urllib.request
|
|
import uuid
|
|
import zlib
|
|
|
|
|
|
def png_chunk(kind: bytes, data: bytes) -> bytes:
|
|
return struct.pack(">I", len(data)) + kind + data + struct.pack(">I", zlib.crc32(kind + data) & 0xFFFFFFFF)
|
|
|
|
|
|
def rgb_png(width: int, height: int, pixels: list[tuple[int, int, int]]) -> bytes:
|
|
rows = []
|
|
for y in range(height):
|
|
start = y * width
|
|
row = b"".join(bytes(pixel) for pixel in pixels[start : start + width])
|
|
rows.append(b"\x00" + row)
|
|
header = struct.pack(">IIBBBBB", width, height, 8, 2, 0, 0, 0)
|
|
return b"\x89PNG\r\n\x1a\n" + png_chunk(b"IHDR", header) + png_chunk(b"IDAT", zlib.compress(b"".join(rows), 6)) + png_chunk(b"IEND", b"")
|
|
|
|
|
|
def make_source_image(width: int, height: int) -> str:
|
|
pixels = []
|
|
for y in range(height):
|
|
for x in range(width):
|
|
if width // 3 < x < 2 * width // 3 and height // 3 < y < 2 * height // 3:
|
|
pixels.append((210, 60, 70))
|
|
else:
|
|
pixels.append((35 + (x * 80 // width), 95 + (y * 80 // height), 145))
|
|
return "data:image/png;base64," + base64.b64encode(rgb_png(width, height, pixels)).decode("ascii")
|
|
|
|
|
|
def make_mask_image(width: int, height: int) -> str:
|
|
pixels = []
|
|
for y in range(height):
|
|
for x in range(width):
|
|
masked = width // 3 < x < 2 * width // 3 and height // 3 < y < 2 * height // 3
|
|
pixels.append((255, 255, 255) if masked else (0, 0, 0))
|
|
return "data:image/png;base64," + base64.b64encode(rgb_png(width, height, pixels)).decode("ascii")
|
|
|
|
|
|
def post_json(url: str, payload: dict, timeout: int = 30) -> dict:
|
|
request = urllib.request.Request(
|
|
url,
|
|
data=json.dumps(payload, ensure_ascii=False).encode("utf-8"),
|
|
headers={"Content-Type": "application/json"},
|
|
method="POST",
|
|
)
|
|
with urllib.request.urlopen(request, timeout=timeout) as response:
|
|
data = json.loads(response.read().decode("utf-8"))
|
|
if not isinstance(data, dict):
|
|
raise ValueError(f"JSON object expected from {url}")
|
|
if data.get("error"):
|
|
raise ValueError(str(data["error"]))
|
|
return data
|
|
|
|
|
|
def get_json(url: str, timeout: int = 30) -> dict:
|
|
with urllib.request.urlopen(url, timeout=timeout) as response:
|
|
data = json.loads(response.read().decode("utf-8"))
|
|
if not isinstance(data, dict):
|
|
raise ValueError(f"JSON object expected from {url}")
|
|
if data.get("error"):
|
|
raise ValueError(str(data["error"]))
|
|
return data
|
|
|
|
|
|
def submit_job(args: argparse.Namespace, operation: str) -> dict:
|
|
payload: dict[str, object] = {
|
|
"operation": operation,
|
|
"base_url": args.image_base_url,
|
|
"model_id": args.model_id,
|
|
"model_mode": args.model_mode,
|
|
"model": args.model,
|
|
"prompt": args.prompt,
|
|
"negative_prompt": args.negative_prompt,
|
|
"width": args.width,
|
|
"height": args.height,
|
|
"steps": args.steps,
|
|
"guidance_scale": args.guidance_scale,
|
|
"seed": args.seed,
|
|
}
|
|
if operation == "edit":
|
|
payload.update(
|
|
{
|
|
"image_base64": make_source_image(args.width, args.height),
|
|
"mask_base64": make_mask_image(args.width, args.height),
|
|
"strength": args.strength,
|
|
}
|
|
)
|
|
return post_json(f"{args.chat_base_url.rstrip('/')}/api/image/submit", payload)
|
|
|
|
|
|
def poll_job(args: argparse.Namespace, job_id: str, base_url: str) -> dict:
|
|
deadline = time.monotonic() + args.timeout
|
|
last_status = ""
|
|
while time.monotonic() < deadline:
|
|
url = (
|
|
f"{args.chat_base_url.rstrip('/')}/api/image/job?"
|
|
f"job_id={urllib.parse.quote(job_id)}&base_url={urllib.parse.quote(base_url, safe='')}"
|
|
)
|
|
job = get_json(url)
|
|
status = str(job.get("status") or "")
|
|
if status != last_status:
|
|
print(f"{job_id}: {status}")
|
|
last_status = status
|
|
if status in {"completed", "cancelled", "error"}:
|
|
return job
|
|
time.sleep(args.poll_interval)
|
|
raise TimeoutError(f"image job {job_id} did not finish within {args.timeout}s")
|
|
|
|
|
|
def cancel_job(args: argparse.Namespace, job_id: str, base_url: str) -> dict:
|
|
return post_json(
|
|
f"{args.chat_base_url.rstrip('/')}/api/image/cancel",
|
|
{"job_id": job_id, "base_url": base_url},
|
|
)
|
|
|
|
|
|
def run_one(args: argparse.Namespace, operation: str) -> int:
|
|
submitted = submit_job(args, operation)
|
|
job_id = str(submitted.get("report_id") or submitted.get("id") or uuid.uuid4())
|
|
base_url = str(submitted.get("base_url") or args.image_base_url)
|
|
print(f"submitted {operation}: {job_id}")
|
|
if args.cancel:
|
|
time.sleep(args.cancel_after)
|
|
cancel_job(args, job_id, base_url)
|
|
job = poll_job(args, job_id, base_url)
|
|
status = job.get("status")
|
|
print(json.dumps({key: job.get(key) for key in ("id", "operation", "status", "latency_ms", "artifact", "error")}, ensure_ascii=False, indent=2))
|
|
if args.cancel:
|
|
return 0 if status in {"cancelled", "cancel_requested"} else 1
|
|
return 0 if status == "completed" and job.get("artifact") else 1
|
|
|
|
|
|
def main() -> int:
|
|
parser = argparse.ArgumentParser(description="Smoke test async image jobs through the model chat API.")
|
|
parser.add_argument("--chat-base-url", default="http://192.168.220.91:8765")
|
|
parser.add_argument("--image-base-url", default="http://docker-gpu.cin.su:8040")
|
|
parser.add_argument("--model", default="sdxl-image")
|
|
parser.add_argument("--model-id", default="sdxl-base-1_0")
|
|
parser.add_argument("--model-mode", default="sdxl")
|
|
parser.add_argument("--operation", choices=["generate", "edit"], default="generate")
|
|
parser.add_argument("--prompt", default="A small clean product photo of a blue ceramic mug on a white desk")
|
|
parser.add_argument("--negative-prompt", default="blurry, distorted, low quality")
|
|
parser.add_argument("--width", type=int, default=512)
|
|
parser.add_argument("--height", type=int, default=512)
|
|
parser.add_argument("--steps", type=int, default=8)
|
|
parser.add_argument("--guidance-scale", type=float, default=5.5)
|
|
parser.add_argument("--strength", type=float, default=0.85)
|
|
parser.add_argument("--seed", type=int, default=123)
|
|
parser.add_argument("--timeout", type=int, default=900)
|
|
parser.add_argument("--poll-interval", type=float, default=2.0)
|
|
parser.add_argument("--cancel", action="store_true")
|
|
parser.add_argument("--cancel-after", type=float, default=1.0)
|
|
args = parser.parse_args()
|
|
|
|
try:
|
|
return run_one(args, args.operation)
|
|
except (OSError, TimeoutError, ValueError, urllib.error.URLError) as exc:
|
|
print(f"Smoke image job failed: {exc}", file=sys.stderr)
|
|
return 1
|
|
|
|
|
|
if __name__ == "__main__":
|
|
raise SystemExit(main())
|