save code

This commit is contained in:
xsl
2026-07-17 01:54:04 +08:00
parent 2b1f528ddd
commit c1bb9614c7
8 changed files with 299 additions and 12 deletions
+67 -10
View File
@@ -4,11 +4,27 @@ import io
import json
import time
import random
import logging
import traceback
import requests
from flask import Flask, request, jsonify, send_file
import numpy as np
from PIL import Image, ImageFilter
# 让 PIL 支持 iPhone 的 HEIC/HEIF 照片(浏览器 accept="image/*" 会允许选中它们)。
try:
from pillow_heif import register_heif_opener
register_heif_opener()
_HEIF_OK = True
except Exception:
_HEIF_OK = False
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s [%(levelname)s] %(message)s",
)
log = logging.getLogger("hair")
app = Flask(__name__)
COMFYUI_URL = "http://127.0.0.1:8188"
@@ -18,6 +34,7 @@ _CORS_HEADERS = {
"Access-Control-Allow-Origin": "*",
"Access-Control-Allow-Methods": "POST, OPTIONS",
"Access-Control-Allow-Headers": "Content-Type",
"Access-Control-Expose-Headers": "X-Generate-Time",
}
@@ -43,7 +60,7 @@ def build_workflow(image_filename, prompt_text, seed=None):
# Loaders
"16": {"class_type": "UNETLoader", "inputs": {
"unet_name": "flux2.0/flux-2-klein-9b-fp8.safetensors",
"weight_dtype": "fp8_e4m3fn"}},
"weight_dtype": "fp8_e4m3fn_fast"}},
"3": {"class_type": "VAELoader", "inputs": {
"vae_name": "flux2-vae.safetensors"}},
"61": {"class_type": "CLIPLoader", "inputs": {
@@ -136,7 +153,7 @@ def build_workflow(image_filename, prompt_text, seed=None):
"1": {"class_type": "BasicScheduler", "inputs": {
"model": ["2", 0],
"scheduler": "simple",
"steps": 6,
"steps": 4,
"denoise": 1}},
"20": {"class_type": "BasicGuider", "inputs": {
"model": ["2", 0],
@@ -181,20 +198,42 @@ def index():
@app.route("/api/generate", methods=["POST"])
def generate():
t_start = time.time()
try:
if "image" not in request.files or "mask" not in request.files:
msg = f"缺少上传文件: files={list(request.files.keys())}"
log.warning(msg)
return jsonify({"error": msg}), 400
image_file = request.files["image"]
mask_file = request.files["mask"]
prompt_text = request.form.get("prompt", "填充遮罩区域的头发,皮肤加一点磨皮")
log.info(
"收到请求: image=%s mask=%s prompt=%r",
image_file.filename, mask_file.filename, prompt_text,
)
# Load original image as RGB
image = Image.open(image_file).convert("RGB")
try:
image = Image.open(image_file).convert("RGB")
except Exception as e:
log.error("无法解码人物图片 %s: %s", image_file.filename, e)
hint = "" if _HEIF_OK else "(当前不支持 HEIC"
return jsonify({
"error": f"无法识别人物图片格式{hint},请改用 JPG/PNG: {e}"
}), 400
# Load mask and extract mask data from ALL channels (R, G, B, A)
# This handles different mask formats:
# - Red mask (R=255 where drawn): user-provided PNG
# - White mask (R=G=B=255 where drawn): frontend canvas
# - Alpha mask (A=255 where drawn): transparent brush
mask_img = Image.open(mask_file).convert("RGBA")
try:
mask_img = Image.open(mask_file).convert("RGBA")
except Exception as e:
log.error("无法解码遮罩图片 %s: %s", mask_file.filename, e)
return jsonify({
"error": f"无法识别遮罩图片格式,请改用 JPG/PNG: {e}"
}), 400
mask_arr = np.array(mask_img)
# Use max of all channels: 255 where any color/alpha is drawn, 0 where empty
mask_data = np.max(mask_arr, axis=2) # (H, W) uint8
@@ -229,6 +268,7 @@ def generate():
)
upload_data = upload_resp.json()
if "name" not in upload_data:
log.error("ComfyUI 上传图片失败: %s", upload_data)
return jsonify({"error": f"Upload failed: {upload_data}"}), 500
image_filename = upload_data["name"]
@@ -241,12 +281,16 @@ def generate():
)
prompt_data = prompt_resp.json()
if "error" in prompt_data:
log.error(
"ComfyUI /prompt 校验失败: error=%s node_errors=%s",
prompt_data.get("error"), prompt_data.get("node_errors"),
)
return jsonify({"error": json.dumps(prompt_data["error"], ensure_ascii=False)}), 500
prompt_id = prompt_data["prompt_id"]
# Poll for completion (5 min timeout)
for _ in range(150):
time.sleep(2)
# Poll for completion (5 min timeout, 0.1s interval)
for _ in range(3000):
time.sleep(0.1)
history_resp = requests.get(
f"{COMFYUI_URL}/history/{prompt_id}", timeout=10
)
@@ -254,7 +298,14 @@ def generate():
if prompt_id in history_data:
status = history_data[prompt_id].get("status", {})
if status.get("status_str") == "error":
return jsonify({"error": "Workflow execution failed"}), 500
log.error(
"ComfyUI 工作流执行失败: %s",
json.dumps(status, ensure_ascii=False),
)
return jsonify({
"error": "Workflow execution failed",
"detail": status.get("messages", status),
}), 500
outputs = history_data[prompt_id].get("outputs", {})
if "17" in outputs: # SaveImage node
image_info = outputs["17"]["images"][0]
@@ -266,16 +317,22 @@ def generate():
params={"filename": filename, "subfolder": subfolder, "type": img_type},
timeout=30,
)
return send_file(
elapsed = time.time() - t_start
log.info("重绘完成,服务端耗时 %.2fs", elapsed)
resp = send_file(
io.BytesIO(view_resp.content), mimetype="image/png"
)
resp.headers["X-Generate-Time"] = f"{elapsed:.2f}"
return resp
return jsonify({"error": "Timeout: workflow did not complete in 5 minutes"}), 500
except requests.ConnectionError:
log.error("无法连接 ComfyUI @ %s", COMFYUI_URL)
return jsonify({"error": "Cannot connect to ComfyUI at " + COMFYUI_URL + ". Is it running?"}), 503
except Exception as e:
return jsonify({"error": str(e)}), 500
log.error("生成失败,未捕获异常:\n%s", traceback.format_exc())
return jsonify({"error": f"{type(e).__name__}: {e}"}), 500
if __name__ == "__main__":