| """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" |
|
|
| |
| |
| 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, "***") |
| |
| if len(k) > 12: |
| s = s.replace(k[:12], "***") |
| |
| 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 |
| 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) |
|
|