feat: implement images API (optimize, generate)

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
Claude
2026-03-24 15:14:00 +08:00
parent 73ceb95bdd
commit ec0fcd3e48
3 changed files with 142 additions and 3 deletions
+77 -1
View File
@@ -1,3 +1,79 @@
from fastapi import APIRouter from fastapi import APIRouter, Depends, HTTPException, status
from sqlalchemy.orm import Session
from app.database import get_db
from app.api.deps import get_current_user
from app.models.user import User
from app.schemas.image import (
PromptOptimizeRequest,
PromptOptimizeResponse,
ImageGenerateRequest,
ImageGenerateResponse,
)
from app.services.prompt import PromptOptimizationService
from app.services.image import ImageGenerationService
from app.services.auth import AuthService
router = APIRouter() router = APIRouter()
@router.post("/optimize", response_model=PromptOptimizeResponse)
async def optimize_prompt(
data: PromptOptimizeRequest,
current_user: User = Depends(get_current_user),
db: Session = Depends(get_db)
):
service = PromptOptimizationService()
try:
if data.mode == "text2image":
optimized = service.optimize_text2image(data.prompt)
else:
optimized = service.optimize_image2image(data.prompt, data.original_description or "")
return PromptOptimizeResponse(optimized_prompt=optimized)
except Exception as e:
raise HTTPException(status_code=500, detail=f"Optimization failed: {str(e)}")
@router.post("/generate", response_model=ImageGenerateResponse)
async def generate_image(
data: ImageGenerateRequest,
current_user: User = Depends(get_current_user),
db: Session = Depends(get_db)
):
service = ImageGenerationService()
try:
if data.mode == "text2image":
result = await service.generate_text2image(data.optimized_prompt, data.model)
else:
if not data.input_image_b64:
raise HTTPException(status_code=400, detail="input_image_b64 required for image2image")
result = await service.generate_image2image(
data.optimized_prompt,
data.input_image_b64,
"image/png",
data.model
)
# Save image
image_url = service.save_image(result["image_b64"], str(current_user.id))
# Save to history
from app.models.history import GenerationHistory
history_item = GenerationHistory(
user_id=current_user.id,
original_prompt=data.prompt,
optimized_prompt=data.optimized_prompt,
mode=data.mode,
model=result["model"],
image_url=image_url,
input_image_url=f"/uploads/{input_image_b64.split('/')[0]}.png" if data.mode == "image2image" and data.input_image_b64 else None
)
db.add(history_item)
db.commit()
return ImageGenerateResponse(
image_url=image_url,
model=result["model"],
optimized_prompt=data.optimized_prompt
)
except Exception as e:
raise HTTPException(status_code=500, detail=f"Generation failed: {str(e)}")
+27 -1
View File
@@ -1 +1,27 @@
"""Image schemas placeholder.""" from pydantic import BaseModel
from typing import Optional
from enum import Enum
class GenerationMode(str, Enum):
TEXT2IMAGE = "text2image"
IMAGE2IMAGE = "image2image"
class PromptOptimizeRequest(BaseModel):
prompt: str
mode: GenerationMode = GenerationMode.TEXT2IMAGE
original_description: Optional[str] = None
class PromptOptimizeResponse(BaseModel):
optimized_prompt: str
class ImageGenerateRequest(BaseModel):
prompt: str
optimized_prompt: str
mode: GenerationMode
model: str = "gemini-2.5-flash-image"
input_image_b64: Optional[str] = None # For I2I
class ImageGenerateResponse(BaseModel):
image_url: str
model: str
optimized_prompt: str
@@ -0,0 +1,37 @@
import pytest
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from app.database import Base, get_db
from app.main import app
from app.models.user import User
from app.core.security import hash_password, create_access_token
engine = create_engine("sqlite:///:memory:")
TestSessionLocal = sessionmaker(bind=engine)
def override_get_db():
db = TestSessionLocal()
try:
yield db
finally:
db.close()
app.dependency_overrides[get_db] = override_get_db
client = TestClient(app)
def get_auth_header(user_id):
token = create_access_token({"sub": str(user_id)})
return {"Authorization": f"Bearer {token}"}
def test_optimize_prompt_requires_auth():
response = client.post("/api/images/optimize", json={"prompt": "a cat"})
assert response.status_code == 403 # No auth header
def test_generate_requires_auth():
response = client.post("/api/images/generate", json={
"prompt": "a cat",
"mode": "text2image",
"model": "gemini-2.5-flash-image"
})
assert response.status_code == 403