Files
video_upscaler/app/main.py
T

2086 lines
72 KiB
Python

import os
import uuid
import queue
import threading
import asyncio
import subprocess
import time
import urllib.request
import urllib.parse
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")
# Middleware to prevent browser caching of static files (dev mode)
from starlette.middleware.base import BaseHTTPMiddleware
from starlette.requests import Request
class NoCacheStaticMiddleware(BaseHTTPMiddleware):
async def dispatch(self, request: Request, call_next):
response = await call_next(request)
path = request.url.path
if path.endswith(('.js', '.css', '.html')) or path == '/':
response.headers["Cache-Control"] = "no-cache, no-store, must-revalidate"
response.headers["Pragma"] = "no-cache"
response.headers["Expires"] = "0"
return response
app.add_middleware(NoCacheStaticMiddleware)
# CORS middleware for local development
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# In-memory databases
class CustomJobQueue:
def __init__(self):
self.queue = []
self.lock = threading.Lock()
self.condition = threading.Condition(self.lock)
def put(self, job_id: str):
with self.lock:
if job_id not in self.queue:
self.queue.append(job_id)
self.condition.notify()
def get(self) -> str:
with self.lock:
while not self.queue:
self.condition.wait()
return self.queue.pop(0)
def remove(self, job_id: str) -> bool:
with self.lock:
if job_id in self.queue:
self.queue.remove(job_id)
return True
return False
def get_all(self) -> List[str]:
with self.lock:
return list(self.queue)
def reorder(self, job_ids: List[str]):
with self.lock:
valid_ids = [jid for jid in job_ids if jid in self.queue]
missing_ids = [jid for jid in self.queue if jid not in valid_ids]
self.queue = valid_ids + missing_ids
def task_done(self):
pass
def empty(self) -> bool:
with self.lock:
return len(self.queue) == 0
def qsize(self) -> int:
with self.lock:
return len(self.queue)
jobs_db: Dict[str, upscaler.UpscaleJob] = {}
ws_connections: Dict[str, List[WebSocket]] = {}
preview_db: Dict[str, Dict[str, str]] = {} # preview_id -> {orig, upscaled}
# Custom thread-safe queue for upscaling jobs to support reordering & cancellation
job_queue = CustomJobQueue()
queue_lock = threading.Lock()
current_running_job_id = None
main_loop = None
JOBS_FILE = os.path.join(upscaler.BASE_DIR, "jobs.json")
SETTINGS_FILE = os.path.join(upscaler.BASE_DIR, "settings.json")
def load_settings() -> dict:
"""Load persistent app settings from settings.json"""
defaults = {
"default_temp_dir": "",
"default_output_dir": "",
"role": "standalone",
"coordinator_url": "",
"worker_node_url": "http://localhost:8000",
"use_shared_storage": False,
"shared_storage_path": "",
"smb_shares": []
}
if os.path.exists(SETTINGS_FILE):
try:
with open(SETTINGS_FILE, "r") as f:
saved = json.load(f)
defaults.update(saved)
except Exception as e:
print(f"Error loading settings: {e}")
return defaults
def save_settings(settings: dict):
"""Save persistent app settings to settings.json"""
try:
with open(SETTINGS_FILE, "w") as f:
json.dump(settings, f, indent=4)
except Exception as e:
print(f"Error saving settings: {e}")
def load_jobs_db():
global jobs_db
if os.path.exists(JOBS_FILE):
try:
with open(JOBS_FILE, "r") as f:
data = json.load(f)
for job_id, job_data in data.items():
job = upscaler.UpscaleJob.from_dict(job_data)
# Automatically put queued items back in the queue
if job.status == "queued":
job_queue.put(job_id)
# Mark active items as interrupted so they can be resumed
elif job.status in ["analyzing", "extracting", "upscaling", "restoring_faces", "interpolating", "assembling"]:
job.status = "interrupted"
job.eta = "Interrupted"
jobs_db[job_id] = job
except Exception as e:
print(f"Error loading jobs database: {e}")
def save_jobs_db():
try:
with open(JOBS_FILE, "w") as f:
data = {job_id: job.to_dict() for job_id, job in jobs_db.items()}
json.dump(data, f, indent=4)
except Exception as e:
print(f"Error saving jobs database: {e}")
@app.on_event("startup")
def startup_event():
global main_loop
main_loop = asyncio.get_event_loop()
load_jobs_db()
restart_dist_worker()
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):
save_jobs_db()
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 | None = None
server_file_path: str | None = None
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
ai_face_restoration: bool = False
ai_rife_interpolation: bool = False
ai_audio_denoise: bool = False
temp_dir: str | None = None
output_dir: str | None = None
class StartBatchUpscaleRequest(BaseModel):
file_ids: List[str] | None = None
server_file_paths: List[str] | None = None
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
ai_face_restoration: bool = False
ai_rife_interpolation: bool = False
ai_audio_denoise: bool = False
temp_dir: str | None = None
output_dir: str | None = None
class PreviewRequest(BaseModel):
file_id: str | None = None
server_file_path: str | None = None
timestamp_sec: float = 1.0
model: str = "realesr-animevideov3"
scale: int = 4
tile_size: int = 256
gpu_ids: str | None = None
# Endpoints
UPLOAD_METADATA_FILE = os.path.join(upscaler.UPLOAD_DIR, "metadata.json")
def load_upload_metadata():
if os.path.exists(UPLOAD_METADATA_FILE):
try:
with open(UPLOAD_METADATA_FILE, "r") as f:
return json.load(f)
except Exception:
pass
return {}
def save_upload_metadata(metadata):
try:
with open(UPLOAD_METADATA_FILE, "w") as f:
json.dump(metadata, f, indent=4)
except Exception:
pass
@app.post("/api/upload")
async def upload_video(file: UploadFile = File(...)):
"""Upload video to workspace directory and parse metadata"""
import time
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.")
# Save upload metadata
metadata = load_upload_metadata()
metadata[file_id] = {
"file_id": file_id,
"original_filename": file.filename,
"ext": ext,
"size_bytes": os.path.getsize(save_path),
"upload_time": time.time(),
"metadata": info
}
save_upload_metadata(metadata)
return {
"file_id": file_id,
"filename": file.filename,
"path": save_path,
"metadata": info
}
@app.get("/api/uploads")
def list_uploads():
"""List all uploaded video source files"""
metadata = load_upload_metadata()
valid_uploads = []
metadata_updated = False
if os.path.exists(upscaler.UPLOAD_DIR):
all_files = os.listdir(upscaler.UPLOAD_DIR)
for filename in all_files:
if filename == "metadata.json":
continue
file_path = os.path.join(upscaler.UPLOAD_DIR, filename)
file_id, ext = os.path.splitext(filename)
if file_id in metadata:
valid_uploads.append(metadata[file_id])
else:
info = upscaler.get_video_info(file_path)
if info:
entry = {
"file_id": file_id,
"original_filename": filename,
"ext": ext,
"size_bytes": os.path.getsize(file_path),
"upload_time": os.path.getmtime(file_path),
"metadata": info
}
metadata[file_id] = entry
valid_uploads.append(entry)
metadata_updated = True
# Clean up missing files from metadata
for fid in list(metadata.keys()):
ext = metadata[fid].get("ext", ".mp4")
expected_file = os.path.join(upscaler.UPLOAD_DIR, f"{fid}{ext}")
if not os.path.exists(expected_file):
del metadata[fid]
metadata_updated = True
if metadata_updated:
save_upload_metadata(metadata)
# Sort by upload time desc
valid_uploads.sort(key=lambda x: x.get("upload_time", 0), reverse=True)
return valid_uploads
@app.delete("/api/uploads/{file_id}")
@app.post("/api/uploads/delete/{file_id}")
def delete_upload(file_id: str):
"""Delete an uploaded source file"""
metadata = load_upload_metadata()
if file_id not in metadata:
exts = [".mp4", ".mkv", ".avi", ".mov", ".webm"]
file_path = None
for ext in exts:
p = os.path.join(upscaler.UPLOAD_DIR, f"{file_id}{ext}")
if os.path.exists(p):
file_path = p
break
if not file_path:
raise HTTPException(status_code=404, detail="Upload file not found.")
os.remove(file_path)
return {"file_id": file_id, "status": "deleted"}
ext = metadata[file_id].get("ext", ".mp4")
file_path = os.path.join(upscaler.UPLOAD_DIR, f"{file_id}{ext}")
if os.path.exists(file_path):
os.remove(file_path)
del metadata[file_id]
save_upload_metadata(metadata)
return {"file_id": file_id, "status": "deleted"}
@app.post("/api/upscale/start")
def start_upscale(req: StartUpscaleRequest):
"""Queue upscaling task"""
# Find uploaded file
file_path = None
if req.server_file_path:
if os.path.exists(req.server_file_path) and os.path.isfile(req.server_file_path):
file_path = req.server_file_path
else:
raise HTTPException(status_code=404, detail="Server file path not found.")
else:
if not req.file_id:
raise HTTPException(status_code=400, detail="Either file_id or server_file_path must be provided.")
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.")
metadata = load_upload_metadata()
source_filename = None
if req.file_id and req.file_id in metadata:
source_filename = metadata[req.file_id].get("original_filename")
if not source_filename and file_path:
source_filename = os.path.basename(file_path)
settings = load_settings()
output_dir = req.output_dir or settings.get("default_output_dir") or None
if output_dir:
output_dir = output_dir.strip()
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,
ai_face_restoration=req.ai_face_restoration,
ai_rife_interpolation=req.ai_rife_interpolation,
ai_audio_denoise=req.ai_audio_denoise,
temp_dir=req.temp_dir or settings.get("default_temp_dir") or None,
output_dir=output_dir,
source_filename=source_filename
)
jobs_db[job_id] = job
job_queue.put(job_id)
save_jobs_db()
# 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.post("/api/upscale/start/batch")
def start_upscale_batch(req: StartBatchUpscaleRequest):
"""Queue multiple upscaling tasks with one set of settings"""
queued_jobs = []
file_sources = [] # list of (file_id, file_path)
if req.server_file_paths:
for path in req.server_file_paths:
if os.path.exists(path) and os.path.isfile(path):
file_sources.append((None, path))
else:
raise HTTPException(status_code=404, detail=f"Server file {path} not found.")
else:
if not req.file_ids:
raise HTTPException(status_code=400, detail="Either file_ids or server_file_paths must be provided.")
for file_id in req.file_ids:
file_path = None
exts = [".mp4", ".mkv", ".avi", ".mov", ".webm"]
for ext in exts:
test_path = os.path.join(upscaler.UPLOAD_DIR, f"{file_id}{ext}")
if os.path.exists(test_path):
file_path = test_path
break
if not file_path:
raise HTTPException(status_code=404, detail=f"Uploaded file {file_id} not found.")
file_sources.append((file_id, file_path))
settings = load_settings()
output_dir = req.output_dir or settings.get("default_output_dir") or None
if output_dir:
output_dir = output_dir.strip()
# 2. Queueing loop
for file_id, file_path in file_sources:
metadata = load_upload_metadata()
source_filename = None
if file_id and file_id in metadata:
source_filename = metadata[file_id].get("original_filename")
if not source_filename and file_path:
source_filename = os.path.basename(file_path)
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,
ai_face_restoration=req.ai_face_restoration,
ai_rife_interpolation=req.ai_rife_interpolation,
ai_audio_denoise=req.ai_audio_denoise,
temp_dir=req.temp_dir or settings.get("default_temp_dir") or None,
output_dir=output_dir,
source_filename=source_filename
)
jobs_db[job_id] = job
job_queue.put(job_id)
queued_jobs.append({"job_id": job_id, "status": "queued"})
save_jobs_db()
# Broadcast initial queued progress for all queued jobs
for qj in queued_jobs:
broadcast_progress(qj["job_id"], {
"status": "queued",
"progress": 0.0,
"current_frame": 0,
"total_frames": 0,
"eta": "Calculating..."
})
return {"jobs": queued_jobs}
@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.")
# Remove from queue if it was queued
job_queue.remove(job_id)
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"
})
save_jobs_db()
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
if req.server_file_path:
if os.path.exists(req.server_file_path) and os.path.isfile(req.server_file_path):
file_path = req.server_file_path
else:
raise HTTPException(status_code=404, detail="Server file path not found.")
else:
if not req.file_id:
raise HTTPException(status_code=400, detail="Either file_id or server_file_path must be provided.")
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, path: str | None = None):
allowed_dirs = {os.path.abspath(upscaler.OUTPUT_DIR)}
settings = load_settings()
default_out = settings.get("default_output_dir", "").strip()
if default_out:
allowed_dirs.add(os.path.abspath(default_out))
for job in jobs_db.values():
if getattr(job, "output_dir", None):
allowed_dirs.add(os.path.abspath(job.output_dir))
if path:
resolved_path = os.path.abspath(path)
if os.path.isdir(resolved_path):
file_path = os.path.join(resolved_path, filename)
else:
file_path = resolved_path
else:
file_path = os.path.join(upscaler.OUTPUT_DIR, filename)
resolved_file_path = os.path.abspath(file_path)
allowed = False
for allowed_dir in allowed_dirs:
if resolved_file_path.startswith(allowed_dir):
allowed = True
break
if not allowed:
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, media_type="application/octet-stream", filename=os.path.basename(file_path))
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 in queue-sorted order"""
active_id = current_running_job_id
queued_ids = job_queue.get_all()
# Sort active first, then queued in order, then history by start time descending
def get_sort_key(job):
if job.job_id == active_id:
return (0, 0)
elif job.job_id in queued_ids:
return (1, queued_ids.index(job.job_id))
else:
t = job.start_time if job.start_time is not None else 0
return (2, -t)
sorted_jobs = sorted(jobs_db.values(), key=get_sort_key)
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),
"queue_position": queued_ids.index(job.job_id) if job.job_id in queued_ids else -1 if job.job_id == active_id else None
}
for job in sorted_jobs
]
@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.")
# Remove from queue if it is queued
job_queue.remove(job_id)
# 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]
save_jobs_db()
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()
# Re-initialize custom queue
global job_queue
job_queue = CustomJobQueue()
# Reset upload metadata file
save_upload_metadata({})
save_jobs_db()
return {"status": "all purged"}
class ReorderQueueRequest(BaseModel):
job_ids: List[str]
@app.post("/api/queue/reorder")
def reorder_queue(req: ReorderQueueRequest):
"""Reorder the job queue with preemption support"""
global current_running_job_id
running_id = current_running_job_id
if running_id and running_id in req.job_ids:
idx = req.job_ids.index(running_id)
if idx > 0:
# Preempt the running job!
job = jobs_db.get(running_id)
if job:
print(f"Preempting running job {running_id} (moved to index {idx} in queue)")
try:
job.pause()
except Exception as e:
print(f"Error pausing job {running_id} for preemption: {e}")
# Re-queue the job
job.status = "queued"
job.error = None
job.eta = "Queued (preempted)..."
if hasattr(job, "_is_paused"):
job._is_paused = False
job_queue.put(running_id)
broadcast_progress(running_id, {
"status": "queued",
"progress": job.progress,
"current_frame": job.current_frame,
"total_frames": job.total_frames,
"eta": "Queued (preempted)..."
})
job_queue.reorder(req.job_ids)
save_jobs_db()
q = job_queue.get_all()
if current_running_job_id:
q = [current_running_job_id] + q
return {"status": "success", "queue": q}
@app.get("/api/queue")
def get_queue():
"""Get the current job queue order"""
q = job_queue.get_all()
if current_running_job_id:
q = [current_running_job_id] + q
return {"queue": q}
@app.post("/api/upscale/resume/{job_id}")
def resume_job(job_id: str):
"""Resume an interrupted/failed/paused upscale job"""
job = jobs_db.get(job_id)
if not job:
raise HTTPException(status_code=404, detail="Job not found.")
# Re-queue the job
job.status = "queued"
job.error = None
job.eta = "Queued for resume..."
if hasattr(job, "_is_paused"):
job._is_paused = False
job_queue.put(job_id)
save_jobs_db()
broadcast_progress(job_id, {
"status": "queued",
"progress": job.progress,
"current_frame": job.current_frame,
"total_frames": job.total_frames,
"eta": "Queued for resume..."
})
return {"job_id": job_id, "status": "queued"}
@app.post("/api/upscale/pause/{job_id}")
def pause_job(job_id: str):
"""Pause a running or queued job"""
job = jobs_db.get(job_id)
if not job:
raise HTTPException(status_code=404, detail="Job not found.")
if job.status not in ["queued", "pending", "analyzing", "extracting", "upscaling", "restoring_faces", "interpolating", "assembling"]:
raise HTTPException(status_code=400, detail=f"Job in status {job.status} cannot be paused.")
# Remove from queue if it is in queue
job_queue.remove(job_id)
# Call pause logic on the job
job.pause()
save_jobs_db()
broadcast_progress(job_id, {
"status": "paused",
"progress": job.progress,
"current_frame": job.current_frame,
"total_frames": job.total_frames,
"eta": "Paused"
})
return {"job_id": job_id, "status": "paused"}
# 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_dirs = {upscaler.OUTPUT_DIR}
settings = load_settings()
default_out = settings.get("default_output_dir", "").strip()
if default_out:
output_dirs.add(default_out)
for job in jobs_db.values():
if job.status == "completed" and job.output_file:
output_dirs.add(os.path.dirname(job.output_file))
elif getattr(job, "output_dir", None):
output_dirs.add(job.output_dir)
files = []
seen_files = set()
for out_dir in output_dirs:
if not os.path.exists(out_dir):
continue
try:
for filename in os.listdir(out_dir):
if filename.startswith("original_"):
continue
file_path = os.path.join(out_dir, filename)
abs_file_path = os.path.abspath(file_path)
if abs_file_path in seen_files:
continue
if os.path.isfile(file_path):
seen_files.add(abs_file_path)
stat = os.stat(file_path)
if os.path.abspath(out_dir) == os.path.abspath(upscaler.OUTPUT_DIR):
url = f"/api/download/{filename}"
else:
url = f"/api/download/{filename}?path={urllib.parse.quote(out_dir)}"
files.append({
"filename": filename,
"size": stat.st_size,
"modified": stat.st_mtime,
"url": url,
"path": file_path
})
except Exception as e:
print(f"Error scanning output dir {out_dir}: {e}")
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, path: str | None = None):
allowed_dirs = {os.path.abspath(upscaler.OUTPUT_DIR)}
settings = load_settings()
default_out = settings.get("default_output_dir", "").strip()
if default_out:
allowed_dirs.add(os.path.abspath(default_out))
for job in jobs_db.values():
if getattr(job, "output_dir", None):
allowed_dirs.add(os.path.abspath(job.output_dir))
if path:
resolved_path = os.path.abspath(path)
if os.path.isdir(resolved_path):
file_path = os.path.join(resolved_path, filename)
else:
file_path = resolved_path
else:
file_path = os.path.join(upscaler.OUTPUT_DIR, filename)
resolved_file_path = os.path.abspath(file_path)
allowed = False
for allowed_dir in allowed_dirs:
if resolved_file_path.startswith(allowed_dir):
allowed = True
break
if not allowed:
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}
# ─── Settings Endpoints ─────────────────────────────────────────────
class UpdateSettingsRequest(BaseModel):
default_temp_dir: str = ""
default_output_dir: str = ""
role: str = "standalone"
coordinator_url: str = ""
worker_node_url: str = "http://localhost:8000"
use_shared_storage: bool = False
shared_storage_path: str = ""
smb_shares: List[Any] = []
@app.get("/api/settings")
def get_settings():
"""Get persistent app settings"""
settings = load_settings()
return settings
def migrate_temp_directory_task(old_dir: str, new_dir: str, active_job_ids: list):
import time
import shutil
# 1. Sleep to allow pipeline threads to clean up and exit
time.sleep(2.0)
old_real = os.path.realpath(old_dir) if old_dir else ""
new_real = os.path.realpath(new_dir) if new_dir else ""
if old_real and new_real and old_real != new_real:
# 2. Migrate files
print(f"Migrating temp files from {old_real} to {new_real}...")
if os.path.exists(old_real):
for item in os.listdir(old_real):
src = os.path.join(old_real, item)
dst = os.path.join(new_real, item)
try:
if os.path.exists(dst):
if os.path.isdir(dst):
shutil.rmtree(dst)
else:
os.remove(dst)
shutil.move(src, dst)
print(f"Moved {src} to {dst}")
except Exception as e:
print(f"Error moving {src} to {dst}: {e}")
# 3. Update job temp_dir attributes in database
for job_id, job in jobs_db.items():
job_temp = getattr(job, "temp_dir", None)
if not job_temp or os.path.realpath(job_temp) == old_real:
job.temp_dir = new_real
save_jobs_db()
# 4. Resume the paused jobs
for job_id in active_job_ids:
job = jobs_db.get(job_id)
if job:
print(f"Auto-resuming job {job_id} at new temp location...")
job.status = "queued"
job.error = None
job.eta = "Queued for resume..."
if hasattr(job, "_is_paused"):
job._is_paused = False
job_queue.put(job_id)
save_jobs_db()
broadcast_progress(job_id, {
"status": "queued",
"progress": job.progress,
"current_frame": job.current_frame,
"total_frames": job.total_frames,
"eta": "Queued for resume..."
})
@app.post("/api/settings")
def update_settings(req: UpdateSettingsRequest, background_tasks: BackgroundTasks):
"""Update persistent app settings"""
settings = load_settings()
old_dir = settings.get("default_temp_dir", "").strip()
if not old_dir:
old_dir = upscaler.TEMP_DIR
new_dir = req.default_temp_dir.strip()
if new_dir:
try:
os.makedirs(new_dir, exist_ok=True)
except OSError as e:
raise HTTPException(status_code=400, detail=f"Invalid default_temp_dir path: {e}")
new_out_dir = req.default_output_dir.strip()
if new_out_dir:
try:
os.makedirs(new_out_dir, exist_ok=True)
except OSError as e:
raise HTTPException(status_code=400, detail=f"Invalid default_output_dir path: {e}")
new_shared_path = req.shared_storage_path.strip()
if req.use_shared_storage and new_shared_path:
try:
os.makedirs(new_shared_path, exist_ok=True)
except OSError as e:
raise HTTPException(status_code=400, detail=f"Invalid shared_storage_path path: {e}")
role_changed = (settings.get("role") != req.role)
# Save settings
settings["default_temp_dir"] = new_dir
settings["default_output_dir"] = new_out_dir
settings["role"] = req.role
settings["coordinator_url"] = req.coordinator_url
settings["worker_node_url"] = req.worker_node_url
settings["use_shared_storage"] = req.use_shared_storage
settings["shared_storage_path"] = req.shared_storage_path
settings["smb_shares"] = req.smb_shares
save_settings(settings)
if role_changed:
restart_dist_worker()
# Identify and pause active jobs
active_jobs = []
old_real = os.path.realpath(old_dir)
new_real = os.path.realpath(new_dir) if new_dir else os.path.realpath(upscaler.TEMP_DIR)
if old_real != new_real:
for job_id, job in list(jobs_db.items()):
if job.status in ["queued", "pending", "analyzing", "extracting", "upscaling", "restoring_faces", "interpolating", "assembling"]:
active_jobs.append(job_id)
try:
job.pause()
job_queue.remove(job_id)
broadcast_progress(job_id, {
"status": "paused",
"progress": job.progress,
"current_frame": job.current_frame,
"total_frames": job.total_frames,
"eta": "Paused"
})
except Exception as e:
print(f"Error pausing job {job_id} for temp migration: {e}")
save_jobs_db()
# Start background migration and resumption task
background_tasks.add_task(migrate_temp_directory_task, old_dir, new_dir, active_jobs)
return {"status": "saved", "settings": settings, "migrating_jobs": active_jobs}
class CreateDirRequest(BaseModel):
parent_path: str
name: str
@app.get("/api/system/browse-dir")
def browse_dir(path: str = "", show_files: bool = False):
"""Browse directories on the system"""
if not path:
path = os.path.expanduser("~")
path = os.path.abspath(path)
if not os.path.exists(path):
raise HTTPException(status_code=404, detail="Path does not exist")
if not os.path.isdir(path):
raise HTTPException(status_code=400, detail="Path is not a directory")
try:
entries = []
# Add parent directory entry if not at root
parent = os.path.dirname(path)
if parent != path:
entries.append({
"name": "..",
"path": parent,
"is_dir": True
})
video_extensions = {'.mp4', '.mkv', '.avi', '.mov', '.webm', '.flv', '.m4v', '.wmv', '.mpg', '.mpeg'}
for name in sorted(os.listdir(path)):
full_path = os.path.join(path, name)
try:
if os.path.isdir(full_path):
entries.append({
"name": name,
"path": full_path,
"is_dir": True
})
elif show_files:
_, ext = os.path.splitext(name.lower())
if ext in video_extensions:
entries.append({
"name": name,
"path": full_path,
"is_dir": False
})
except PermissionError:
continue
return {
"current_path": path,
"entries": entries
}
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
@app.get("/api/system/file-metadata")
def get_file_metadata(path: str):
"""Get video metadata for a local server file"""
if not os.path.exists(path) or not os.path.isfile(path):
raise HTTPException(status_code=404, detail="File not found.")
info = upscaler.get_video_info(path)
if not info:
raise HTTPException(status_code=400, detail="Could not read video file metadata.")
return {
"original_filename": os.path.basename(path),
"path": path,
"size_bytes": os.path.getsize(path),
"metadata": info
}
@app.post("/api/system/create-dir")
def create_dir(req: CreateDirRequest):
"""Create a new directory"""
full_path = os.path.abspath(os.path.join(req.parent_path, req.name))
try:
os.makedirs(full_path, exist_ok=True)
return {"status": "ok", "path": full_path}
except Exception as e:
raise HTTPException(status_code=400, detail=str(e))
@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)
class CleanupTempRequest(BaseModel):
paths: List[str] | None = None
def get_dir_stats(path: str) -> tuple[int, int]:
total_size = 0
total_files = 0
if not os.path.exists(path):
return 0, 0
for root, dirs, files in os.walk(path):
for f in files:
fp = os.path.join(root, f)
try:
if os.path.exists(fp):
total_size += os.path.getsize(fp)
total_files += 1
except Exception:
pass
return total_size, total_files
@app.get("/api/system/temp-info")
def get_temp_info():
temp_folders = []
# 1. Collect all potential temp root directories
# Default temp root
temp_roots = {upscaler.TEMP_DIR}
# Custom temp roots from jobs db
for job in jobs_db.values():
if getattr(job, "temp_dir", None):
temp_roots.add(job.temp_dir)
# 2. Scan folders inside the temp roots
seen_paths = set()
for root_dir in temp_roots:
if not os.path.exists(root_dir):
continue
try:
for item in os.listdir(root_dir):
item_path = os.path.join(root_dir, item)
if not os.path.isdir(item_path):
continue
if item_path in seen_paths:
continue
seen_paths.add(item_path)
# Check if it corresponds to a job_id (which is usually a UUID)
job_id = item
job = jobs_db.get(job_id)
size_bytes, file_count = get_dir_stats(item_path)
if job:
job_status = job.status
video_name = os.path.basename(job.video_path)
else:
job_status = "orphaned"
video_name = "Unknown Video (Orphaned)"
temp_folders.append({
"job_id": job_id,
"path": item_path,
"size_bytes": size_bytes,
"file_count": file_count,
"status": job_status,
"video_name": video_name
})
except Exception as e:
print(f"Error scanning temp root {root_dir}: {e}")
return {"temp_folders": temp_folders}
@app.post("/api/system/temp-cleanup")
def cleanup_temp(req: CleanupTempRequest):
import shutil
cleaned = []
errors = []
# Identify active job IDs
active_statuses = ["analyzing", "extracting", "upscaling", "restoring_faces", "interpolating", "assembling"]
active_job_ids = {job.job_id for job in jobs_db.values() if job.status in active_statuses}
# Get current temp folder status
temp_info = get_temp_info()
folders = temp_info["temp_folders"]
target_paths = req.paths if req.paths is not None else [f["path"] for f in folders]
for f in folders:
path = f["path"]
if path not in target_paths:
continue
# Safety checks:
if f["job_id"] in active_job_ids:
errors.append(f"Cannot clean up active job {f['job_id']}")
continue
resolved_path = os.path.realpath(path)
is_sub_of_temp = resolved_path.startswith(os.path.realpath(upscaler.TEMP_DIR))
if not is_sub_of_temp:
for job in jobs_db.values():
if getattr(job, "temp_dir", None):
if resolved_path.startswith(os.path.realpath(job.temp_dir)):
is_sub_of_temp = True
break
if not is_sub_of_temp:
errors.append(f"Safety constraint: Path {path} is not in a valid temp root.")
continue
try:
if os.path.exists(path):
shutil.rmtree(path)
cleaned.append(path)
# Update DB state if needed
job = jobs_db.get(f["job_id"])
if job and job.status in ["paused", "interrupted", "failed"]:
job.status = "interrupted"
job.progress = 0.0
job.current_frame = 0
job.eta = "Temp files cleaned. Will restart on resume."
else:
cleaned.append(path)
except Exception as e:
errors.append(f"Error removing {path}: {str(e)}")
class ReimportTempRequest(BaseModel):
video_path: str
model: str = "realesrgan-x4plus"
scale: int = 4
tile_size: int = 256
preserve_audio: bool = True
ai_face_restoration: bool = False
ai_rife_interpolation: bool = False
ai_audio_denoise: bool = False
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
transcode_format: str = "mp4"
temp_dir: str | None = None
@app.post("/api/system/temp-reimport/{job_id}")
def reimport_temp(job_id: str, req: ReimportTempRequest):
if job_id in jobs_db:
raise HTTPException(status_code=400, detail="Job already exists in database.")
# Verify the video path is valid
if not os.path.exists(req.video_path):
raise HTTPException(status_code=400, detail=f"Video file not found at: {req.video_path}")
# Reconstruct the job dictionary
job_dict = {
"job_id": job_id,
"video_path": req.video_path,
"model": req.model,
"scale": req.scale,
"tile_size": req.tile_size,
"preserve_audio": req.preserve_audio,
"ai_face_restoration": req.ai_face_restoration,
"ai_rife_interpolation": req.ai_rife_interpolation,
"ai_audio_denoise": req.ai_audio_denoise,
"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,
"transcode_format": req.transcode_format,
"temp_dir": req.temp_dir,
"status": "paused", # Start as paused so it can be resumed
"progress": 0.0,
"current_frame": 0,
"total_frames": 0,
"eta": "Ready to resume"
}
job = upscaler.UpscaleJob.from_dict(job_dict)
# Calculate progress based on existing frames in the temp directory if possible
job_temp_dir = os.path.join(req.temp_dir if req.temp_dir else upscaler.TEMP_DIR, job_id)
input_frames_dir = os.path.join(job_temp_dir, "input_frames")
output_frames_dir = os.path.join(job_temp_dir, "output_frames")
total_frames = 0
if os.path.exists(input_frames_dir):
total_frames = len([f for f in os.listdir(input_frames_dir) if f.startswith("frame_")])
job.total_frames = total_frames
if os.path.exists(output_frames_dir):
out_frames = len([f for f in os.listdir(output_frames_dir) if f.startswith("frame_")])
job.current_frame = out_frames
if total_frames > 0:
job.progress = min(99.0, round((out_frames / total_frames) * 100.0, 2))
jobs_db[job_id] = job
save_jobs_db()
# Auto-resume the job by adding it to the queue
job_queue.put(job_id)
job.update_status("queued", eta="Queued for resume...")
save_jobs_db()
return {"status": "success", "job_id": job_id}
# ─── Distributed Processing Coordinator & Worker Endpoints ────────────────
class ChunkDoneRequest(BaseModel):
job_id: str
chunk_id: str
@app.get("/api/dist/request-chunk")
def request_chunk(worker_url: str = ""):
settings = load_settings()
if settings.get("role") != "coordinator":
raise HTTPException(status_code=400, detail="Node is not configured as a coordinator.")
with upscaler.dist_chunks_lock:
job_id = current_running_job_id
if not job_id:
return {"status": "idle"}
job = jobs_db.get(job_id)
if not job or job.status != "upscaling":
return {"status": "idle"}
chunks = upscaler.dist_chunks.get(job_id, [])
for chunk in chunks:
if chunk["status"] == "pending":
chunk["status"] = "assigned"
chunk["worker_url"] = worker_url
chunk["updated_at"] = time.time()
return {
"status": "assigned",
"job_id": job_id,
"chunk_id": chunk["chunk_id"],
"model": job.model,
"scale": job.scale,
"tile_size": job.tile_size,
"gpu_ids": job.gpu_ids,
"tta": job.tta,
"files": chunk["files"],
"use_shared_storage": settings.get("use_shared_storage", False),
"shared_storage_path": settings.get("shared_storage_path", ""),
"job_temp_dir": os.path.join(settings.get("shared_storage_path") if settings.get("use_shared_storage") else upscaler.TEMP_DIR, job_id)
}
return {"status": "no_chunks"}
@app.get("/api/dist/download-frame/{job_id}/{filename}")
def download_frame(job_id: str, filename: str):
settings = load_settings()
if settings.get("role") != "coordinator":
raise HTTPException(status_code=400, detail="Node is not configured as a coordinator.")
if "/" in filename or "\\" in filename or filename in [".", ".."]:
raise HTTPException(status_code=400, detail="Invalid filename")
base_temp = settings.get("shared_storage_path") if settings.get("use_shared_storage") else upscaler.TEMP_DIR
file_path = os.path.join(base_temp, job_id, "input_frames", filename)
if not os.path.exists(file_path):
raise HTTPException(status_code=404, detail="Frame not found")
return FileResponse(file_path)
@app.post("/api/dist/upload-chunk/{job_id}")
async def upload_chunk(job_id: str, files: List[UploadFile] = File(...)):
settings = load_settings()
if settings.get("role") != "coordinator":
raise HTTPException(status_code=400, detail="Node is not configured as a coordinator.")
base_temp = settings.get("shared_storage_path") if settings.get("use_shared_storage") else upscaler.TEMP_DIR
output_frames_dir = os.path.join(base_temp, job_id, "output_frames")
os.makedirs(output_frames_dir, exist_ok=True)
for file in files:
filename = os.path.basename(file.filename)
if "/" in filename or "\\" in filename or filename in [".", ".."]:
continue
save_path = os.path.join(output_frames_dir, filename)
with open(save_path, "wb") as buffer:
content = await file.read()
buffer.write(content)
return {"status": "success"}
@app.post("/api/dist/chunk-done")
def chunk_done(req: ChunkDoneRequest):
settings = load_settings()
if settings.get("role") != "coordinator":
raise HTTPException(status_code=400, detail="Node is not configured as a coordinator.")
with upscaler.dist_chunks_lock:
chunks = upscaler.dist_chunks.get(req.job_id, [])
for chunk in chunks:
if chunk["chunk_id"] == req.chunk_id:
chunk["status"] = "completed"
chunk["updated_at"] = time.time()
return {"status": "completed"}
raise HTTPException(status_code=404, detail="Chunk not found.")
# ─── SMB Mounting & Background Worker Thread ───────────────────────────────
def mount_smb_share(share: dict) -> bool:
"""Mount an SMB share if not already mounted"""
local_path = share.get("local_path", "").strip()
share_path = share.get("share_path", "").strip()
username = share.get("username", "").strip()
password = share.get("password", "").strip()
if not local_path or not share_path:
print("SMB Share: missing local_path or share_path")
return False
try:
os.makedirs(local_path, exist_ok=True)
except Exception as e:
print(f"SMB Share: failed to create mount point {local_path}: {e}")
return False
try:
with open("/proc/mounts", "r") as f:
mounts = f.read()
if os.path.realpath(local_path) in mounts or share_path in mounts:
print(f"SMB Share: {share_path} is already mounted at {local_path}")
return True
except Exception:
pass
options = []
if username:
options.append(f"username={username}")
if password:
options.append(f"password={password}")
cmd = ["mount", "-t", "cifs", share_path, local_path]
if options:
cmd.extend(["-o", ",".join(options)])
try:
print(f"Mounting SMB share: {' '.join(cmd)}")
res = subprocess.run(cmd, capture_output=True, text=True)
if res.returncode == 0:
print(f"SMB Share: successfully mounted {share_path} to {local_path}")
return True
else:
print(f"SMB Share: failed to mount (exit {res.returncode}): {res.stderr}")
return False
except Exception as e:
print(f"SMB Share: exception mounting: {e}")
return False
def upload_files_multipart(url: str, files: List[tuple[str, str]]):
"""Upload files using multipart/form-data via urllib"""
import uuid
import mimetypes
boundary = f"----WebKitFormBoundary{uuid.uuid4().hex}"
parts = []
for field_name, file_path in files:
filename = os.path.basename(file_path)
mime_type = mimetypes.guess_type(file_path)[0] or "image/jpeg"
parts.append(f"--{boundary}".encode('utf-8'))
parts.append(f'Content-Disposition: form-data; name="{field_name}"; filename="{filename}"'.encode('utf-8'))
parts.append(f"Content-Type: {mime_type}".encode('utf-8'))
parts.append(b"")
with open(file_path, "rb") as f:
parts.append(f.read())
parts.append(f"--{boundary}--".encode('utf-8'))
body = b"\r\n".join(parts)
headers = {
"Content-Type": f"multipart/form-data; boundary={boundary}",
"Content-Length": str(len(body))
}
req = urllib.request.Request(url, data=body, headers=headers, method="POST")
with urllib.request.urlopen(req) as res:
return res.read()
def process_dist_chunk(chunk: dict) -> bool:
job_id = chunk["job_id"]
chunk_id = chunk["chunk_id"]
model = chunk["model"]
scale = chunk["scale"]
tile_size = chunk["tile_size"]
gpu_ids = chunk.get("gpu_ids")
files = chunk["files"]
use_shared_storage = chunk.get("use_shared_storage", False)
coordinator_url = load_settings().get("coordinator_url", "").rstrip("/")
print(f"Worker: Processing chunk {chunk_id} for job {job_id} ({len(files)} files)")
local_job_temp = os.path.join(upscaler.TEMP_DIR, f"worker_{job_id}_{chunk_id}")
local_input_dir = os.path.join(local_job_temp, "input_frames")
local_output_dir = os.path.join(local_job_temp, "output_frames")
if not use_shared_storage:
os.makedirs(local_input_dir, exist_ok=True)
os.makedirs(local_output_dir, exist_ok=True)
try:
input_paths = []
output_paths = []
for filename in files:
if use_shared_storage:
job_temp_dir = chunk["job_temp_dir"]
in_path = os.path.join(job_temp_dir, "input_frames", filename)
out_path = os.path.join(job_temp_dir, "output_frames", filename)
else:
in_path = os.path.join(local_input_dir, filename)
out_path = os.path.join(local_output_dir, filename)
download_url = f"{coordinator_url}/api/dist/download-frame/{job_id}/{filename}"
try:
urllib.request.urlretrieve(download_url, in_path)
except Exception as dl_err:
print(f"Worker: Failed to download frame {filename}: {dl_err}")
return False
input_paths.append(in_path)
output_paths.append(out_path)
for in_p, out_p in zip(input_paths, output_paths):
success = upscaler.upscale_image_file(in_p, out_p, model, scale, tile_size, gpu_ids)
if not success:
print(f"Worker: Failed to upscale frame {in_p}")
return False
if not use_shared_storage:
upload_url = f"{coordinator_url}/api/dist/upload-chunk/{job_id}"
file_tuples = [("files", out_p) for out_p in output_paths]
try:
upload_files_multipart(upload_url, file_tuples)
except Exception as up_err:
print(f"Worker: Failed to upload output frames: {up_err}")
return False
done_url = f"{coordinator_url}/api/dist/chunk-done"
try:
req = urllib.request.Request(
done_url,
data=json.dumps({"job_id": job_id, "chunk_id": chunk_id}).encode('utf-8'),
headers={"Content-Type": "application/json"},
method="POST"
)
with urllib.request.urlopen(req) as res:
res.read()
except Exception as done_err:
print(f"Worker: Failed to notify done: {done_err}")
return False
return True
finally:
if not use_shared_storage and os.path.exists(local_job_temp):
try:
shutil.rmtree(local_job_temp)
except Exception:
pass
dist_worker_running = False
dist_worker_thread = None
dist_worker_stop_event = threading.Event()
def dist_worker_loop():
global dist_worker_running
dist_worker_running = True
print("Background worker loop started.")
settings = load_settings()
for share in settings.get("smb_shares", []):
mount_smb_share(share)
while not dist_worker_stop_event.is_set():
settings = load_settings()
role = settings.get("role", "standalone")
if role != "worker":
time.sleep(5.0)
continue
coordinator_url = settings.get("coordinator_url", "").strip().rstrip("/")
if not coordinator_url:
time.sleep(10.0)
continue
worker_node_url = settings.get("worker_node_url", "http://localhost:8000")
try:
req_url = f"{coordinator_url}/api/dist/request-chunk?worker_url={urllib.parse.quote(worker_node_url)}"
req = urllib.request.Request(req_url, method="GET")
with urllib.request.urlopen(req, timeout=10) as res:
chunk_data = json.loads(res.read().decode('utf-8'))
if chunk_data.get("status") == "assigned":
success = process_dist_chunk(chunk_data)
if not success:
print(f"Worker: failed to process chunk {chunk_data.get('chunk_id')}. Will retry later.")
else:
print(f"Worker: successfully processed chunk {chunk_data.get('chunk_id')}")
else:
time.sleep(3.0)
except Exception as e:
print(f"Worker loop error: {e}")
time.sleep(5.0)
dist_worker_running = False
print("Background worker loop stopped.")
def restart_dist_worker():
global dist_worker_thread, dist_worker_stop_event
dist_worker_stop_event.set()
if dist_worker_thread and dist_worker_thread.is_alive():
dist_worker_thread.join(timeout=2.0)
dist_worker_stop_event.clear()
dist_worker_thread = threading.Thread(target=dist_worker_loop, daemon=True)
dist_worker_thread.start()
# Mount static folder
app.mount("/", StaticFiles(directory=os.path.join(upscaler.BASE_DIR, "static"), html=True), name="static")