297 lines
10 KiB
Python
297 lines
10 KiB
Python
#!/usr/bin/env python3
|
||
"""
|
||
ComfyUI 换装服务
|
||
- 接受3个base64图片(模特、上衣、裤子)
|
||
- 上传到 ComfyUI,运行工作流
|
||
- 返回生成结果 base64 图片
|
||
"""
|
||
|
||
import asyncio
|
||
import base64
|
||
import io
|
||
import json
|
||
import os
|
||
import tempfile
|
||
import time
|
||
import uuid
|
||
from pathlib import Path
|
||
|
||
import httpx
|
||
import oss2
|
||
from PIL import Image
|
||
from fastapi import FastAPI, HTTPException
|
||
from fastapi.middleware.cors import CORSMiddleware
|
||
from fastapi.responses import FileResponse
|
||
from fastapi.staticfiles import StaticFiles
|
||
from pydantic import BaseModel
|
||
|
||
# COMFYUI_URL = "http://127.0.0.1:8188"
|
||
COMFYUI_URL = "http://112.126.94.241:28188"
|
||
COMFYUI_URL_BACKUP = "http://112.126.94.241:38188"
|
||
WORKFLOW_PATH = Path(__file__).parent / "change2_2_0308.json"
|
||
|
||
_cred_dir = Path(__file__).parent.parent
|
||
COMFYUI_USER = (_cred_dir / "user.txt").read_text(encoding="utf-8").strip()
|
||
COMFYUI_PASS = (_cred_dir / "password.txt").read_text(encoding="utf-8").strip()
|
||
COMFYUI_AUTH = (COMFYUI_USER, COMFYUI_PASS)
|
||
|
||
# WORKFLOW_PATH = Path(__file__).parent / "change_fast3.json"
|
||
|
||
# 节点 ID 映射
|
||
NODE_MODEL = "11" # 模特
|
||
NODE_SHIRT = "10" # 上衣
|
||
NODE_PANTS = "4" # 裤子
|
||
NODE_OUTPUT = "33" # 输出
|
||
NODE_PROMPT = "31" # 文本提示(PrimitiveStringMultiline)
|
||
|
||
app = FastAPI(title="ComfyUI 换装服务")
|
||
|
||
app.add_middleware(
|
||
CORSMiddleware,
|
||
allow_origins=["*"],
|
||
allow_methods=["*"],
|
||
allow_headers=["*"],
|
||
)
|
||
|
||
|
||
def upload_to_oss(image_path, object_name=None):
|
||
access_key_id = 'LTAI5tGp1sLzedqxihcNC1eb'
|
||
access_key_secret = 'IFZE1b8YYreCP6zfA6GaZ9uBT678qO'
|
||
endpoint = 'oss-cn-beijing.aliyuncs.com'
|
||
bucket_name = 'xiangsilian'
|
||
|
||
auth = oss2.Auth(access_key_id, access_key_secret)
|
||
bucket = oss2.Bucket(auth, endpoint, bucket_name)
|
||
|
||
if object_name is None:
|
||
object_name = image_path.split('/')[-1]
|
||
|
||
try:
|
||
bucket.put_object_from_file(object_name, image_path)
|
||
url = bucket.sign_url('GET', object_name, 3600 * 24 * 365 * 10)
|
||
print(f"文件上传成功,URL: {url}")
|
||
return url
|
||
except Exception as e:
|
||
print(f"上传失败: {str(e)}")
|
||
return None
|
||
|
||
|
||
class TryOnRequest(BaseModel):
|
||
model_image: str # base64 模特图片
|
||
shirt_image: str # base64 上衣图片
|
||
pants_image: str # base64 裤子图片
|
||
change_desc: str = "" # 换装/编辑描述,可空;非空时替换节点 31 的 inputs
|
||
|
||
|
||
class TryOnResponse(BaseModel):
|
||
result_url: str # OSS 图片 URL
|
||
|
||
|
||
def decode_base64_image(b64_str: str) -> bytes:
|
||
"""解码 base64 图片,自动处理 data:image/... 前缀"""
|
||
if "," in b64_str:
|
||
b64_str = b64_str.split(",", 1)[1]
|
||
return base64.b64decode(b64_str)
|
||
|
||
|
||
async def check_comfyui_alive(base_url: str, timeout: float = 3.0) -> bool:
|
||
"""检查指定的 ComfyUI 服务器是否可用(能连通且非 5xx 即视为可用)"""
|
||
try:
|
||
async with httpx.AsyncClient(timeout=timeout, auth=COMFYUI_AUTH) as client:
|
||
resp = await client.get(f"{base_url}/system_stats")
|
||
return resp.status_code < 500
|
||
except Exception as e:
|
||
print(f"ComfyUI 健康检查失败 {base_url}: {e}")
|
||
return False
|
||
|
||
|
||
async def pick_comfyui_url() -> str:
|
||
"""优先返回正式服务器,不可用时切换到备份服务器"""
|
||
if await check_comfyui_alive(COMFYUI_URL):
|
||
print(f"使用正式服务器: {COMFYUI_URL}")
|
||
return COMFYUI_URL
|
||
print(f"正式服务器不可用,尝试切换到备份服务器: {COMFYUI_URL_BACKUP}")
|
||
if await check_comfyui_alive(COMFYUI_URL_BACKUP):
|
||
print(f"使用备份服务器: {COMFYUI_URL_BACKUP}")
|
||
return COMFYUI_URL_BACKUP
|
||
raise HTTPException(status_code=503, detail="正式与备份 ComfyUI 服务器均不可用")
|
||
|
||
|
||
async def upload_image(client: httpx.AsyncClient, base_url: str, image_bytes: bytes, filename: str) -> str:
|
||
"""上传图片到 ComfyUI,返回服务器文件名"""
|
||
files = {
|
||
"image": (filename, io.BytesIO(image_bytes), "image/jpeg"),
|
||
}
|
||
data = {"overwrite": "true"}
|
||
resp = await client.post(f"{base_url}/upload/image", files=files, data=data)
|
||
resp.raise_for_status()
|
||
result = resp.json()
|
||
return result["name"]
|
||
|
||
|
||
async def queue_prompt(client: httpx.AsyncClient, base_url: str, workflow: dict) -> str:
|
||
"""提交工作流到队列,返回 prompt_id"""
|
||
payload = {"prompt": workflow, "client_id": str(uuid.uuid4())}
|
||
resp = await client.post(f"{base_url}/prompt", json=payload)
|
||
resp.raise_for_status()
|
||
return resp.json()["prompt_id"]
|
||
|
||
|
||
async def wait_for_result(client: httpx.AsyncClient, base_url: str, prompt_id: str, timeout: int = 300) -> dict:
|
||
"""轮询历史记录,等待任务完成,返回输出节点数据"""
|
||
deadline = time.time() + timeout
|
||
while time.time() < deadline:
|
||
resp = await client.get(f"{base_url}/history/{prompt_id}")
|
||
resp.raise_for_status()
|
||
history = resp.json()
|
||
if prompt_id in history:
|
||
entry = history[prompt_id]
|
||
status = entry.get("status", {})
|
||
if status.get("completed"):
|
||
return entry.get("outputs", {})
|
||
if status.get("status_str") == "error":
|
||
messages = status.get("messages", [])
|
||
raise HTTPException(status_code=500, detail=f"ComfyUI 工作流执行出错: {messages}")
|
||
await asyncio.sleep(2)
|
||
raise HTTPException(status_code=504, detail="等待 ComfyUI 超时(300秒)")
|
||
|
||
|
||
async def fetch_image_bytes(client: httpx.AsyncClient, base_url: str, filename: str, subfolder: str = "", type_: str = "output") -> bytes:
|
||
"""从 ComfyUI 下载图片,返回原始字节"""
|
||
params = {"filename": filename, "subfolder": subfolder, "type": type_}
|
||
resp = await client.get(f"{base_url}/view", params=params)
|
||
resp.raise_for_status()
|
||
return resp.content
|
||
|
||
|
||
@app.post("/try-on", response_model=TryOnResponse)
|
||
async def try_on(req: TryOnRequest):
|
||
"""
|
||
换装接口
|
||
- model_image: 模特图片 base64
|
||
- shirt_image: 上衣图片 base64
|
||
- pants_image: 裤子图片 base64
|
||
返回 result_image: 换装结果 base64
|
||
"""
|
||
print(f"接收到换装请求,正在处理...")
|
||
|
||
comfyui_url = await pick_comfyui_url()
|
||
|
||
workflow = json.loads(WORKFLOW_PATH.read_text(encoding="utf-8"))
|
||
|
||
desc = (req.change_desc or "").strip()
|
||
if desc:
|
||
workflow[NODE_PROMPT]["inputs"] = {"value": desc}
|
||
|
||
async with httpx.AsyncClient(timeout=60.0, auth=COMFYUI_AUTH) as client:
|
||
try:
|
||
model_bytes = decode_base64_image(req.model_image)
|
||
shirt_bytes = decode_base64_image(req.shirt_image)
|
||
pants_bytes = decode_base64_image(req.pants_image)
|
||
except Exception as e:
|
||
raise HTTPException(status_code=400, detail=f"图片解码失败: {e}")
|
||
|
||
uid = uuid.uuid4().hex[:8]
|
||
model_fname = f"model_{uid}.jpg"
|
||
shirt_fname = f"shirt_{uid}.jpg"
|
||
pants_fname = f"pants_{uid}.jpg"
|
||
|
||
try:
|
||
model_name, shirt_name, pants_name = await asyncio.gather(
|
||
upload_image(client, comfyui_url, model_bytes, model_fname),
|
||
upload_image(client, comfyui_url, shirt_bytes, shirt_fname),
|
||
upload_image(client, comfyui_url, pants_bytes, pants_fname),
|
||
)
|
||
except httpx.HTTPError as e:
|
||
raise HTTPException(status_code=502, detail=f"上传图片到 ComfyUI 失败: {e}")
|
||
|
||
workflow[NODE_MODEL]["inputs"]["image"] = model_name
|
||
workflow[NODE_SHIRT]["inputs"]["image"] = shirt_name
|
||
workflow[NODE_PANTS]["inputs"]["image"] = pants_name
|
||
|
||
try:
|
||
prompt_id = await queue_prompt(client, comfyui_url, workflow)
|
||
except httpx.HTTPError as e:
|
||
raise HTTPException(status_code=502, detail=f"提交工作流失败: {e}")
|
||
|
||
async with httpx.AsyncClient(timeout=30.0, auth=COMFYUI_AUTH) as poll_client:
|
||
outputs = await wait_for_result(poll_client, comfyui_url, prompt_id, timeout=300)
|
||
|
||
node_output = outputs.get(NODE_OUTPUT)
|
||
if not node_output:
|
||
raise HTTPException(status_code=500, detail=f"工作流未返回节点 {NODE_OUTPUT} 的输出")
|
||
|
||
images = node_output.get("images", [])
|
||
if not images:
|
||
raise HTTPException(status_code=500, detail="输出节点没有图片")
|
||
|
||
img_info = images[0]
|
||
filename = img_info["filename"]
|
||
subfolder = img_info.get("subfolder", "")
|
||
type_ = img_info.get("type", "output")
|
||
|
||
async with httpx.AsyncClient(timeout=30.0, auth=COMFYUI_AUTH) as dl_client:
|
||
result_bytes = await fetch_image_bytes(dl_client, comfyui_url, filename, subfolder, type_)
|
||
|
||
# 缩放到 768x1024
|
||
# img = Image.open(io.BytesIO(result_bytes))
|
||
# img = img.resize((768, 1024), Image.LANCZOS)
|
||
# buf = io.BytesIO()
|
||
# fmt = (Path(filename).suffix.lstrip(".") or "png").upper()
|
||
# fmt = "JPEG" if fmt in ("JPG", "JPEG") else fmt
|
||
# img.save(buf, format=fmt)
|
||
# result_bytes = buf.getvalue()
|
||
|
||
# 将图片保存到临时文件,上传到 OSS
|
||
suffix = Path(filename).suffix or ".png"
|
||
object_name = f"tryon/{uuid.uuid4().hex}{suffix}"
|
||
with tempfile.NamedTemporaryFile(suffix=suffix, delete=False) as tmp:
|
||
tmp.write(result_bytes)
|
||
tmp_path = tmp.name
|
||
|
||
try:
|
||
result_url = upload_to_oss(tmp_path, object_name)
|
||
finally:
|
||
os.unlink(tmp_path)
|
||
|
||
if not result_url:
|
||
raise HTTPException(status_code=500, detail="上传结果图片到 OSS 失败")
|
||
|
||
return TryOnResponse(result_url=result_url)
|
||
|
||
|
||
@app.get("/health")
|
||
async def health():
|
||
"""健康检查"""
|
||
primary_ok = await check_comfyui_alive(COMFYUI_URL)
|
||
backup_ok = await check_comfyui_alive(COMFYUI_URL_BACKUP)
|
||
if primary_ok:
|
||
active = COMFYUI_URL
|
||
elif backup_ok:
|
||
active = COMFYUI_URL_BACKUP
|
||
else:
|
||
active = None
|
||
return {
|
||
"status": "ok" if active else "degraded",
|
||
"comfyui_active": active,
|
||
"comfyui_primary": {"url": COMFYUI_URL, "alive": primary_ok},
|
||
"comfyui_backup": {"url": COMFYUI_URL_BACKUP, "alive": backup_ok},
|
||
}
|
||
|
||
|
||
# 挂载静态文件(前端页面)
|
||
static_dir = Path(__file__).parent / "static"
|
||
static_dir.mkdir(exist_ok=True)
|
||
print(f"静态文件目录: {static_dir}")
|
||
app.mount("/static", StaticFiles(directory=str(static_dir)), name="static")
|
||
|
||
|
||
@app.get("/")
|
||
async def index():
|
||
return FileResponse(str(static_dir / "index.html"))
|
||
|
||
|
||
if __name__ == "__main__":
|
||
import uvicorn
|
||
uvicorn.run(app, host="0.0.0.0", port=12223, log_level="info")
|