Files

294 lines
10 KiB
Python

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)