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.
 
 
ai-image/main.py

103 lines
3.6 KiB

#!/usr/bin/env python3
"""FastAPI 入口:AI 图片生成服务。"""
from dotenv import load_dotenv
load_dotenv()
from fastapi import FastAPI
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)
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=8000)