Files
dev 3a345cbefd feat: AI photo enhancement (Real-ESRGAN container) + restore original
- new photo-ai service: FastAPI + Real-ESRGAN x2plus on CPU, internal only
- async job queue in server (POST enhance-ai / GET status), 5min timeout, sequential processing
- keep original photo backup (entries.photo_original_path, uploads/.originals), restore-original endpoint
- UI: AI button and restore-original button in enhance modal
- db: photo_original_path column (init.sql, migration.sql, runtime ensure)
2026-09-17 15:37:49 +03:00

74 lines
2.1 KiB
Python

import io
import os
import asyncio
import threading
import urllib.request
from fastapi import FastAPI, File, Form, UploadFile
from fastapi.responses import Response
import numpy as np
import cv2
import torch
from basicsr.archs.rrdbnet_arch import RRDBNet
from realesrgan import RealESRGANer
MODEL_PATH = os.environ.get('MODEL_PATH', '/models/RealESRGAN_x2plus.pth')
MODEL_URL = 'https://github.com/xinntao/Real-ESRGAN/releases/download/v0.2.1/RealESRGAN_x2plus.pth'
MAX_PIXELS = int(os.environ.get('MAX_INPUT_PIXELS', str(4_000_000)))
lock = threading.Lock()
upsampler = None
app = FastAPI()
def ensure_model():
if not os.path.exists(MODEL_PATH):
os.makedirs(os.path.dirname(MODEL_PATH), exist_ok=True)
tmp = MODEL_PATH + '.tmp'
urllib.request.urlretrieve(MODEL_URL, tmp)
os.replace(tmp, MODEL_PATH)
def load_model():
global upsampler
ensure_model()
model = RRDBNet(num_in_ch=3, num_out_ch=3, scale=2, num_feat=64, num_block=23, num_grow_ch=32)
upsampler = RealESRGANer(
scale=2,
model_path=MODEL_PATH,
model=model,
tile=256,
tile_pad=10,
pre_pad=0,
half=False,
device='cpu',
)
@app.on_event('startup')
async def startup():
await asyncio.to_thread(load_model)
@app.get('/health')
def health():
return {'ok': upsampler is not None}
@app.post('/enhance')
async def enhance(image: UploadFile = File(...), scale: int = Form(2)):
data = await image.read()
img = cv2.imdecode(np.frombuffer(data, np.uint8), cv2.IMREAD_COLOR)
if img is None:
return Response('bad image', status_code=400)
if img.shape[0] * img.shape[1] > MAX_PIXELS:
r = (MAX_PIXELS / (img.shape[0] * img.shape[1])) ** 0.5
img = cv2.resize(img, (int(img.shape[1] * r), int(img.shape[0] * r)), interpolation=cv2.INTER_AREA)
outscale = min(max(int(scale), 2), 4)
with lock:
out, _ = upsampler.enhance(img, outscale=outscale)
ok, enc = cv2.imencode('.jpg', out, [int(cv2.IMWRITE_JPEG_QUALITY), 92])
if not ok:
return Response('encode failed', status_code=500)
return Response(enc.tobytes(), media_type='image/jpeg')