#!/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// 下") @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// 下") @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)