Initial commit of video upscaler
This commit is contained in:
+684
@@ -0,0 +1,684 @@
|
||||
import os
|
||||
import uuid
|
||||
import queue
|
||||
import threading
|
||||
import asyncio
|
||||
import subprocess
|
||||
import urllib.request
|
||||
import json
|
||||
from typing import Dict, List, Any
|
||||
from fastapi import FastAPI, UploadFile, File, Form, BackgroundTasks, HTTPException, WebSocket, WebSocketDisconnect
|
||||
from fastapi.responses import FileResponse, JSONResponse
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from pydantic import BaseModel
|
||||
|
||||
import app.upscaler as upscaler
|
||||
|
||||
app = FastAPI(title="AI Video Upscaler Dashboard")
|
||||
|
||||
# CORS middleware for local development
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
# In-memory databases
|
||||
jobs_db: Dict[str, upscaler.UpscaleJob] = {}
|
||||
ws_connections: Dict[str, List[WebSocket]] = {}
|
||||
preview_db: Dict[str, Dict[str, str]] = {} # preview_id -> {orig, upscaled}
|
||||
|
||||
# FIFO queue for upscaling jobs to prevent GPU memory overload
|
||||
job_queue = queue.Queue()
|
||||
queue_lock = threading.Lock()
|
||||
current_running_job_id = None
|
||||
|
||||
main_loop = None
|
||||
|
||||
@app.on_event("startup")
|
||||
def startup_event():
|
||||
global main_loop
|
||||
main_loop = asyncio.get_event_loop()
|
||||
|
||||
global_webhook_url = None
|
||||
|
||||
class WebhookSetupRequest(BaseModel):
|
||||
url: str
|
||||
|
||||
def send_webhook_notification(url: str, payload: dict):
|
||||
def post():
|
||||
try:
|
||||
req = urllib.request.Request(
|
||||
url,
|
||||
data=json.dumps(payload).encode('utf-8'),
|
||||
headers={'Content-Type': 'application/json'}
|
||||
)
|
||||
with urllib.request.urlopen(req, timeout=5) as response:
|
||||
response.read()
|
||||
except Exception as e:
|
||||
print(f"Webhook notification failed to {url}: {e}")
|
||||
threading.Thread(target=post, daemon=True).start()
|
||||
|
||||
# Broadcast updates to websockets and webhooks
|
||||
def broadcast_progress(job_id: str, data: dict):
|
||||
job = jobs_db.get(job_id)
|
||||
if job:
|
||||
data["is_preview"] = getattr(job, "is_preview", False)
|
||||
if getattr(job, "is_preview", False):
|
||||
data["original_preview_file"] = f"original_{job_id}.{getattr(job, 'transcode_format', 'mp4')}"
|
||||
|
||||
if job_id in ws_connections and main_loop is not None:
|
||||
active_sockets = list(ws_connections[job_id])
|
||||
for websocket in active_sockets:
|
||||
try:
|
||||
# Schedule the coroutine in the main event loop
|
||||
asyncio.run_coroutine_threadsafe(websocket.send_json(data), main_loop)
|
||||
except Exception as e:
|
||||
print(f"Error sending websocket message: {e}")
|
||||
|
||||
# Process webhooks
|
||||
job = jobs_db.get(job_id)
|
||||
webhook_url = None
|
||||
if job and getattr(job, "webhook_url", None):
|
||||
webhook_url = job.webhook_url
|
||||
elif global_webhook_url:
|
||||
webhook_url = global_webhook_url
|
||||
|
||||
if webhook_url:
|
||||
payload = {
|
||||
"job_id": job_id,
|
||||
"status": data.get("status"),
|
||||
"progress": data.get("progress"),
|
||||
"current_frame": data.get("current_frame"),
|
||||
"total_frames": data.get("total_frames"),
|
||||
"eta": data.get("eta"),
|
||||
"error": data.get("error"),
|
||||
"output_file": data.get("output_file")
|
||||
}
|
||||
send_webhook_notification(webhook_url, payload)
|
||||
|
||||
# Sequential queue processor thread
|
||||
def queue_worker():
|
||||
global current_running_job_id
|
||||
while True:
|
||||
try:
|
||||
job_id = job_queue.get()
|
||||
if job_id is None:
|
||||
break
|
||||
|
||||
job = jobs_db.get(job_id)
|
||||
if not job or job.status == "cancelled":
|
||||
job_queue.task_done()
|
||||
continue
|
||||
|
||||
with queue_lock:
|
||||
current_running_job_id = job_id
|
||||
|
||||
print(f"Starting job {job_id}...")
|
||||
upscaler.run_upscale_pipeline(job, broadcast_progress)
|
||||
|
||||
with queue_lock:
|
||||
current_running_job_id = None
|
||||
|
||||
job_queue.task_done()
|
||||
except Exception as e:
|
||||
print(f"Error in queue worker: {e}")
|
||||
|
||||
# Start the background queue worker thread
|
||||
worker_thread = threading.Thread(target=queue_worker, daemon=True)
|
||||
worker_thread.start()
|
||||
|
||||
|
||||
class StartUpscaleRequest(BaseModel):
|
||||
file_id: str
|
||||
model: str = "realesr-animevideov3"
|
||||
scale: int = 4
|
||||
tile_size: int = 256
|
||||
preserve_audio: bool = True
|
||||
ss: str | None = None
|
||||
t: str | None = None
|
||||
gpu_ids: str | None = None
|
||||
tta: bool = False
|
||||
unsharp: bool = False
|
||||
double_fps: bool = False
|
||||
preserve_subtitles: bool = True
|
||||
start_sec: float | None = None
|
||||
end_sec: float | None = None
|
||||
crf: int = 18
|
||||
preset: str = "medium"
|
||||
denoise: bool = False
|
||||
sharpen: bool = False
|
||||
interpolation: bool = False
|
||||
webhook_url: str | None = None
|
||||
transcode_format: str = "mp4"
|
||||
is_preview: bool = False
|
||||
|
||||
class PreviewRequest(BaseModel):
|
||||
file_id: str
|
||||
timestamp_sec: float = 1.0
|
||||
model: str = "realesr-animevideov3"
|
||||
scale: int = 4
|
||||
tile_size: int = 256
|
||||
gpu_ids: str | None = None
|
||||
|
||||
|
||||
# Endpoints
|
||||
@app.post("/api/upload")
|
||||
async def upload_video(file: UploadFile = File(...)):
|
||||
"""Upload video to workspace directory and parse metadata"""
|
||||
file_id = str(uuid.uuid4())
|
||||
ext = os.path.splitext(file.filename)[1].lower()
|
||||
if ext not in [".mp4", ".mkv", ".avi", ".mov", ".webm"]:
|
||||
raise HTTPException(status_code=400, detail="Invalid video format. Supported: .mp4, .mkv, .avi, .mov, .webm")
|
||||
|
||||
save_path = os.path.join(upscaler.UPLOAD_DIR, f"{file_id}{ext}")
|
||||
|
||||
with open(save_path, "wb") as buffer:
|
||||
content = await file.read()
|
||||
buffer.write(content)
|
||||
|
||||
info = upscaler.get_video_info(save_path)
|
||||
if not info:
|
||||
# cleanup if reading failed
|
||||
os.remove(save_path)
|
||||
raise HTTPException(status_code=400, detail="Could not read video file metadata. File may be corrupted.")
|
||||
|
||||
return {
|
||||
"file_id": file_id,
|
||||
"filename": file.filename,
|
||||
"path": save_path,
|
||||
"metadata": info
|
||||
}
|
||||
|
||||
@app.post("/api/upscale/start")
|
||||
def start_upscale(req: StartUpscaleRequest):
|
||||
"""Queue upscaling task"""
|
||||
# Find uploaded file
|
||||
file_path = None
|
||||
exts = [".mp4", ".mkv", ".avi", ".mov", ".webm"]
|
||||
for ext in exts:
|
||||
test_path = os.path.join(upscaler.UPLOAD_DIR, f"{req.file_id}{ext}")
|
||||
if os.path.exists(test_path):
|
||||
file_path = test_path
|
||||
break
|
||||
|
||||
if not file_path:
|
||||
raise HTTPException(status_code=404, detail="Uploaded file not found.")
|
||||
|
||||
job_id = str(uuid.uuid4())
|
||||
job = upscaler.UpscaleJob(
|
||||
job_id=job_id,
|
||||
video_path=file_path,
|
||||
model=req.model,
|
||||
scale=req.scale,
|
||||
tile_size=req.tile_size,
|
||||
preserve_audio=req.preserve_audio,
|
||||
ss=req.ss,
|
||||
t=req.t,
|
||||
gpu_ids=req.gpu_ids,
|
||||
tta=req.tta,
|
||||
unsharp=req.unsharp,
|
||||
double_fps=req.double_fps,
|
||||
preserve_subtitles=req.preserve_subtitles,
|
||||
start_sec=req.start_sec,
|
||||
end_sec=req.end_sec,
|
||||
crf=req.crf,
|
||||
preset=req.preset,
|
||||
denoise=req.denoise,
|
||||
sharpen=req.sharpen,
|
||||
interpolation=req.interpolation,
|
||||
webhook_url=req.webhook_url,
|
||||
transcode_format=req.transcode_format,
|
||||
is_preview=req.is_preview
|
||||
)
|
||||
|
||||
jobs_db[job_id] = job
|
||||
job_queue.put(job_id)
|
||||
|
||||
# Broadcast initial queued progress
|
||||
broadcast_progress(job_id, {
|
||||
"status": "queued",
|
||||
"progress": 0.0,
|
||||
"current_frame": 0,
|
||||
"total_frames": 0,
|
||||
"eta": "Calculating..."
|
||||
})
|
||||
|
||||
return {
|
||||
"job_id": job_id,
|
||||
"status": "queued"
|
||||
}
|
||||
|
||||
@app.get("/api/upscale/status/{job_id}")
|
||||
def get_status(job_id: str):
|
||||
"""Get status of upscale job"""
|
||||
job = jobs_db.get(job_id)
|
||||
if not job:
|
||||
raise HTTPException(status_code=404, detail="Job not found.")
|
||||
|
||||
return {
|
||||
"job_id": job.job_id,
|
||||
"status": job.status,
|
||||
"progress": job.progress,
|
||||
"current_frame": job.current_frame,
|
||||
"total_frames": job.total_frames,
|
||||
"eta": job.eta,
|
||||
"error": job.error,
|
||||
"output_file": os.path.basename(job.output_file) if job.output_file else None,
|
||||
"is_preview": getattr(job, "is_preview", False),
|
||||
"original_preview_file": f"original_{job_id}.{getattr(job, 'transcode_format', 'mp4')}" if getattr(job, "is_preview", False) else None
|
||||
}
|
||||
|
||||
@app.post("/api/upscale/cancel/{job_id}")
|
||||
def cancel_job(job_id: str):
|
||||
"""Cancel a running or queued job"""
|
||||
job = jobs_db.get(job_id)
|
||||
if not job:
|
||||
raise HTTPException(status_code=404, detail="Job not found.")
|
||||
|
||||
job.cancel()
|
||||
# Broadcast cancellation status
|
||||
broadcast_progress(job_id, {
|
||||
"status": "cancelled",
|
||||
"progress": job.progress,
|
||||
"current_frame": job.current_frame,
|
||||
"total_frames": job.total_frames,
|
||||
"eta": "N/A"
|
||||
})
|
||||
return {"job_id": job_id, "status": "cancelled"}
|
||||
|
||||
@app.post("/api/preview/generate")
|
||||
async def generate_preview(req: PreviewRequest):
|
||||
"""Generate high quality upscaled single frame preview"""
|
||||
file_path = None
|
||||
exts = [".mp4", ".mkv", ".avi", ".mov", ".webm"]
|
||||
for ext in exts:
|
||||
test_path = os.path.join(upscaler.UPLOAD_DIR, f"{req.file_id}{ext}")
|
||||
if os.path.exists(test_path):
|
||||
file_path = test_path
|
||||
break
|
||||
|
||||
if not file_path:
|
||||
raise HTTPException(status_code=404, detail="Video file not found.")
|
||||
|
||||
preview_id = str(uuid.uuid4())
|
||||
preview_temp_dir = os.path.join(upscaler.TEMP_DIR, f"preview_{preview_id}")
|
||||
os.makedirs(preview_temp_dir, exist_ok=True)
|
||||
|
||||
orig_path = os.path.join(preview_temp_dir, "orig.jpg")
|
||||
upscaled_path = os.path.join(preview_temp_dir, "upscaled.jpg")
|
||||
|
||||
# Extract frame
|
||||
extracted = upscaler.extract_single_frame(file_path, req.timestamp_sec, orig_path)
|
||||
if not extracted:
|
||||
raise HTTPException(status_code=500, detail="Failed to extract preview frame from video.")
|
||||
|
||||
# Upscale frame
|
||||
upscaled = upscaler.upscale_image_file(orig_path, upscaled_path, req.model, req.scale, req.tile_size, req.gpu_ids)
|
||||
if not upscaled:
|
||||
raise HTTPException(status_code=500, detail="Failed to upscale preview frame.")
|
||||
|
||||
preview_db[preview_id] = {
|
||||
"orig": orig_path,
|
||||
"upscaled": upscaled_path
|
||||
}
|
||||
|
||||
return {"preview_id": preview_id}
|
||||
|
||||
@app.get("/api/preview/original/{preview_id}")
|
||||
def get_preview_original(preview_id: str):
|
||||
paths = preview_db.get(preview_id)
|
||||
if not paths or not os.path.exists(paths["orig"]):
|
||||
raise HTTPException(status_code=404, detail="Preview frame not found.")
|
||||
return FileResponse(paths["orig"])
|
||||
|
||||
@app.get("/api/preview/upscaled/{preview_id}")
|
||||
def get_preview_upscaled(preview_id: str):
|
||||
paths = preview_db.get(preview_id)
|
||||
if not paths or not os.path.exists(paths["upscaled"]):
|
||||
raise HTTPException(status_code=404, detail="Preview frame not found.")
|
||||
return FileResponse(paths["upscaled"])
|
||||
|
||||
@app.get("/api/models")
|
||||
def get_models():
|
||||
"""Scan models directory and return list of available models"""
|
||||
models_dir = os.path.join(upscaler.BASE_DIR, "realesrgan-bin", "models")
|
||||
if not os.path.exists(models_dir):
|
||||
return []
|
||||
|
||||
models = set()
|
||||
for filename in os.listdir(models_dir):
|
||||
if filename.endswith(".param"):
|
||||
name = filename[:-6] # strip '.param'
|
||||
# Strip scale suffix if present
|
||||
for suffix in ["-x2", "-x3", "-x4"]:
|
||||
if name.endswith(suffix):
|
||||
name = name[:-len(suffix)]
|
||||
break
|
||||
models.add(name)
|
||||
|
||||
return sorted(list(models))
|
||||
|
||||
@app.get("/api/download/{filename}")
|
||||
def download_file(filename: str):
|
||||
file_path = os.path.join(upscaler.OUTPUT_DIR, filename)
|
||||
if not os.path.exists(file_path):
|
||||
raise HTTPException(status_code=404, detail="File not found.")
|
||||
return FileResponse(file_path, media_type="application/octet-stream", filename=filename)
|
||||
|
||||
def safe_delete_file(file_path: str):
|
||||
if file_path and os.path.exists(file_path):
|
||||
try:
|
||||
os.remove(file_path)
|
||||
print(f"Deleted file: {file_path}")
|
||||
except Exception as e:
|
||||
print(f"Failed to delete {file_path}: {e}")
|
||||
|
||||
@app.get("/api/jobs")
|
||||
def list_jobs():
|
||||
"""List details of all submitted jobs"""
|
||||
return [
|
||||
{
|
||||
"job_id": job.job_id,
|
||||
"status": job.status,
|
||||
"progress": job.progress,
|
||||
"current_frame": job.current_frame,
|
||||
"total_frames": job.total_frames,
|
||||
"eta": job.eta,
|
||||
"error": job.error,
|
||||
"model": job.model,
|
||||
"scale": job.scale,
|
||||
"output_file": os.path.basename(job.output_file) if job.output_file else None,
|
||||
"video_path": job.video_path,
|
||||
"is_preview": getattr(job, "is_preview", False)
|
||||
}
|
||||
for job in jobs_db.values()
|
||||
]
|
||||
|
||||
@app.delete("/api/jobs/{job_id}")
|
||||
@app.post("/api/jobs/delete/{job_id}")
|
||||
def delete_job(job_id: str):
|
||||
job = jobs_db.get(job_id)
|
||||
if not job:
|
||||
raise HTTPException(status_code=404, detail="Job not found.")
|
||||
|
||||
# Safely delete uploaded input file
|
||||
if job.video_path:
|
||||
safe_delete_file(job.video_path)
|
||||
|
||||
# Safely delete original preview video if present
|
||||
for ext in [".mp4", ".mkv", ".avi", ".mov", ".webm"]:
|
||||
orig_prev_path = os.path.join(upscaler.OUTPUT_DIR, f"original_{job_id}{ext}")
|
||||
if os.path.exists(orig_prev_path):
|
||||
safe_delete_file(orig_prev_path)
|
||||
|
||||
# Safely delete upscaled output file
|
||||
if job.output_file:
|
||||
safe_delete_file(job.output_file)
|
||||
|
||||
# Clean up any residual temp folder
|
||||
job_temp_dir = os.path.join(upscaler.TEMP_DIR, job_id)
|
||||
if os.path.exists(job_temp_dir):
|
||||
try:
|
||||
import shutil
|
||||
shutil.rmtree(job_temp_dir)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Delete from in-memory db
|
||||
if job_id in jobs_db:
|
||||
del jobs_db[job_id]
|
||||
|
||||
return {"job_id": job_id, "status": "purged"}
|
||||
|
||||
@app.post("/api/jobs/purge-all")
|
||||
def purge_all_jobs():
|
||||
# Purge uploads folder
|
||||
for f in os.listdir(upscaler.UPLOAD_DIR):
|
||||
safe_delete_file(os.path.join(upscaler.UPLOAD_DIR, f))
|
||||
|
||||
# Purge outputs folder
|
||||
for f in os.listdir(upscaler.OUTPUT_DIR):
|
||||
safe_delete_file(os.path.join(upscaler.OUTPUT_DIR, f))
|
||||
|
||||
# Purge temp folder
|
||||
for f in os.listdir(upscaler.TEMP_DIR):
|
||||
path = os.path.join(upscaler.TEMP_DIR, f)
|
||||
try:
|
||||
if os.path.isdir(path):
|
||||
import shutil
|
||||
shutil.rmtree(path)
|
||||
else:
|
||||
os.remove(path)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Reset in-memory database
|
||||
jobs_db.clear()
|
||||
|
||||
return {"status": "all purged"}
|
||||
|
||||
# Websocket endpoint for real-time progress updates
|
||||
@app.websocket("/ws/progress/{job_id}")
|
||||
async def websocket_progress(websocket: WebSocket, job_id: str):
|
||||
await websocket.accept()
|
||||
if job_id not in ws_connections:
|
||||
ws_connections[job_id] = []
|
||||
ws_connections[job_id].append(websocket)
|
||||
|
||||
# Send current state immediately
|
||||
job = jobs_db.get(job_id)
|
||||
if job:
|
||||
await websocket.send_json({
|
||||
"status": job.status,
|
||||
"progress": job.progress,
|
||||
"current_frame": job.current_frame,
|
||||
"total_frames": job.total_frames,
|
||||
"eta": job.eta,
|
||||
"error": job.error,
|
||||
"output_file": os.path.basename(job.output_file) if job.output_file else None
|
||||
})
|
||||
|
||||
try:
|
||||
while True:
|
||||
# Just keep the connection alive
|
||||
await websocket.receive_text()
|
||||
except WebSocketDisconnect:
|
||||
if job_id in ws_connections:
|
||||
ws_connections[job_id].remove(websocket)
|
||||
if not ws_connections[job_id]:
|
||||
del ws_connections[job_id]
|
||||
|
||||
# Clean up previews on exit
|
||||
@app.on_event("shutdown")
|
||||
def cleanup_previews():
|
||||
for preview_id, paths in preview_db.items():
|
||||
try:
|
||||
parent_dir = os.path.dirname(paths["orig"])
|
||||
if os.path.exists(parent_dir):
|
||||
import shutil
|
||||
shutil.rmtree(parent_dir)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# System Diagnostics Endpoint
|
||||
@app.get("/api/diagnostics")
|
||||
def get_diagnostics():
|
||||
# RAM Diagnostics
|
||||
ram = {"total": 0, "available": 0, "used": 0, "percent": 0.0}
|
||||
try:
|
||||
if os.path.exists('/proc/meminfo'):
|
||||
mem_info = {}
|
||||
with open('/proc/meminfo', 'r') as f:
|
||||
for line in f:
|
||||
parts = line.split(':')
|
||||
if len(parts) == 2:
|
||||
name = parts[0].strip()
|
||||
val = parts[1].replace('kB', '').strip()
|
||||
mem_info[name] = int(val)
|
||||
total = mem_info.get('MemTotal', 0) * 1024
|
||||
free = mem_info.get('MemFree', 0) * 1024
|
||||
available = mem_info.get('MemAvailable', total - free) * 1024
|
||||
used = total - available
|
||||
percent = round((used / total) * 100, 1) if total > 0 else 0.0
|
||||
ram = {
|
||||
"total": total,
|
||||
"available": available,
|
||||
"used": used,
|
||||
"percent": percent
|
||||
}
|
||||
except Exception as e:
|
||||
print(f"Error getting RAM diagnostics: {e}")
|
||||
|
||||
# CPU Diagnostics
|
||||
cpu = {"count": os.cpu_count(), "load_avg": [], "percent": 0.0}
|
||||
try:
|
||||
load_1, load_5, load_15 = os.getloadavg()
|
||||
cpu["load_avg"] = [load_1, load_5, load_15]
|
||||
cpu["percent"] = min(100.0, round((load_1 / cpu["count"]) * 100, 1)) if cpu["count"] else 0.0
|
||||
except Exception as e:
|
||||
print(f"Error getting CPU diagnostics: {e}")
|
||||
|
||||
# GPU Diagnostics
|
||||
gpu = {"available": False, "gpus": []}
|
||||
try:
|
||||
cmd = [
|
||||
"nvidia-smi",
|
||||
"--query-gpu=utilization.gpu,utilization.memory,memory.total,memory.free,memory.used,name,temperature.gpu",
|
||||
"--format=csv,noheader,nounits"
|
||||
]
|
||||
res = subprocess.run(cmd, capture_output=True, text=True, check=True)
|
||||
lines = res.stdout.strip().split('\n')
|
||||
gpu_list = []
|
||||
for line in lines:
|
||||
if not line.strip():
|
||||
continue
|
||||
parts = [p.strip() for p in line.split(',')]
|
||||
if len(parts) >= 7:
|
||||
gpu_list.append({
|
||||
"name": parts[5],
|
||||
"gpu_utilization_percent": float(parts[0]),
|
||||
"memory_utilization_percent": float(parts[1]),
|
||||
"memory_total_mb": float(parts[2]),
|
||||
"memory_free_mb": float(parts[3]),
|
||||
"memory_used_mb": float(parts[4]),
|
||||
"temperature_c": float(parts[6])
|
||||
})
|
||||
if gpu_list:
|
||||
gpu["available"] = True
|
||||
gpu["gpus"] = gpu_list
|
||||
except Exception as e:
|
||||
gpu["error"] = str(e)
|
||||
|
||||
return {
|
||||
"ram": ram,
|
||||
"cpu": cpu,
|
||||
"gpu": gpu
|
||||
}
|
||||
|
||||
# Outputs Gallery Endpoints
|
||||
@app.get("/api/outputs")
|
||||
def list_outputs():
|
||||
output_dir = upscaler.OUTPUT_DIR
|
||||
if not os.path.exists(output_dir):
|
||||
return []
|
||||
|
||||
files = []
|
||||
for filename in os.listdir(output_dir):
|
||||
if filename.startswith("original_"):
|
||||
continue
|
||||
file_path = os.path.join(output_dir, filename)
|
||||
if os.path.isfile(file_path):
|
||||
stat = os.stat(file_path)
|
||||
files.append({
|
||||
"filename": filename,
|
||||
"size": stat.st_size,
|
||||
"modified": stat.st_mtime,
|
||||
"url": f"/api/download/{filename}"
|
||||
})
|
||||
files.sort(key=lambda x: x["modified"], reverse=True)
|
||||
return files
|
||||
|
||||
@app.delete("/api/outputs/delete/{filename}")
|
||||
@app.post("/api/outputs/delete/{filename}")
|
||||
def delete_output(filename: str):
|
||||
file_path = os.path.join(upscaler.OUTPUT_DIR, filename)
|
||||
resolved_path = os.path.abspath(file_path)
|
||||
if not resolved_path.startswith(os.path.abspath(upscaler.OUTPUT_DIR)):
|
||||
raise HTTPException(status_code=403, detail="Access denied.")
|
||||
|
||||
if not os.path.exists(file_path):
|
||||
raise HTTPException(status_code=404, detail="File not found.")
|
||||
|
||||
try:
|
||||
os.remove(file_path)
|
||||
return {"filename": filename, "status": "deleted"}
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=500, detail=f"Failed to delete file: {e}")
|
||||
|
||||
# Custom Model Uploading Endpoint
|
||||
@app.post("/api/models/upload")
|
||||
async def upload_model(
|
||||
param_file: UploadFile = File(None),
|
||||
bin_file: UploadFile = File(None),
|
||||
files: List[UploadFile] = File(None)
|
||||
):
|
||||
models_dir = os.path.join(upscaler.BASE_DIR, "realesrgan-bin", "models")
|
||||
os.makedirs(models_dir, exist_ok=True)
|
||||
|
||||
uploaded_files = []
|
||||
all_files = []
|
||||
if files:
|
||||
all_files.extend(files)
|
||||
if param_file:
|
||||
all_files.append(param_file)
|
||||
if bin_file:
|
||||
all_files.append(bin_file)
|
||||
|
||||
if not all_files:
|
||||
raise HTTPException(status_code=400, detail="No files uploaded. Please upload .param or .bin files.")
|
||||
|
||||
for f in all_files:
|
||||
filename = f.filename
|
||||
ext = os.path.splitext(filename)[1].lower()
|
||||
if ext not in [".param", ".bin"]:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Invalid file extension: {ext}. Only .param and .bin files are allowed for models."
|
||||
)
|
||||
|
||||
safe_filename = os.path.basename(filename)
|
||||
save_path = os.path.join(models_dir, safe_filename)
|
||||
|
||||
with open(save_path, "wb") as buffer:
|
||||
content = await f.read()
|
||||
buffer.write(content)
|
||||
uploaded_files.append(safe_filename)
|
||||
|
||||
return {
|
||||
"status": "success",
|
||||
"uploaded": uploaded_files
|
||||
}
|
||||
|
||||
# Webhook Setup/Retrieval Endpoints
|
||||
@app.post("/api/webhook/setup")
|
||||
def setup_webhook(req: WebhookSetupRequest):
|
||||
global global_webhook_url
|
||||
global_webhook_url = req.url
|
||||
return {"status": "success", "webhook_url": global_webhook_url}
|
||||
|
||||
@app.get("/api/uploads/{filename}")
|
||||
def get_upload_file(filename: str):
|
||||
file_path = os.path.join(upscaler.UPLOAD_DIR, filename)
|
||||
resolved_path = os.path.abspath(file_path)
|
||||
if not resolved_path.startswith(os.path.abspath(upscaler.UPLOAD_DIR)):
|
||||
raise HTTPException(status_code=403, detail="Access denied.")
|
||||
if not os.path.exists(file_path):
|
||||
raise HTTPException(status_code=404, detail="File not found.")
|
||||
return FileResponse(file_path)
|
||||
|
||||
# Mount static folder
|
||||
app.mount("/", StaticFiles(directory=os.path.join(upscaler.BASE_DIR, "static"), html=True), name="static")
|
||||
Reference in New Issue
Block a user