feat: implement images API (optimize, generate)
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -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.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)}")
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user