File size: 17,050 Bytes
3cc5b15
 
 
 
 
 
 
 
 
 
 
 
 
 
5da7b89
 
 
 
 
 
 
 
 
 
 
 
3cc5b15
 
 
 
 
 
 
 
5da7b89
3cc5b15
 
 
 
 
5da7b89
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8a9a9dd
 
 
 
5da7b89
 
 
 
3cc5b15
 
 
 
 
 
 
160614f
7a57c31
 
 
 
 
 
 
 
160614f
7a57c31
 
160614f
 
 
 
 
 
 
 
3cc5b15
f3537fe
 
 
 
 
 
 
 
3cc5b15
 
 
f3537fe
 
3cc5b15
f3537fe
 
3cc5b15
f3537fe
 
 
 
 
 
 
 
3cc5b15
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
160614f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7a57c31
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
160614f
7a57c31
 
 
 
160614f
 
 
 
 
 
 
 
 
 
 
 
 
 
7a57c31
 
 
 
 
 
 
 
 
 
 
 
 
160614f
7a57c31
160614f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3cc5b15
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
160614f
3cc5b15
 
160614f
 
 
 
 
 
 
3cc5b15
 
 
 
 
 
 
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
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
"""Local audio-gen server β€” reference example for audio-brief's "Local server" model.

When you select "Local server (your machine)" in audio-brief's Generate tab,
your browser POSTs `{prompt, duration}` to `<this server>/generate` and
expects audio bytes back (WAV or MP3). The audio never leaves your machine
on its way to a remote gen provider β€” the HF Spaces backend only receives
the resulting audio (uploaded by the browser, not the server).

Run me:
    pip install fastapi uvicorn
    # plus whatever your gen backend needs (mlx-audio, diffusers, etc.)
    python local-gen-server-example.py
    # serves at http://localhost:7864

CORS + Private Network Access:
    The browser fetch comes from `https://<your-space>.hf.space` β€” a
    different origin from `http://localhost:7864`. The server must return
    permissive CORS headers AND the PNA-Allow header. The middleware
    below does both. Without the PNA header the browser silently fails
    the request with `TypeError: Failed to fetch` β€” even though Chrome
    explicitly exempts http://localhost from mixed-content blocking.

Mixed content:
    Chrome / Brave / Edge allow http://localhost from HTTPS pages
    (special exception). Safari is stricter β€” if Safari blocks, run
    this server with a self-signed cert on HTTPS, or test in Chrome.
"""
from __future__ import annotations

import io
import os
import subprocess
from pathlib import Path

from fastapi import FastAPI, Request
from fastapi.responses import Response
from pydantic import BaseModel

app = FastAPI()


# CORS + Private Network Access in one explicit middleware.
#
# Chrome enforces Private Network Access (PNA) on cross-origin requests
# from public origins (e.g. an HTTPS HF Space) to private-network IPs
# (localhost, 127.0.0.1, 10/8, 192.168/16). It fires a CORS-style
# preflight OPTIONS with header `Access-Control-Request-Private-Network:
# true`, and the local server MUST answer with
# `Access-Control-Allow-Private-Network: true` or the browser refuses
# the real request with TypeError: Failed to fetch.
#
# FastAPI's CORSMiddleware does NOT add the PNA-Allow header, so we
# replace it with this hand-rolled middleware that returns both. We
# answer OPTIONS preflights directly (204, no body) so the SA3 binary
# never spins up on a preflight.
#
# Spec: https://wicg.github.io/private-network-access/
@app.middleware("http")
async def cors_and_pna(request: Request, call_next):
    if request.method == "OPTIONS":
        response = Response(status_code=204)
    else:
        response = await call_next(request)

    origin = request.headers.get("origin", "*")
    response.headers["Access-Control-Allow-Origin"] = origin
    response.headers["Vary"] = "Origin"
    response.headers["Access-Control-Allow-Methods"] = "GET, POST, OPTIONS"
    # CORS header matching is case-insensitive per the Fetch spec, but
    # some browser builds/extensions have been observed comparing
    # case-sensitively. Cheap to widen the allowlist.
    response.headers["Access-Control-Allow-Headers"] = "Content-Type, content-type"
    # The PNA-Allow header is the load-bearing one β€” without it Chrome
    # silently fails the request before any body is sent. Keep this on.
    response.headers["Access-Control-Allow-Private-Network"] = "true"
    return response


class GenRequest(BaseModel):
    prompt: str
    duration: int = 15


