Files
change_cloth_server/change2/service2.py
T
2026-03-13 03:21:38 +00:00

278 lines
9.1 KiB
Python

#!/usr/bin/env python3
"""
ComfyUI 换装服务
- 接受3个base64图片(模特、上衣、裤子)
- 上传到 ComfyUI,运行工作流
- 返回生成结果 base64 图片
"""
import asyncio
import base64
import io
import json
import os
import sys
import tempfile
import time
import uuid
from pathlib import Path
# 添加父目录到路径以导入config
sys.path.insert(0, str(Path(__file__).parent.parent))
import config
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"
# 节点 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_)
# 根据配置转换图片格式并上传到 OSS
output_format = config.OUTPUT_FORMAT.lower()
file_ext = '.jpg' if output_format == 'jpg' else '.png'
object_name = f"tryon/{uuid.uuid4().hex}{file_ext}"
# 使用PIL处理图片
img = Image.open(io.BytesIO(result_bytes))
if output_format == 'jpg':
# JPG格式:转换为RGB模式
if img.mode in ('RGBA', 'LA', 'P'):
rgb_img = Image.new('RGB', img.size, (255, 255, 255))
if img.mode == 'P':
img = img.convert('RGBA')
rgb_img.paste(img, mask=img.split()[-1] if img.mode in ('RGBA', 'LA') else None)
img = rgb_img
elif img.mode != 'RGB':
img = img.convert('RGB')
# 保存为JPG
with tempfile.NamedTemporaryFile(suffix='.jpg', delete=False) as tmp:
img.save(tmp, format='JPEG', quality=config.JPG_QUALITY)
tmp_path = tmp.name
else:
# PNG格式:保持原样
with tempfile.NamedTemporaryFile(suffix='.png', delete=False) as tmp:
img.save(tmp, format='PNG')
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")