wan-image-to-video / wavespeed.py
chengzeyi's picture
Rebuild on shared BYOK template: user-supplied API key, v3 API, Python 3.11
40f9a0b
Raw
History Blame Contribute Delete
5.47 kB
"""Minimal WaveSpeed v3 API client.
Shared verbatim by every Space in the wavespeed org. Generated from
_shared/spaceapp/ — edit there and re-run _shared/build_apps.py, never edit the
copy inside a Space.
The user supplies their own API key through the UI. It is used to sign requests
to api.wavespeed.ai and nothing else: it is never logged, never written to
disk, never placed in a Gradio component value, and is stripped out of every
error message before that message can reach a browser (see `redact`). An
exception raised by `requests` can carry the full request headers in its text,
which is exactly how a key ends up in a user-visible traceback, so every raise
in this module goes through `redact` first.
API reference: https://wavespeed.ai/docs/rest-api
"""
from __future__ import annotations
import time
from typing import Any
import requests
API_BASE = "https://api.wavespeed.ai/api/v3"
UPLOAD_URL = f"{API_BASE}/media/upload/binary"
# The docs ask for >= 2s between polls of the same task, easing toward 5-10s
# for long jobs. Anything faster risks being throttled.
POLL_START = 2.0
POLL_MAX = 8.0
POLL_GROWTH = 1.25
POLL_TIMEOUT = 600
TERMINAL_OK = "completed"
TERMINAL_BAD = ("failed", "cancelled", "timeout")
class WaveSpeedError(Exception):
"""User-facing error. The message is always key-free."""
def redact(text: Any, key: str | None) -> str:
"""Remove the API key (and any bearer token) from text headed for a user."""
s = str(text)
if key:
k = key.strip()
if k:
s = s.replace(k, "***")
# Defend against a partially-quoted key in a repr.
if len(k) > 12:
s = s.replace(k[:12], "***")
# Catch any Authorization header echoed by a library.
import re
s = re.sub(r"(?i)(bearer\s+)[A-Za-z0-9._\-]+", r"\1***", s)
s = re.sub(r"(?i)('authorization':\s*')[^']*", r"\1***", s)
return s
def _headers(key: str, json: bool = False) -> dict:
h = {"Authorization": f"Bearer {key.strip()}"}
if json:
h["Content-Type"] = "application/json"
return h
def _check(resp: requests.Response, key: str) -> dict:
if resp.status_code == 401:
raise WaveSpeedError("Invalid API key. Check the key and try again.")
if resp.status_code == 402:
raise WaveSpeedError("This account is out of credit.")
if resp.status_code == 429:
raise WaveSpeedError("Rate limit or quota exceeded. Wait and retry.")
if resp.status_code >= 400:
raise WaveSpeedError(
redact(f"API error {resp.status_code}: {resp.text[:300]}", key)
)
try:
body = resp.json()
except ValueError:
raise WaveSpeedError("API returned a non-JSON response.") from None
if body.get("code") != 200:
raise WaveSpeedError(redact(body.get("message", "Unknown API error"), key))
return body.get("data", {}) or {}
def upload(key: str, path: str) -> str:
"""Upload a local file, returning the URL to reference it by."""
try:
with open(path, "rb") as fh:
resp = requests.post(
UPLOAD_URL, headers=_headers(key), files={"file": fh}, timeout=120
)
except requests.RequestException as e:
raise WaveSpeedError(redact(f"Upload failed: {e}", key)) from None
data = _check(resp, key)
url = data.get("download_url") or data.get("url")
if not url:
raise WaveSpeedError("Upload succeeded but returned no URL.")
return url
def submit(key: str, model: str, payload: dict) -> str:
"""Start a job and return its request id.
Deliberately not retried: the docs warn that repeating a POST can bill the
caller twice. Only the GET poll below is safe to retry.
"""
try:
resp = requests.post(
f"{API_BASE}/{model}", headers=_headers(key, json=True), json=payload, timeout=60
)
except requests.RequestException as e:
raise WaveSpeedError(redact(f"Could not reach the API: {e}", key)) from None
data = _check(resp, key)
rid = data.get("id")
if not rid:
raise WaveSpeedError("API accepted the request but returned no task id.")
return rid
def poll(key: str, request_id: str, on_tick=None) -> list[str]:
"""Poll a task to completion and return its output URLs."""
url = f"{API_BASE}/predictions/{request_id}/result"
deadline = time.time() + POLL_TIMEOUT
delay = POLL_START
while time.time() < deadline:
time.sleep(delay)
delay = min(delay * POLL_GROWTH, POLL_MAX)
try:
resp = requests.get(url, headers=_headers(key), timeout=60)
except requests.RequestException:
continue # transient; a GET is safe to repeat
data = _check(resp, key)
status = data.get("status", "")
if status == TERMINAL_OK:
outputs = data.get("outputs") or []
if not outputs:
raise WaveSpeedError("Generation finished but produced no output.")
return outputs
if status in TERMINAL_BAD:
detail = redact(data.get("error") or status, key)
raise WaveSpeedError(f"Generation {status}: {detail}")
if on_tick:
on_tick(status)
raise WaveSpeedError("Timed out waiting for the result. The job may still finish.")
def run(key: str, model: str, payload: dict, on_tick=None) -> list[str]:
return poll(key, submit(key, model, payload), on_tick=on_tick)