class TransformRequest(BaseModel):
    """Audio-to-audio (init-audio) gen. The browser POSTs from the Space
    page; it sends EITHER `audio_url` (preferred β€” the Space's public
    `/gradio_api/file=…` URL, which the bridge fetches itself) OR
    `audio_b64` (fallback β€” base64 of the raw bytes). audio_url is
    strongly preferred: avoids a 1.4Γ— base64 inflation across the JSON
    boundary and keeps the bridge stateless about who couriered the
    bytes. Bridge converts to SA3's required WAV format (44.1 kHz /
    16-bit PCM) via ffmpeg either way, then invokes SA3 --init-audio."""
    prompt: str
    audio_url: str = ""          # preferred: URL the bridge fetches itself
    audio_b64: str = ""          # fallback: base64 of source audio bytes
    duration: int = 15
    init_strength: float = 0.7   # β†’ --init-noise-level. 0.4-0.8 typical;
                                 # 1.0 = full regen (init ignored).
    steps: int = 8               # β†’ --steps. SA3 default; sweet spot.
    cfg: float = 1.0             # β†’ --cfg. >1.0 enables negative_prompt.
    negative_prompt: str = ""    # β†’ --negative-prompt. Ignored when cfg==1.0.


def run_mlx_sa3(prompt: str, duration: int) -> bytes:
    """Shell out to the MLX Stable Audio 3 Small-Music binary at
    `~/sa3_mlx/sa3`. CLI verified against the pj-battle engine
    (engine/audio-providers.js:302). Output is WAV.

    Default flags: `--dit sm-music --decoder same-s`. These match the
    engine's "small music" defaults β€” good quality, fits in 8 GB RAM.
    Spawn-per-request only; never two SA3 processes at once on this
    machine (each holds ~2 GB)."""
    binary = os.environ.get("SA3_BIN", str(Path.home() / "sa3_mlx" / "sa3"))
    if not Path(binary).exists():
        raise RuntimeError(
            f"SA3_BIN not found at {binary}. Expected ~/sa3_mlx/sa3 β€” "
            "see ~/sa3_mlx/README.md for install."
        )
    out_path = Path("/tmp") / f"local-gen-{os.getpid()}-{duration}.wav"
    # Match the engine's CLI: --dit sm-music --decoder same-s --seconds N
    subprocess.run(
        [
            binary,
            "--prompt", prompt,
            "--dit", "sm-music",
            "--decoder", "same-s",
            "--seconds", str(duration),
            "--out", str(out_path),
        ],
        check=True, timeout=300,
    )
    return out_path.read_bytes()


def run_stub_sine(prompt: str, duration: int) -> bytes:
    """Fallback: render a single sine tone so you can verify the bridge
    works end-to-end before plugging in a real gen backend."""
    import math, struct, wave
    sr = 44100
    n_samples = int(sr * duration)
    freq = 440.0
    buf = io.BytesIO()
    with wave.open(buf, "wb") as w:
        w.setnchannels(1)
        w.setsampwidth(2)
        w.setframerate(sr)
        for i in range(n_samples):
            v = int(20000 * math.sin(2.0 * math.pi * freq * i / sr))
            w.writeframesraw(struct.pack("<h", v))
    return buf.getvalue()


def _convert_to_sa3_wav(audio_bytes: bytes, out_wav: Path) -> None:
    """SA3's --init-audio requires WAV at 44.1 kHz, 16-bit PCM, mono OR
    stereo. The caller may post any container ffmpeg can read (mp3, wav,
    flac, m4a β€” Pollinations gens are mp3, our /generate output is wav,
    user uploads could be anything). Always re-encode via ffmpeg so SA3
    gets the exact format it expects.

    Why not soundfile-only: soundfile can read wav/flac but NOT mp3.
    Adding a graceful-fallback codepath makes the bridge more useful to
    less-technical users than a 400 'wav-only' rejection.

    ffmpeg must be installed (Homebrew: `brew install ffmpeg`)."""
    import shutil
    ffmpeg = shutil.which("ffmpeg")
    if not ffmpeg:
        raise RuntimeError(
            "ffmpeg not found on PATH. /transform needs it to convert "
            "the source audio into SA3's required WAV format. "
            "Install with `brew install ffmpeg` (or your distro equivalent)."
        )
    in_tmp = Path("/tmp") / f"sa3-in-{os.getpid()}.bin"
    in_tmp.write_bytes(audio_bytes)
    try:
        # -y                 overwrite output
        # -i in              read whatever the bytes are
        # -ar 44100          resample to SA3's required rate
        # -acodec pcm_s16le  16-bit PCM as required
        # -ac 2              force stereo (SA3 accepts mono OR stereo;
        #                    stereo is the safer default for music tiles)
        subprocess.run(
            [ffmpeg, "-y", "-loglevel", "error",
             "-i", str(in_tmp),
             "-ar", "44100",
             "-acodec", "pcm_s16le",
             "-ac", "2",
             str(out_wav)],
            check=True, capture_output=True, text=True, timeout=30,
        )
    except subprocess.CalledProcessError as e:
        raise RuntimeError(
            f"ffmpeg failed to convert source audio "
            f"(exit {e.returncode}): {(e.stderr or '')[-300:]}"
        ) from None
    finally:
        try:
            in_tmp.unlink()
        except FileNotFoundError:
            pass


