You can not select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
144 lines
5.5 KiB
144 lines
5.5 KiB
#!/usr/bin/env python3
|
|
"""FastAPI 入口:AI 图片生成服务。"""
|
|
|
|
from dotenv import load_dotenv
|
|
|
|
load_dotenv()
|
|
import base64
|
|
import re
|
|
from fastapi import FastAPI, File, Form, HTTPException, UploadFile
|
|
from fastapi.staticfiles import StaticFiles
|
|
from pydantic import BaseModel, Field, model_validator
|
|
|
|
from utils.mm_api_t2i import OUTPUT_DIR, PROVIDERS, generate_images
|
|
from utils.agnes_image import VALID_RATIOS, generate_images as agnes_generate_images
|
|
|
|
app = FastAPI(title="AI Image T2I API", version="1.0.0")
|
|
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
|
|
app.mount("/output", StaticFiles(directory=OUTPUT_DIR), name="output")
|
|
|
|
|
|
class ImageGenerationRequest(BaseModel):
|
|
prompt: str = Field(min_length=1, description="图片提示词")
|
|
provider: str = Field(default="grok", pattern="^(grok|gpt)$", description="模型提供方")
|
|
model: str | None = Field(default=None, description="覆盖提供方默认模型")
|
|
size: str = Field(default="1024x1024", description="图片尺寸")
|
|
n: int = Field(default=1, ge=1, le=10, description="生成数量")
|
|
profile_name: str = Field(pattern=r"^[\w-]+$", description="调用方标识(必填),图片会保存到 output/<profile_name>/ 下")
|
|
|
|
@model_validator(mode="before")
|
|
@classmethod
|
|
def check_profile_name(cls, data):
|
|
if isinstance(data, dict) and not data.get("profile_name"):
|
|
raise ValueError("缺少 profile_name 参数,请在请求中带上你的 agent/profile 名称")
|
|
return data
|
|
|
|
|
|
class GeneratedImage(BaseModel):
|
|
filename: str
|
|
url: str
|
|
bytes: int
|
|
|
|
|
|
class ImageGenerationResponse(BaseModel):
|
|
provider: str
|
|
model: str
|
|
images: list[GeneratedImage]
|
|
|
|
|
|
@app.get("/health")
|
|
def health_check():
|
|
return {"status": "ok"}
|
|
|
|
|
|
@app.post("/images/generations", response_model=ImageGenerationResponse)
|
|
def generate_images_endpoint(request: ImageGenerationRequest):
|
|
result = generate_images(
|
|
provider=request.provider,
|
|
model=request.model,
|
|
prompt=request.prompt,
|
|
size=request.size,
|
|
n=request.n,
|
|
profile_name=request.profile_name,
|
|
)
|
|
return ImageGenerationResponse(**result)
|
|
|
|
|
|
class AgnesImageGenerationRequest(BaseModel):
|
|
prompt: str = Field(min_length=1, description="图片提示词")
|
|
size: str = Field(default="1K", description="输出尺寸档位: 1K / 2K / 3K / 4K,也支持精确尺寸如 1024x1024")
|
|
ratio: str | None = Field(default=None, description=f"宽高比,支持 {sorted(VALID_RATIOS)};与 size 档位配合使用")
|
|
images: list[str] | None = Field(default=None, description="图生图 / 多图合成的输入图像 URL 或 Data URI Base64")
|
|
profile_name: str = Field(pattern=r"^[\w-]+$", description="调用方标识(必填),图片会保存到 output/<profile_name>/ 下")
|
|
|
|
@model_validator(mode="before")
|
|
@classmethod
|
|
def check_profile_name(cls, data):
|
|
if isinstance(data, dict) and not data.get("profile_name"):
|
|
raise ValueError("缺少 profile_name 参数,请在请求中带上你的 agent/profile 名称")
|
|
return data
|
|
|
|
|
|
class AgnesImageGenerationResponse(BaseModel):
|
|
model: str
|
|
size: str
|
|
ratio: str
|
|
images: list[GeneratedImage]
|
|
|
|
|
|
@app.post("/agnes/images/generations", response_model=AgnesImageGenerationResponse)
|
|
def agnes_generate_images_endpoint(request: AgnesImageGenerationRequest):
|
|
result = agnes_generate_images(
|
|
prompt=request.prompt,
|
|
size=request.size,
|
|
ratio=request.ratio,
|
|
images=request.images,
|
|
profile_name=request.profile_name,
|
|
)
|
|
return AgnesImageGenerationResponse(**result)
|
|
|
|
|
|
_ALLOWED_IMAGE_TYPES = {"image/png", "image/jpeg", "image/jpg", "image/webp"}
|
|
|
|
|
|
@app.post("/agnes/images/generations/upload", response_model=AgnesImageGenerationResponse)
|
|
def agnes_generate_images_upload_endpoint(
|
|
prompt: str = Form(..., description="图片提示词"),
|
|
profile_name: str = Form(..., description="调用方标识(必填)"),
|
|
size: str = Form(default="1K", description="输出尺寸档位: 1K / 2K / 3K / 4K"),
|
|
ratio: str | None = Form(default=None, description=f"宽高比,支持 {sorted(VALID_RATIOS)}"),
|
|
images: list[UploadFile] = File(default=[], description="参考图像文件,可传多个"),
|
|
):
|
|
"""multipart/form-data 版图生图:上传本地图片文件直接测试,无需 base64。"""
|
|
if not prompt.strip():
|
|
raise HTTPException(status_code=422, detail="prompt 不能为空")
|
|
if not re.match(r"^[\w-]+$", profile_name):
|
|
raise HTTPException(status_code=422, detail="profile_name 只能包含字母数字下划线和短横线")
|
|
if not images or all(not f.filename for f in images):
|
|
raise HTTPException(status_code=422, detail="必须至少上传一张参考图 images")
|
|
|
|
data_uris: list[str] = []
|
|
for f in images:
|
|
if f.content_type and f.content_type not in _ALLOWED_IMAGE_TYPES:
|
|
raise HTTPException(status_code=400, detail=f"不支持的图片类型: {f.content_type}")
|
|
raw = f.file.read()
|
|
if not raw:
|
|
raise HTTPException(status_code=400, detail=f"上传的图片 {f.filename} 为空")
|
|
mime = f.content_type or "image/png"
|
|
b64 = base64.b64encode(raw).decode("ascii")
|
|
data_uris.append(f"data:{mime};base64,{b64}")
|
|
|
|
result = agnes_generate_images(
|
|
prompt=prompt,
|
|
size=size,
|
|
ratio=ratio,
|
|
images=data_uris,
|
|
profile_name=profile_name,
|
|
)
|
|
return AgnesImageGenerationResponse(**result)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
import uvicorn
|
|
|
|
uvicorn.run(app, host="0.0.0.0", port=8000)
|
|
|