editify-backend / main.py
Nisha56's picture
Add PostgreSQL authentication and UI improvements
6dd7eae
Raw
History Blame Contribute Delete
8.43 kB
import os
import time
import uuid
import shutil
from fastapi import FastAPI, File, UploadFile, HTTPException, Depends
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse
from sqlalchemy.orm import Session
from datetime import timezone, timedelta
from config import UPLOAD_DIR, PROJECT_NAME
from database import Base, engine, get_db
from models import Project, User
from schemas import PipelineResponse, ProjectResponse
from crud import (
create_project,
update_project_ready,
update_project_error,
get_all_projects,
get_project_by_file_id,
save_project_state,
get_saved_state,
delete_project,
)
from pipeline import run_pipeline
from auth_utils import get_current_user
from auth import router as auth_router
Base.metadata.create_all(bind=engine)
app = FastAPI(title=PROJECT_NAME)
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
app.include_router(auth_router)
@app.get("/")
def root():
return {"message": f"{PROJECT_NAME} API is running"}
@app.get("/health")
def health_check():
return {"status": "ok", "project": PROJECT_NAME}
@app.post("/upload")
async def upload_image(
file: UploadFile = File(...),
db: Session = Depends(get_db),
current_user: User = Depends(get_current_user),
):
allowed_types = ["image/jpeg", "image/png", "image/webp"]
if file.content_type not in allowed_types:
raise HTTPException(400, "Only JPEG, PNG, WEBP images allowed")
file_id = str(uuid.uuid4())
extension = file.filename.split(".")[-1]
save_path = os.path.join(UPLOAD_DIR, f"{file_id}.{extension}")
with open(save_path, "wb") as buffer:
shutil.copyfileobj(file.file, buffer)
project = create_project(db, file.filename, file_id, save_path, current_user.id)
start = time.time()
try:
result = run_pipeline(save_path)
processing_ms = int((time.time() - start) * 1000)
update_project_ready(db, file_id, len(result["layers"]), processing_ms)
except Exception as e:
update_project_error(db, file_id)
raise HTTPException(500, f"Pipeline failed: {str(e)}")
return JSONResponse({
"project_id": project.id,
"file_id": file_id,
"filename": file.filename,
"image_w": result["image_w"],
"image_h": result["image_h"],
"background_base64": result["background_base64"],
"original_base64": result["original_base64"],
"layers": result["layers"],
"processing_time_s": result["processing_time_s"],
})
@app.get("/projects")
def list_projects(
db: Session = Depends(get_db),
current_user: User = Depends(get_current_user),
):
projects = get_all_projects(db, current_user.id)
return [
{
"id": p.id,
"filename": p.filename,
"file_id": p.file_id,
"layer_count": p.layer_count,
"status": p.status,
"processing_ms": p.processing_ms,
"created_at": (
p.created_at.replace(tzinfo=timezone.utc)
.astimezone(timezone(timedelta(hours=5, minutes=30)))
.isoformat()
),
"thumbnail_b64": p.thumbnail_b64,
}
for p in projects
]
@app.get("/projects/{file_id}")
def get_project(
file_id: str,
db: Session = Depends(get_db),
current_user: User = Depends(get_current_user),
):
project = get_project_by_file_id(db, file_id, current_user.id)
if not project:
raise HTTPException(404, "Project not found")
return {
"id": project.id,
"filename": project.filename,
"file_id": project.file_id,
"layer_count": project.layer_count,
"status": project.status,
"processing_ms": project.processing_ms,
"created_at": (
project.created_at.replace(tzinfo=timezone.utc)
.astimezone(timezone(timedelta(hours=5, minutes=30)))
.isoformat()
),
}
@app.post("/inpaint")
async def inpaint_element(
request: dict,
current_user: User = Depends(get_current_user),
):
import base64
from inpaint import inpaint_region, unload_inpaint_model, create_mask_from_bbox
file_id = request.get("file_id")
x, y = request.get("x", 0), request.get("y", 0)
w, h = request.get("w", 0), request.get("h", 0)
if not file_id:
raise HTTPException(400, "file_id is required")
upload_files = os.listdir(UPLOAD_DIR)
image_path = None
for f in upload_files:
if f.startswith(file_id):
image_path = os.path.join(UPLOAD_DIR, f)
break
if not image_path:
raise HTTPException(404, "Original image not found")
try:
mask = create_mask_from_bbox(image_path, x, y, w, h)
result = inpaint_region(image_path, mask)
unload_inpaint_model()
out_path = os.path.join("outputs", f"inpaint_{file_id}.jpg")
result.save(out_path, quality=95)
with open(out_path, "rb") as f:
b64 = base64.b64encode(f.read()).decode("utf-8")
return JSONResponse({"inpainted_image_base64": b64})
except Exception as e:
unload_inpaint_model()
raise HTTPException(500, f"Inpainting failed: {str(e)}")
@app.get("/projects/{file_id}/layers")
async def get_project_layers(
file_id: str,
db: Session = Depends(get_db),
current_user: User = Depends(get_current_user),
):
project = get_project_by_file_id(db, file_id, current_user.id)
if not project:
raise HTTPException(404, "Project not found")
if not os.path.exists(project.upload_path):
raise HTTPException(404, "Original file no longer exists")
try:
result = run_pipeline(project.upload_path)
return JSONResponse({
"project_id": project.id,
"file_id": file_id,
"filename": project.filename,
"image_w": result["image_w"],
"image_h": result["image_h"],
"background_base64": result["background_base64"],
"original_base64": result["original_base64"],
"layers": result["layers"],
"processing_time_s": result["processing_time_s"],
})
except Exception as e:
raise HTTPException(500, f"Pipeline failed: {str(e)}")
@app.post("/projects/{file_id}/save")
async def save_project(
file_id: str,
request: dict,
db: Session = Depends(get_db),
current_user: User = Depends(get_current_user),
):
project = get_project_by_file_id(db, file_id, current_user.id)
if not project:
raise HTTPException(404, "Project not found")
saved_state = request.get("saved_state")
thumbnail_b64 = request.get("thumbnail_b64")
layer_count = request.get("layer_count", project.layer_count)
if not saved_state:
raise HTTPException(400, "saved_state is required")
save_project_state(db, file_id, saved_state, thumbnail_b64, layer_count, current_user.id)
return JSONResponse({"status": "saved", "file_id": file_id})
@app.get("/projects/{file_id}/saved")
async def load_saved_project(
file_id: str,
db: Session = Depends(get_db),
current_user: User = Depends(get_current_user),
):
data = get_saved_state(db, file_id, current_user.id)
if not data:
raise HTTPException(404, "Project not found")
return JSONResponse(data)
@app.delete("/projects/{file_id}")
async def remove_project(
file_id: str,
db: Session = Depends(get_db),
current_user: User = Depends(get_current_user),
):
project = get_project_by_file_id(db, file_id, current_user.id)
if not project:
raise HTTPException(404, "Project not found")
if project.upload_path and os.path.exists(project.upload_path):
try:
os.remove(project.upload_path)
except:
pass
inpaint_path = os.path.join("outputs", f"inpaint_{file_id}.jpg")
if os.path.exists(inpaint_path):
try:
os.remove(inpaint_path)
except:
pass
delete_project(db, file_id, current_user.id)
return {"status": "deleted"}