# Default origin allowlist for audio_url. Override with the env var
# ALLOWED_AUDIO_ORIGINS (comma-separated, case-insensitive) when running
# against a forked / private Space β€” e.g.
#   ALLOWED_AUDIO_ORIGINS="myname-myfork.hf.space,myname.hf.space"
_DEFAULT_AUDIO_ORIGINS = (
    "kalamishere-audio-brief.hf.space",
    # *.hf.space + the canonical wrapper aren't included by default β€”
    # tighter is safer. Users running forks should set the env var.
)
_AUDIO_URL_PATH_PREFIX = "/gradio_api/file="


def _fetch_audio_url(url: str) -> bytes:
    """Download source audio from the Space's file-proxy URL.

    Security perimeter β€” only fetches URLs that pass ALL of:
      β€’ scheme = https
      β€’ host is in ALLOWED_AUDIO_ORIGINS env (or built-in default)
      β€’ path starts with `/gradio_api/file=`
      β€’ response size <= 100 MB

    Anything else is rejected with 400 before the bridge touches the
    network. Without this gate, /transform would be a generic
    request-forwarder happy to fetch http://internal-router/admin or
    file:// URLs the FastAPI process can read.
    """
    import urllib.request
    import urllib.parse
    from urllib.parse import urlparse

    parsed = urlparse(url)
    if parsed.scheme != "https":
        raise RuntimeError(
            f"audio_url must be https (got scheme '{parsed.scheme}')")
    if not parsed.path.startswith(_AUDIO_URL_PATH_PREFIX):
        raise RuntimeError(
            f"audio_url path must start with {_AUDIO_URL_PATH_PREFIX!r}; "
            f"got {parsed.path[:80]!r}")
    env_allow = os.environ.get("ALLOWED_AUDIO_ORIGINS", "").strip()
    if env_allow:
        allowed = tuple(h.strip().lower() for h in env_allow.split(",") if h.strip())
    else:
        allowed = _DEFAULT_AUDIO_ORIGINS
    host = (parsed.netloc or "").lower()
    # Strip any :port the user might have included.
    if ":" in host:
        host = host.split(":", 1)[0]
    if host not in allowed:
        raise RuntimeError(
            f"audio_url host '{host}' not in allowlist {allowed}. "
            "Override via ALLOWED_AUDIO_ORIGINS env var.")
    # Cap download size β€” 100 MB is a generous ceiling for an audio
    # tile (a 180s 44.1k stereo WAV is ~60 MB; mp3 is <10 MB).
    MAX = 100 * 1024 * 1024
    req = urllib.request.Request(
        url, headers={"User-Agent": "audio-brief-local-bridge/1.0"})
    with urllib.request.urlopen(req, timeout=30) as r:
        # Some content-length headers are missing or lying β€” read up to
        # MAX+1 and reject if we hit the ceiling.
        data = r.read(MAX + 1)
    if len(data) > MAX:
        raise RuntimeError(f"audio_url response exceeded {MAX} bytes cap")
    if not data:
        raise RuntimeError("audio_url returned empty body")
    return data


