Add video generation API with job tracking and model loading
This commit is contained in:
+293
@@ -0,0 +1,293 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import socket
|
||||
import sys
|
||||
import time
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import httpx
|
||||
|
||||
from app.config import settings
|
||||
|
||||
|
||||
DEFAULT_HOST = "http://localhost:4033"
|
||||
|
||||
|
||||
def _valid_image(path: str) -> str:
|
||||
if not os.path.isfile(path):
|
||||
raise argparse.ArgumentTypeError(f"Image file not found: {path}")
|
||||
return path
|
||||
|
||||
|
||||
def _positive_int(val: str) -> int:
|
||||
v = int(val)
|
||||
if v <= 0:
|
||||
raise argparse.ArgumentTypeError(f"Must be > 0, got {v}")
|
||||
return v
|
||||
|
||||
|
||||
def _positive_float(val: str) -> float:
|
||||
v = float(val)
|
||||
if v <= 0:
|
||||
raise argparse.ArgumentTypeError(f"Must be > 0, got {v}")
|
||||
return v
|
||||
|
||||
|
||||
def _valid_frames(val: str) -> int:
|
||||
v = int(val)
|
||||
if v % 8 != 1:
|
||||
n = v // 8
|
||||
candidates = sorted(str(8 * i + 1) for i in range(max(0, n - 2), n + 3) if 0 < 8 * i + 1 <= settings.max_frame_count)
|
||||
raise argparse.ArgumentTypeError(f"Must be 8n+1, got {v}. Closest: {candidates}")
|
||||
if v > settings.max_frame_count:
|
||||
raise argparse.ArgumentTypeError(f"Must be <= {settings.max_frame_count}, got {v}")
|
||||
return v
|
||||
|
||||
|
||||
def _build_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(
|
||||
prog="revids",
|
||||
description="CLI for the Revids video generation API",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--host",
|
||||
default=DEFAULT_HOST,
|
||||
help="API base URL (default: http://localhost:4033)",
|
||||
)
|
||||
sub = parser.add_subparsers(dest="command")
|
||||
|
||||
# health
|
||||
health = sub.add_parser("health", help="Check server health")
|
||||
|
||||
# generate
|
||||
gen = sub.add_parser("generate", help="Submit a video generation job")
|
||||
gen.add_argument("image", type=_valid_image, help="Path to input image")
|
||||
gen.add_argument("--prompt", default=None, help="Text prompt")
|
||||
gen.add_argument("--width", type=_positive_int, default=768, help="Video width (default: 768)")
|
||||
gen.add_argument("--height", type=_positive_int, default=512, help="Video height (default: 512)")
|
||||
gen.add_argument("--frames", type=_valid_frames, default=65, help="Frame count (must be 8n+1, default: 65)")
|
||||
gen.add_argument("--fps", type=_positive_float, default=24.0, help="Frames per second (default: 24)")
|
||||
gen.add_argument("--seed", type=int, default=None, help="Random seed")
|
||||
gen.add_argument("--wait", action="store_true", help="Wait for job completion and download")
|
||||
gen.add_argument("--timeout", type=int, default=0, help="Timeout in seconds for --wait (0 = no limit)")
|
||||
gen.add_argument("--poll-interval", type=int, default=3, help="Seconds between status polls (default: 3)")
|
||||
gen.add_argument("--output", default=None, help="Output path for downloaded video (default: <job_id>.mp4)")
|
||||
|
||||
# status
|
||||
stat = sub.add_parser("status", help="Check job status")
|
||||
stat.add_argument("job_id", help="Job ID")
|
||||
|
||||
# download
|
||||
dl = sub.add_parser("download", help="Download generated video")
|
||||
dl.add_argument("job_id", help="Job ID")
|
||||
dl.add_argument("--output", default=None, help="Output file path (default: <job_id>.mp4)")
|
||||
|
||||
# delete
|
||||
del_cmd = sub.add_parser("delete", help="Delete a job")
|
||||
del_cmd.add_argument("job_id", help="Job ID")
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
def _check_server(host: str) -> None:
|
||||
try:
|
||||
host_addr = socket.gethostbyname(socket.gethostname())
|
||||
hostname = urlparse(host).hostname or host.split(":")[0]
|
||||
server_addr = socket.gethostbyname(hostname)
|
||||
if host_addr != server_addr and "127.0.0.1" not in host:
|
||||
print(f"[WARN] API host {host} appears to be on a different machine.")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def cmd_health(host: str) -> None:
|
||||
print(f"Checking health at {host}...")
|
||||
try:
|
||||
with httpx.Client(timeout=10) as client:
|
||||
resp = client.get(f"{host}/health")
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
print(f" Status: {data['status']}")
|
||||
print(f" GPU available: {data['gpu']}")
|
||||
if data.get("gpu_name"):
|
||||
print(f" GPU name: {data['gpu_name']}")
|
||||
except httpx.ConnectError:
|
||||
print(f" ERROR: Cannot connect to {host}. Is the server running?", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
except httpx.HTTPStatusError as e:
|
||||
print(f" ERROR: Server responded with {e.response.status_code}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
def cmd_generate(args: argparse.Namespace) -> None:
|
||||
_check_server(args.host)
|
||||
|
||||
print(f"Submitting job to {args.host}...")
|
||||
try:
|
||||
with httpx.Client(timeout=60) as client:
|
||||
with open(args.image, "rb") as f:
|
||||
data: dict[str, object] = {
|
||||
"prompt": args.prompt or "",
|
||||
"width": args.width,
|
||||
"height": args.height,
|
||||
"num_frames": args.frames,
|
||||
"fps": args.fps,
|
||||
}
|
||||
if args.seed is not None:
|
||||
data["seed"] = args.seed
|
||||
resp = client.post(
|
||||
f"{args.host}/generate",
|
||||
files={"image": (os.path.basename(args.image), f)},
|
||||
data=data,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
except httpx.ConnectError:
|
||||
print(f"ERROR: Cannot connect to {args.host}. Is the server running?", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
except httpx.HTTPStatusError as e:
|
||||
print(f"ERROR: Server error ({e.response.status_code}): {e.response.text}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
job_id = data["job_id"]
|
||||
print(f"Job submitted: {job_id}")
|
||||
|
||||
if args.wait:
|
||||
poll_status(args, job_id)
|
||||
download_video(args.host, job_id, args.output)
|
||||
|
||||
|
||||
def poll_status(args: argparse.Namespace, job_id: str) -> None:
|
||||
deadline = None
|
||||
if args.timeout > 0:
|
||||
deadline = time.monotonic() + args.timeout
|
||||
|
||||
with httpx.Client(timeout=10) as client:
|
||||
first = True
|
||||
while True:
|
||||
if deadline and time.monotonic() > deadline:
|
||||
print(f"ERROR: Timed out waiting for job {job_id}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
if not first:
|
||||
time.sleep(args.poll_interval)
|
||||
first = False
|
||||
|
||||
try:
|
||||
resp = client.get(f"{args.host}/jobs/{job_id}")
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
except ValueError as e:
|
||||
print(f"ERROR: Invalid server response: {e}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
except httpx.RequestError:
|
||||
print(f"ERROR: Cannot reach server at {args.host}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
except httpx.HTTPStatusError as e:
|
||||
print(f"ERROR: Server error ({e.response.status_code})", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
status = data["status"]
|
||||
ts = datetime.now(timezone.utc).strftime("%H:%M:%S")
|
||||
if status in ("pending", "processing"):
|
||||
print(f" [{ts}] Job {job_id}: {status}...")
|
||||
elif status == "completed":
|
||||
print(f" [{ts}] Job {job_id}: completed")
|
||||
return
|
||||
elif status == "failed":
|
||||
err = data.get("error", "unknown error")
|
||||
print(f" [{ts}] Job {job_id}: FAILED - {err}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
def cmd_status(host: str, job_id: str) -> None:
|
||||
try:
|
||||
with httpx.Client(timeout=10) as client:
|
||||
resp = client.get(f"{host}/jobs/{job_id}")
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
except httpx.ConnectError:
|
||||
print(f"ERROR: Cannot connect to {host}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
except httpx.HTTPStatusError as e:
|
||||
if e.response.status_code == 404:
|
||||
print(f"Job {job_id} not found", file=sys.stderr)
|
||||
else:
|
||||
print(f"ERROR: Server error ({e.response.status_code})", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
print(f" Job ID: {data['job_id']}")
|
||||
print(f" Status: {data['status']}")
|
||||
print(f" Prompt: {data.get('prompt', '')}")
|
||||
print(f" Created: {data.get('created_at', 'N/A')}")
|
||||
print(f" Completed: {data.get('completed_at', 'N/A')}")
|
||||
if data.get("params"):
|
||||
print(f" Params: {data['params']}")
|
||||
if data.get("error"):
|
||||
print(f" Error: {data['error']}")
|
||||
|
||||
|
||||
def download_video(host: str, job_id: str, output: Optional[str] = None) -> None:
|
||||
out_path = output or f"{job_id}.mp4"
|
||||
|
||||
try:
|
||||
with httpx.Client(timeout=120) as client:
|
||||
resp = client.get(f"{host}/jobs/{job_id}/download")
|
||||
resp.raise_for_status()
|
||||
with open(out_path, "wb") as f:
|
||||
f.write(resp.content)
|
||||
except httpx.ConnectError:
|
||||
print(f"ERROR: Cannot connect to {host}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
except httpx.HTTPStatusError as e:
|
||||
print(f"ERROR: Cannot download video ({e.response.status_code}): {e.response.text}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
print(f"Video saved to {out_path}")
|
||||
|
||||
|
||||
def cmd_download(host: str, job_id: str, output: Optional[str] = None) -> None:
|
||||
download_video(host, job_id, output)
|
||||
|
||||
|
||||
def cmd_delete(host: str, job_id: str) -> None:
|
||||
try:
|
||||
with httpx.Client(timeout=10) as client:
|
||||
resp = client.delete(f"{host}/jobs/{job_id}")
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
except httpx.ConnectError:
|
||||
print(f"ERROR: Cannot connect to {host}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
except httpx.HTTPStatusError as e:
|
||||
if e.response.status_code == 404:
|
||||
print(f"Job {job_id} not found", file=sys.stderr)
|
||||
else:
|
||||
print(f"ERROR: Server error ({e.response.status_code})", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
print(f"Deleted job: {data['deleted']}")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = _build_parser()
|
||||
args = parser.parse_args()
|
||||
|
||||
if not args.command:
|
||||
parser.print_help()
|
||||
sys.exit(1)
|
||||
|
||||
if args.command == "health":
|
||||
cmd_health(args.host)
|
||||
elif args.command == "generate":
|
||||
cmd_generate(args)
|
||||
elif args.command == "status":
|
||||
cmd_status(args.host, args.job_id)
|
||||
elif args.command == "download":
|
||||
cmd_download(args.host, args.job_id, args.output)
|
||||
elif args.command == "delete":
|
||||
cmd_delete(args.host, args.job_id)
|
||||
Reference in New Issue
Block a user