File size: 5,465 Bytes
40f9a0b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
"""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)