262 lines
8.6 KiB
Python
262 lines
8.6 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
|
|
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
|
|
|
|
import oss2
|
|
|
|
COMFYUI_URL = "http://127.0.0.1:8188"
|
|
# COMFYUI_URL = "http://117.50.80.187:47697"
|
|
# WORKFLOW_PATH = Path(__file__).parent / "change2_1_ex.json"
|
|
|
|
WORKFLOW_PATH = Path(__file__).parent / "change_fast2.json"
|
|
|
|
# 节点 ID 映射
|
|
NODE_MODEL = "11" # 模特
|
|
NODE_SHIRT = "10" # 上衣
|
|
# NODE_PANTS = "4" # 裤子
|
|
NODE_OUTPUT = "33" # 输出
|
|
|
|
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' # 替换为你的Endpoint
|
|
bucket_name = 'xiangsilian'
|
|
|
|
# 创建Bucket实例
|
|
auth = oss2.Auth(access_key_id, access_key_secret)
|
|
bucket = oss2.Bucket(auth, endpoint, bucket_name)
|
|
|
|
# 如果没有指定OSS文件名,则使用本地文件名
|
|
if object_name is None:
|
|
object_name = image_path.split('/')[-1] # 取本地文件名
|
|
|
|
try:
|
|
# 上传文件
|
|
bucket.put_object_from_file(object_name, image_path)
|
|
|
|
# 获取文件URL(有效期默认10年)
|
|
url = bucket.sign_url('GET', object_name, 3600 * 24 * 365 * 10)
|
|
|
|
# 或者使用公共读Bucket的URL(如果Bucket是公共读权限)
|
|
# url = f"https://{bucket_name}.{endpoint}/{object_name}"
|
|
|
|
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 上衣图片
|
|
|
|
|
|
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 upload_image(client: httpx.AsyncClient, image_bytes: bytes, filename: str) -> str:
|
|
"""上传图片到 ComfyUI,返回服务器文件名"""
|
|
files = {
|
|
"image": (filename, io.BytesIO(image_bytes), "image/jpeg"),
|
|
}
|
|
data = {"overwrite": "true"}
|
|
resp = await client.post(f"{COMFYUI_URL}/upload/image", files=files, data=data)
|
|
resp.raise_for_status()
|
|
result = resp.json()
|
|
return result["name"]
|
|
|
|
|
|
async def queue_prompt(client: httpx.AsyncClient, workflow: dict) -> str:
|
|
"""提交工作流到队列,返回 prompt_id"""
|
|
payload = {"prompt": workflow, "client_id": str(uuid.uuid4())}
|
|
resp = await client.post(f"{COMFYUI_URL}/prompt", json=payload)
|
|
resp.raise_for_status()
|
|
return resp.json()["prompt_id"]
|
|
|
|
|
|
async def wait_for_result(client: httpx.AsyncClient, prompt_id: str, timeout: int = 300) -> dict:
|
|
"""轮询历史记录,等待任务完成,返回输出节点数据"""
|
|
deadline = time.time() + timeout
|
|
while time.time() < deadline:
|
|
resp = await client.get(f"{COMFYUI_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, filename: str, subfolder: str = "", type_: str = "output") -> bytes:
|
|
"""从 ComfyUI 下载图片,返回原始字节"""
|
|
params = {"filename": filename, "subfolder": subfolder, "type": type_}
|
|
resp = await client.get(f"{COMFYUI_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
|
|
返回 result_image: 换装结果 base64
|
|
"""
|
|
|
|
print(f"接收到换装请求,正在处理...")
|
|
# 加载工作流模板
|
|
workflow = json.loads(WORKFLOW_PATH.read_text(encoding="utf-8"))
|
|
|
|
|
|
|
|
async with httpx.AsyncClient(timeout=60.0) as client:
|
|
# 解码两张图片
|
|
try:
|
|
model_bytes = decode_base64_image(req.model_image)
|
|
shirt_bytes = decode_base64_image(req.shirt_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"
|
|
|
|
# 并发上传两张图片
|
|
try:
|
|
model_name, shirt_name = await asyncio.gather(
|
|
upload_image(client, model_bytes, model_fname),
|
|
upload_image(client, shirt_bytes, shirt_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
|
|
|
|
# 提交工作流
|
|
try:
|
|
prompt_id = await queue_prompt(client, workflow)
|
|
except httpx.HTTPError as e:
|
|
raise HTTPException(status_code=502, detail=f"提交工作流失败: {e}")
|
|
|
|
# 等待结果(最多5分钟)
|
|
async with httpx.AsyncClient(timeout=30.0) as poll_client:
|
|
outputs = await wait_for_result(poll_client, prompt_id, timeout=300)
|
|
|
|
# 从节点33的输出获取图片
|
|
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) as dl_client:
|
|
result_bytes = await fetch_image_bytes(dl_client, 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():
|
|
"""健康检查"""
|
|
return {"status": "ok", "comfyui": COMFYUI_URL}
|
|
|
|
|
|
# 挂载静态文件(前端页面)
|
|
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=12222, log_level="info")
|