def run_mlx_sa3_transform(req: "TransformRequest") -> bytes:
    """Audio-to-audio via SA3 init-audio. Resolves the source bytes
    (URL-fetch preferred over b64-payload), re-encodes to WAV at
    44.1 kHz / 16-bit PCM via ffmpeg, then shells out to ~/sa3_mlx/sa3
    with --init-audio + --init-noise-level. WAV bytes back.

    init_strength is exposed as --init-noise-level in [0.01, 1.0]:
        0.4-0.6  Β· subtle variation, mostly preserves source
        0.7      Β· sweet spot for noticeable but recognisable transform
        0.8-1.0  Β· heavy regeneration; init becomes a loose suggestion
    """
    import base64
    binary = os.environ.get("SA3_BIN", str(Path.home() / "sa3_mlx" / "sa3"))
    if not Path(binary).exists():
        raise RuntimeError(
            f"SA3_BIN not found at {binary}. Expected ~/sa3_mlx/sa3 β€” "
            "see ~/sa3_mlx/README.md for install."
        )

    # Resolve the source bytes. URL preferred when present β€” saves a
    # base64 round-trip across the browser β†’ bridge JSON boundary.
    audio_url = (req.audio_url or "").strip()
    if audio_url:
        audio_bytes = _fetch_audio_url(audio_url)
    elif req.audio_b64:
        try:
            audio_bytes = base64.b64decode(req.audio_b64)
        except Exception as e:
            raise RuntimeError(f"audio_b64 decode failed: {e}") from None
    else:
        raise RuntimeError(
            "neither audio_url nor audio_b64 was provided")
    if not audio_bytes:
        raise RuntimeError("source audio resolved to empty bytes")

    init_wav = Path("/tmp") / f"sa3-init-{os.getpid()}.wav"
    out_wav  = Path("/tmp") / f"sa3-transform-{os.getpid()}.wav"
    _convert_to_sa3_wav(audio_bytes, init_wav)

    # Clamp init_strength to SA3's valid range [0.01, ∞). 1.0 is the
    # "ignore init" boundary; above 1.0 is allowed by SA3 but rarely
    # useful, so cap at 1.0 to avoid surprising users.
    init_noise = max(0.01, min(1.0, float(req.init_strength)))
    steps = max(1, int(req.steps))
    cfg = max(0.0, float(req.cfg))

    args = [
        binary,
        "--prompt", req.prompt or "",
        "--init-audio", str(init_wav),
        "--init-noise-level", str(init_noise),
        "--dit", "sm-music",
        "--decoder", "same-s",
        "--seconds", str(int(req.duration)),
        "--steps", str(steps),
        "--out", str(out_wav),
    ]
    # CFG only enables a second uncond branch when > 1.0 (SA3 skips it
    # otherwise). negative_prompt is only useful when CFG kicks in.
    if cfg != 1.0:
        args += ["--cfg", str(cfg)]
        if req.negative_prompt:
            args += ["--negative-prompt", req.negative_prompt]

    try:
        subprocess.run(args, check=True, timeout=600)
    except subprocess.CalledProcessError as e:
        raise RuntimeError(f"SA3 transform exit {e.returncode}") from None
    finally:
        try:
            init_wav.unlink()
        except FileNotFoundError:
            pass

    try:
        return out_wav.read_bytes()
    finally:
        try:
            out_wav.unlink()
        except FileNotFoundError:
            pass


@app.post("/transform")
def transform(req: TransformRequest):
    """Audio-to-audio transform. Browser POSTs base64'd source audio +
    prompt; server returns WAV bytes via SA3's --init-audio mode. See
    TransformRequest for the request shape.

    Only the mlx-sa3 backend supports this β€” the stub sine has no
    audio-input mode. Returns 503 when LOCAL_GEN_BACKEND != mlx-sa3
    so the UI can surface the right reason instead of a vague 500."""
    backend = os.environ.get("LOCAL_GEN_BACKEND", "stub").lower()
    if backend != "mlx-sa3":
        return Response(
            content=(f"transform requires LOCAL_GEN_BACKEND=mlx-sa3 "
                     f"(currently '{backend}'). Restart the bridge with "
                     "LOCAL_GEN_BACKEND=mlx-sa3 python3 ...").encode(),
            status_code=503,
        )
    try:
        audio_bytes = run_mlx_sa3_transform(req)
    except Exception as e:
        return Response(
            content=f"local transform failed: {e}".encode(),
            status_code=500,
        )
    return Response(content=audio_bytes, media_type="audio/wav")


@app.post("/generate")
def generate(req: GenRequest):
    """Generate audio from a prompt. Returns audio/wav (or audio/mpeg if
    your backend produces MP3). The browser uploads these bytes to
    audio-brief's HF Spaces backend for crate-add + analysis."""
    # Pick your backend:
    backend = os.environ.get("LOCAL_GEN_BACKEND", "stub").lower()
    try:
        if backend == "mlx-sa3":
            audio_bytes = run_mlx_sa3(req.prompt, req.duration)
        else:
            # Default to the stub sine so first-time users can verify the
            # bridge before wiring real gen.
            audio_bytes = run_stub_sine(req.prompt, req.duration)
    except Exception as e:
        return Response(content=f"local gen failed: {e}".encode(), status_code=500)

    media_type = "audio/wav"  # change to "audio/mpeg" if your backend yields MP3
    return Response(content=audio_bytes, media_type=media_type)


@app.get("/")
def root():
    backend = os.environ.get("LOCAL_GEN_BACKEND", "stub").lower()
    return {
        "service": "audio-brief local gen example",
        "endpoints": [
            "POST /generate  {prompt, duration}                            β†’ WAV",
            "POST /transform {prompt, audio_b64, duration, init_strength,"
            " steps, cfg, negative_prompt}  β†’ WAV  (mlx-sa3 backend only)",
        ],
        "backend": backend,
        "transform_supported": backend == "mlx-sa3",
    }


if __name__ == "__main__":
    import uvicorn
    port = int(os.environ.get("PORT", 7864))
    uvicorn.run("local-gen-server-example:app", host="127.0.0.1", port=port, reload=False)