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 = 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