spriteforge/tests/test_comfyui_provider.py

59 lines
3.4 KiB
Python

from __future__ import annotations
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
import io, json
from pathlib import Path
from threading import Thread
from PIL import Image
from spriteforge.studio.models import GenerationRecipe, GenerationRequest
from spriteforge.studio.providers import ComfyUIProvider
class FakeComfy(BaseHTTPRequestHandler):
prompt = None
def log_message(self,*args): pass
def do_GET(self):
if self.path=="/system_stats": return self.reply({"system":{"os":"fake"}})
if self.path=="/history/job1": return self.reply({"job1":{"outputs":{"9":{"images":[{"filename":"out.png","subfolder":"","type":"output"}]}}}})
if self.path.startswith("/view?"):
image=Image.new("RGBA",(12,10),(1,2,3,255)); data=io.BytesIO();image.save(data,"PNG")
blob=data.getvalue();self.send_response(200);self.send_header("Content-Length",str(len(blob)));self.end_headers();self.wfile.write(blob);return
self.send_error(404)
def do_POST(self):
length=int(self.headers["Content-Length"]);body=self.rfile.read(length)
if self.path=="/upload/image": return self.reply({"name":"uploaded.png"})
if self.path=="/prompt":
FakeComfy.prompt=json.loads(body)["prompt"];return self.reply({"prompt_id":"job1"})
self.send_error(404)
def reply(self,value):
body=json.dumps(value).encode();self.send_response(200);self.send_header("Content-Type","application/json");self.send_header("Content-Length",str(len(body)));self.end_headers();self.wfile.write(body)
def test_comfyui_workflow_substitution_upload_and_download(tmp_path:Path):
server=ThreadingHTTPServer(("127.0.0.1",0),FakeComfy);thread=Thread(target=server.serve_forever,daemon=True);thread.start()
workflow=tmp_path/"workflow.json";workflow.write_text(json.dumps({"1":{"class_type":"TestNode","inputs":{"text":"{{PROMPT}}","negative":"{{NEGATIVE_PROMPT}}","seed":"{{SEED}}","image":"{{REFERENCE_0}}","width":"{{WIDTH}}","height":"{{HEIGHT}}"}}}))
provider=ComfyUIProvider(f"http://127.0.0.1:{server.server_port}",workflow,poll_interval=.001,timeout=2)
request=GenerationRequest(project_id="p",asset_id="a",shot_id="s",count=1,recipe=GenerationRecipe(provider="comfyui",model="sd15",seed=42,prompt="horror",width=256,height=256,reference_hashes=["a"*64],controls={"mask":"b"*64}))
try:
assert provider.check()["system"]["os"]=="fake"
output=provider.generate(request,lambda _: b"fake png")
assert output[0].startswith(b"\x89PNG")
inputs=FakeComfy.prompt["1"]["inputs"]
assert inputs=={"text":"horror","negative":"","seed":42,"image":"uploaded.png","width":256,"height":256}
finally:server.shutdown();server.server_close()
def test_workflow_preflight_rejects_ui_format_and_missing_markers(tmp_path:Path):
from spriteforge.studio.providers import workflow_issues
ui=tmp_path/"ui.json";ui.write_text(json.dumps({"nodes":[],"links":[]}));assert "UI format" in workflow_issues(ui)[0]
incomplete=tmp_path/"incomplete.json";incomplete.write_text(json.dumps({"1":{"class_type":"X","inputs":{"text":"{{PROMPT}}"}}}))
assert "missing required" in workflow_issues(incomplete)[0]
def test_bundled_baseline_workflow_passes_preflight():
from spriteforge.studio.providers import workflow_issues
path=Path(__file__).parents[1]/"examples"/"comfyui"/"txt2img_api.json"
assert workflow_issues(path)==[]