feat: implement FastAPI backend and /api/clean endpoint (#3)
This commit is contained in:
parent
932352d85c
commit
d7c43f7f2f
3 changed files with 335 additions and 1 deletions
114
app.py
Normal file
114
app.py
Normal file
|
|
@ -0,0 +1,114 @@
|
||||||
|
"""FastAPI backend for garment photo cleaning."""
|
||||||
|
|
||||||
|
import os
|
||||||
|
from pathlib import Path
|
||||||
|
from fastapi import FastAPI, File, Form, HTTPException, UploadFile, status
|
||||||
|
from fastapi.middleware.cors import CORSMiddleware
|
||||||
|
from fastapi.responses import FileResponse, JSONResponse
|
||||||
|
from fastapi.staticfiles import StaticFiles
|
||||||
|
|
||||||
|
from cleaner import (
|
||||||
|
DEFAULT_MODEL,
|
||||||
|
CleanerError,
|
||||||
|
_get_image_info,
|
||||||
|
build_prompt,
|
||||||
|
clean_garment,
|
||||||
|
encode_image_bytes,
|
||||||
|
)
|
||||||
|
|
||||||
|
DEFAULT_PROMPT = build_prompt("beige")
|
||||||
|
STATIC_DIR = Path(__file__).resolve().parent / "static"
|
||||||
|
STATIC_DIR.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
app = FastAPI(title="unwrap-clothes")
|
||||||
|
|
||||||
|
app.add_middleware(
|
||||||
|
CORSMiddleware,
|
||||||
|
allow_origins=["*"],
|
||||||
|
allow_credentials=True,
|
||||||
|
allow_methods=["*"],
|
||||||
|
allow_headers=["*"],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@app.get("/api/health")
|
||||||
|
def health_check():
|
||||||
|
return {"status": "ok"}
|
||||||
|
|
||||||
|
|
||||||
|
@app.get("/")
|
||||||
|
def read_root():
|
||||||
|
index_file = STATIC_DIR / "index.html"
|
||||||
|
if index_file.is_file():
|
||||||
|
return FileResponse(index_file)
|
||||||
|
return JSONResponse({"status": "ok", "message": "unwrap-clothes backend is running"})
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
@app.post("/api/clean")
|
||||||
|
async def clean_endpoint(
|
||||||
|
file: UploadFile = File(...),
|
||||||
|
prompt: str = Form(default=DEFAULT_PROMPT),
|
||||||
|
model: str = Form(default=DEFAULT_MODEL),
|
||||||
|
api_key: str | None = Form(default=None),
|
||||||
|
restore_res: bool = Form(default=True),
|
||||||
|
):
|
||||||
|
effective_api_key = (api_key.strip() if api_key and api_key.strip() else None) or os.environ.get(
|
||||||
|
"OPENROUTER_API_KEY"
|
||||||
|
)
|
||||||
|
if not effective_api_key:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_400_BAD_REQUEST,
|
||||||
|
detail="API key is required. Provide it in the form or set OPENROUTER_API_KEY.",
|
||||||
|
)
|
||||||
|
|
||||||
|
file_bytes = await file.read()
|
||||||
|
if not file_bytes:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_400_BAD_REQUEST,
|
||||||
|
detail="Uploaded file is empty.",
|
||||||
|
)
|
||||||
|
|
||||||
|
orig_dims, _ = _get_image_info(file_bytes)
|
||||||
|
if orig_dims is None:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_400_BAD_REQUEST,
|
||||||
|
detail="Invalid image file. Could not decode image.",
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
clean_result = clean_garment(
|
||||||
|
image_data=file_bytes,
|
||||||
|
prompt=prompt,
|
||||||
|
api_key=effective_api_key,
|
||||||
|
model=model,
|
||||||
|
restore_res=restore_res,
|
||||||
|
)
|
||||||
|
except CleanerError as exc:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||||
|
detail=str(exc),
|
||||||
|
)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"image": encode_image_bytes(clean_result.image_bytes, mime=clean_result.media_type),
|
||||||
|
"cost": clean_result.cost,
|
||||||
|
"width": clean_result.width,
|
||||||
|
"height": clean_result.height,
|
||||||
|
"original_dimensions": list(clean_result.original_dimensions)
|
||||||
|
if clean_result.original_dimensions
|
||||||
|
else None,
|
||||||
|
"model_dimensions": list(clean_result.model_dimensions)
|
||||||
|
if clean_result.model_dimensions
|
||||||
|
else None,
|
||||||
|
"was_rescaled": clean_result.was_rescaled,
|
||||||
|
"media_type": clean_result.media_type,
|
||||||
|
}
|
||||||
|
|
||||||
|
if STATIC_DIR.is_dir():
|
||||||
|
app.mount("/static", StaticFiles(directory=STATIC_DIR), name="static")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
import uvicorn
|
||||||
|
|
||||||
|
uvicorn.run("app:app", host="127.0.0.1", port=8000, reload=True)
|
||||||
|
|
@ -1,2 +1,5 @@
|
||||||
requests>=2.28.0
|
requests>=2.28.0
|
||||||
Pillow>=9.0.0
|
Pillow>=9.0.0
|
||||||
|
fastapi>=0.100.0
|
||||||
|
uvicorn>=0.20.0
|
||||||
|
python-multipart>=0.0.6
|
||||||
|
|
|
||||||
217
tests/test_app.py
Normal file
217
tests/test_app.py
Normal file
|
|
@ -0,0 +1,217 @@
|
||||||
|
"""Tests for FastAPI backend in app.py."""
|
||||||
|
|
||||||
|
import base64
|
||||||
|
import io
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
from fastapi.testclient import TestClient
|
||||||
|
from PIL import Image
|
||||||
|
|
||||||
|
from cleaner import CleanerError, CleanResult
|
||||||
|
|
||||||
|
|
||||||
|
def test_health_check():
|
||||||
|
"""GET /api/health returns 200 OK with status ok."""
|
||||||
|
from app import app
|
||||||
|
|
||||||
|
client = TestClient(app)
|
||||||
|
response = client.get("/api/health")
|
||||||
|
assert response.status_code == 200
|
||||||
|
assert response.json() == {"status": "ok"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_root_fallback_when_no_index_html(monkeypatch, tmp_path):
|
||||||
|
"""GET / returns a graceful fallback when static/index.html does not exist."""
|
||||||
|
import app as app_module
|
||||||
|
|
||||||
|
# Point STATIC_DIR to a temporary directory without index.html
|
||||||
|
monkeypatch.setattr(app_module, "STATIC_DIR", tmp_path)
|
||||||
|
client = TestClient(app_module.app)
|
||||||
|
response = client.get("/")
|
||||||
|
assert response.status_code == 200
|
||||||
|
assert response.json()["status"] == "ok"
|
||||||
|
|
||||||
|
|
||||||
|
def test_root_serves_index_html_when_present(monkeypatch, tmp_path):
|
||||||
|
"""GET / serves static/index.html when it exists."""
|
||||||
|
import app as app_module
|
||||||
|
|
||||||
|
index_file = tmp_path / "index.html"
|
||||||
|
index_file.write_text("<!DOCTYPE html><html><body>Unwrap Clothes</body></html>")
|
||||||
|
monkeypatch.setattr(app_module, "STATIC_DIR", tmp_path)
|
||||||
|
|
||||||
|
client = TestClient(app_module.app)
|
||||||
|
response = client.get("/")
|
||||||
|
assert response.status_code == 200
|
||||||
|
assert "Unwrap Clothes" in response.text
|
||||||
|
assert "text/html" in response.headers.get("content-type", "")
|
||||||
|
|
||||||
|
|
||||||
|
def _create_test_image_bytes() -> bytes:
|
||||||
|
buf = io.BytesIO()
|
||||||
|
Image.new("RGB", (50, 50), color="blue").save(buf, format="JPEG")
|
||||||
|
return buf.getvalue()
|
||||||
|
|
||||||
|
|
||||||
|
def test_clean_missing_api_key_returns_400(monkeypatch):
|
||||||
|
"""POST /api/clean without API key when OPENROUTER_API_KEY is unset returns 400."""
|
||||||
|
from app import app
|
||||||
|
|
||||||
|
monkeypatch.delenv("OPENROUTER_API_KEY", raising=False)
|
||||||
|
client = TestClient(app)
|
||||||
|
image_bytes = _create_test_image_bytes()
|
||||||
|
files = {"file": ("test.jpg", image_bytes, "image/jpeg")}
|
||||||
|
|
||||||
|
response = client.post("/api/clean", files=files)
|
||||||
|
assert response.status_code == 400
|
||||||
|
assert "API key" in response.json()["detail"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_clean_uses_env_api_key_when_form_omitted(monkeypatch):
|
||||||
|
"""POST /api/clean uses OPENROUTER_API_KEY from environment if api_key form field is omitted."""
|
||||||
|
from app import app
|
||||||
|
|
||||||
|
monkeypatch.setenv("OPENROUTER_API_KEY", "sk-or-v1-env-key")
|
||||||
|
client = TestClient(app)
|
||||||
|
image_bytes = _create_test_image_bytes()
|
||||||
|
files = {"file": ("test.jpg", image_bytes, "image/jpeg")}
|
||||||
|
|
||||||
|
mock_result = CleanResult(
|
||||||
|
image_bytes=image_bytes,
|
||||||
|
media_type="image/jpeg",
|
||||||
|
cost=0.005,
|
||||||
|
width=50,
|
||||||
|
height=50,
|
||||||
|
original_dimensions=(50, 50),
|
||||||
|
model_dimensions=(50, 50),
|
||||||
|
)
|
||||||
|
with patch("app.clean_garment", return_value=mock_result) as mock_clean:
|
||||||
|
response = client.post("/api/clean", files=files)
|
||||||
|
assert response.status_code == 200
|
||||||
|
mock_clean.assert_called_once()
|
||||||
|
_, kwargs = mock_clean.call_args
|
||||||
|
assert kwargs["api_key"] == "sk-or-v1-env-key"
|
||||||
|
|
||||||
|
|
||||||
|
def test_clean_empty_file_returns_400(monkeypatch):
|
||||||
|
"""POST /api/clean with empty file returns 400 Bad Request."""
|
||||||
|
from app import app
|
||||||
|
|
||||||
|
monkeypatch.setenv("OPENROUTER_API_KEY", "sk-or-v1-env-key")
|
||||||
|
client = TestClient(app)
|
||||||
|
files = {"file": ("empty.jpg", b"", "image/jpeg")}
|
||||||
|
|
||||||
|
response = client.post("/api/clean", files=files)
|
||||||
|
assert response.status_code == 400
|
||||||
|
assert "empty" in response.json()["detail"].lower()
|
||||||
|
|
||||||
|
|
||||||
|
def test_clean_invalid_image_file_returns_400(monkeypatch):
|
||||||
|
"""POST /api/clean with invalid image content returns 400 Bad Request."""
|
||||||
|
from app import app
|
||||||
|
|
||||||
|
monkeypatch.setenv("OPENROUTER_API_KEY", "sk-or-v1-env-key")
|
||||||
|
client = TestClient(app)
|
||||||
|
files = {"file": ("corrupt.jpg", b"not-a-valid-image-content", "image/jpeg")}
|
||||||
|
|
||||||
|
response = client.post("/api/clean", files=files)
|
||||||
|
assert response.status_code == 400
|
||||||
|
assert "invalid" in response.json()["detail"].lower()
|
||||||
|
|
||||||
|
|
||||||
|
def test_clean_cleaner_error_returns_502(monkeypatch):
|
||||||
|
"""POST /api/clean when clean_garment raises CleanerError returns 502 Bad Gateway."""
|
||||||
|
from app import app
|
||||||
|
|
||||||
|
monkeypatch.setenv("OPENROUTER_API_KEY", "sk-or-v1-env-key")
|
||||||
|
client = TestClient(app)
|
||||||
|
image_bytes = _create_test_image_bytes()
|
||||||
|
files = {"file": ("test.jpg", image_bytes, "image/jpeg")}
|
||||||
|
|
||||||
|
with patch("app.clean_garment", side_effect=CleanerError("OpenRouter upstream failure")):
|
||||||
|
response = client.post("/api/clean", files=files)
|
||||||
|
assert response.status_code == 502
|
||||||
|
assert "OpenRouter upstream failure" in response.json()["detail"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_clean_success_happy_path():
|
||||||
|
"""POST /api/clean with valid inputs returns full JSON result with image data URI."""
|
||||||
|
from app import app
|
||||||
|
|
||||||
|
client = TestClient(app)
|
||||||
|
image_bytes = _create_test_image_bytes()
|
||||||
|
output_bytes = b"cleaned-image-data-bytes"
|
||||||
|
mock_result = CleanResult(
|
||||||
|
image_bytes=output_bytes,
|
||||||
|
media_type="image/png",
|
||||||
|
cost=0.0125,
|
||||||
|
width=1000,
|
||||||
|
height=1500,
|
||||||
|
original_dimensions=(1000, 1500),
|
||||||
|
model_dimensions=(800, 1200),
|
||||||
|
)
|
||||||
|
|
||||||
|
files = {"file": ("garment.jpg", image_bytes, "image/jpeg")}
|
||||||
|
data = {
|
||||||
|
"prompt": "Custom studio prompt with white background.",
|
||||||
|
"model": "custom/test-model",
|
||||||
|
"api_key": "sk-or-v1-custom-form-key",
|
||||||
|
"restore_res": "true",
|
||||||
|
}
|
||||||
|
|
||||||
|
with patch("app.clean_garment", return_value=mock_result) as mock_clean:
|
||||||
|
response = client.post("/api/clean", files=files, data=data)
|
||||||
|
assert response.status_code == 200
|
||||||
|
payload = response.json()
|
||||||
|
|
||||||
|
# Assert response keys and values
|
||||||
|
expected_data_uri = f"data:image/png;base64,{base64.b64encode(output_bytes).decode()}"
|
||||||
|
assert payload["image"] == expected_data_uri
|
||||||
|
assert payload["cost"] == 0.0125
|
||||||
|
assert payload["width"] == 1000
|
||||||
|
assert payload["height"] == 1500
|
||||||
|
assert payload["original_dimensions"] == [1000, 1500]
|
||||||
|
assert payload["model_dimensions"] == [800, 1200]
|
||||||
|
assert payload["was_rescaled"] is True
|
||||||
|
assert payload["media_type"] == "image/png"
|
||||||
|
|
||||||
|
# Assert arguments passed to clean_garment
|
||||||
|
mock_clean.assert_called_once_with(
|
||||||
|
image_data=image_bytes,
|
||||||
|
prompt="Custom studio prompt with white background.",
|
||||||
|
api_key="sk-or-v1-custom-form-key",
|
||||||
|
model="custom/test-model",
|
||||||
|
restore_res=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_clean_default_parameters():
|
||||||
|
"""POST /api/clean uses default prompt, model, and restore_res if omitted in form."""
|
||||||
|
from app import app, DEFAULT_PROMPT
|
||||||
|
from cleaner import DEFAULT_MODEL
|
||||||
|
|
||||||
|
client = TestClient(app)
|
||||||
|
image_bytes = _create_test_image_bytes()
|
||||||
|
mock_result = CleanResult(
|
||||||
|
image_bytes=image_bytes,
|
||||||
|
media_type="image/jpeg",
|
||||||
|
cost=None,
|
||||||
|
width=50,
|
||||||
|
height=50,
|
||||||
|
original_dimensions=(50, 50),
|
||||||
|
model_dimensions=(50, 50),
|
||||||
|
)
|
||||||
|
|
||||||
|
files = {"file": ("garment.jpg", image_bytes, "image/jpeg")}
|
||||||
|
data = {"api_key": "sk-or-v1-key"}
|
||||||
|
|
||||||
|
with patch("app.clean_garment", return_value=mock_result) as mock_clean:
|
||||||
|
response = client.post("/api/clean", files=files, data=data)
|
||||||
|
assert response.status_code == 200
|
||||||
|
mock_clean.assert_called_once_with(
|
||||||
|
image_data=image_bytes,
|
||||||
|
prompt=DEFAULT_PROMPT,
|
||||||
|
api_key="sk-or-v1-key",
|
||||||
|
model=DEFAULT_MODEL,
|
||||||
|
restore_res=True,
|
||||||
|
)
|
||||||
Loading…
Add table
Reference in a new issue