infinitetalk / wavespeed.py
chengzeyi's picture
Rebuild on shared BYOK template: user-supplied API key, v3 API, Python 3.11
2d02944
Raw
History Blame
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)