294 lines
10 KiB
Python
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)
|