"""The encode path for streamed audio output: chunk bytes in, ADTS frames out.""" from __future__ import annotations import io import logging import os import subprocess import threading import time import wave from collections import deque from pydub import AudioSegment from gradio import processing_utils from gradio.data_classes import MediaStreamChunk logger = logging.getLogger(__name__) AAC_FRAME_SAMPLES = 1024 # The only sample rates an ADTS stream can declare. Anything else gets # resampled by the encoder, and a frame is 1024 samples at the rate that comes # out, not the one that went in, so the playlist has to be told which is which. ADTS_SAMPLE_RATES = ( 96000, 88200, 64000, 48000, 44100, 32000, 24000, 22050, 16000, 12000, 11025, 8000, 7350, ) def nearest_adts_rate(sample_rate: int) -> int: return min(ADTS_SAMPLE_RATES, key=lambda rate: abs(rate - sample_rate)) # How long the first chunk waits for the encoder process to come up, and how # long a later chunk that completes no frame waits before giving up on one. # See `AacStreamEncoder.take`. STARTUP_WAIT = 0.1 STEADY_WAIT = 0.005 def parse_adts_frames(data: bytes | bytearray) -> tuple[list[bytes], int]: """Split `data` on ADTS frame headers. Returns the complete frames and how many bytes of `data` they consumed, so a partial frame at the end can be carried over to the next read. """ frames: list[bytes] = [] offset = 0 while offset + 7 <= len(data): if data[offset] != 0xFF or data[offset + 1] & 0xF0 != 0xF0: offset += 1 continue length = ( ((data[offset + 3] & 0x03) << 11) | (data[offset + 4] << 3) | ((data[offset + 5] & 0xE0) >> 5) ) if length < 7: offset += 1 continue if offset + length > len(data): break frames.append(bytes(data[offset : offset + length])) offset += length return frames, offset def _read_wav_pcm(data: bytes) -> tuple[int, int, bytes] | None: """`(sample_rate, channels, pcm)` for 16-bit PCM wav bytes, else None.""" if data[:4] != b"RIFF" or data[8:12] != b"WAVE": return None try: with wave.open(io.BytesIO(data), "rb") as reader: if reader.getsampwidth() != 2 or reader.getcomptype() != "NONE": return None pcm = reader.readframes(reader.getnframes()) if not pcm: # A wav written for a stream often declares a `data` size of 0 # because the length is not known yet. Trusting it would drop # the chunk's audio silently; ffmpeg reads such a file to EOF. return None return reader.getframerate(), reader.getnchannels(), pcm except (wave.Error, EOFError): return None def _ffmpeg_decode(source: bytes | str, sample_rate: int, channels: int) -> bytes: data = source if isinstance(source, bytes) else None result = subprocess.run( [ "ffmpeg", "-v", "error", "-nostdin", "-i", "pipe:0" if data is not None else source, "-vn", "-f", "s16le", "-acodec", "pcm_s16le", "-ar", str(sample_rate), "-ac", str(channels), "pipe:1", ], input=data, capture_output=True, check=False, ) # fmt: skip if result.returncode != 0: raise processing_utils.ffmpeg_failed( "ffmpeg", result.returncode, result.stderr, "Decoding the streamed audio chunk", ) # Empty is a chunk that carries no samples, which a generator yields for a # tick that produced no audio; the encoder takes an empty write in stride. return result.stdout def decode_to_pcm( data: bytes, sample_rate: int | None = None, channels: int | None = None ) -> tuple[int, int, bytes]: """Decode one streamed chunk to signed 16-bit little-endian PCM. Chunks gradio wrote itself are 16-bit wav, which the stdlib reads with no subprocess at all. Everything else - a `bytes` yield in an unknown format, a non-wav `format=`, or a chunk whose parameters differ from the ones the stream started with - goes through one ffmpeg process, which resamples to `sample_rate` and `channels` on the way. """ parsed = _read_wav_pcm(data) if sample_rate is None or channels is None: if parsed is not None: return parsed segment = AudioSegment.from_file(io.BytesIO(data)).set_sample_width(2) return segment.frame_rate, segment.channels, segment.raw_data if parsed is not None and parsed[0] == sample_rate and parsed[1] == channels: return parsed return sample_rate, channels, _ffmpeg_decode(data, sample_rate, channels) def decode_file_to_pcm(path: str, sample_rate: int, channels: int) -> bytes: """Decode a file's audio track to PCM without reading it into memory first.""" return _ffmpeg_decode(path, sample_rate, channels) class AacStreamEncoder: """One ffmpeg process for the whole lifetime of a streamed output. An AAC frame only reconstructs once overlap-added with its neighbours, so an encoder started per chunk makes every chunk boundary a discontinuity that decodes as roughly 36 ms of near-silence. One encoder never creates it, and keeps the sub-frame remainder between writes, so chunks need not arrive in multiples of 1024 samples. """ def __init__( self, sample_rate: int, channels: int, operation: str = "Streaming audio output", ): # Encoding only, so ffprobe is not wanted here. processing_utils.require_ffmpeg(operation, "ffmpeg") self.sample_rate = sample_rate self.channels = channels # Asking for the resample rather than letting the encoder pick one: # `frame_duration` has to match what the frames actually carry, and it # is needed before the first frame exists to read it from. self.output_rate = nearest_adts_rate(sample_rate) self.process = subprocess.Popen( [ "ffmpeg", "-v", "error", "-nostdin", # Without these two, ffmpeg probes the input before emitting # anything and holds back roughly 64 KB of PCM, which is 2 # seconds of 16 kHz mono audio. The input format is fully # described below, so there is nothing to probe for. "-probesize", "32", "-analyzeduration", "0", "-f", "s16le", "-ar", str(sample_rate), "-ac", str(channels), "-i", "pipe:0", "-c:a", "aac", "-ar", str(self.output_rate), "-flush_packets", "1", "-f", "adts", "pipe:1", ], stdin=subprocess.PIPE, stdout=subprocess.PIPE, ) self._buffer = bytearray() self._ready: deque[bytes] = deque() self._at_eof = False self._waited_for_startup = False self._stdin_closed = False self._closed = False self._condition = threading.Condition() # A full stdout pipe blocks the encoder, and at 48 kHz stereo the pipe # holds only a third of a second of audio, so it has to be drained # continuously rather than between writes. self._reader = threading.Thread( target=self._read_loop, name="gradio-aac-encoder", daemon=True ) try: self._reader.start() except BaseException: # The process is already running; with no reader it never exits. self.process.kill() self.process.wait() for pipe in (self.process.stdin, self.process.stdout): if pipe is not None: pipe.close() raise @property def frame_duration(self) -> float: return AAC_FRAME_SAMPLES / self.output_rate def _read_loop(self) -> None: stdout = self.process.stdout assert stdout is not None # noqa: S101 # os.read rather than the buffered reader's read(), which would block # until the requested size is filled instead of returning what has # arrived, and works the same way on Windows. fd = stdout.fileno() while not self._closed: try: data = os.read(fd, 1 << 16) except (OSError, ValueError): break if not data: break with self._condition: self._buffer += data frames, consumed = parse_adts_frames(self._buffer) del self._buffer[:consumed] self._ready.extend(frames) self._condition.notify_all() with self._condition: self._at_eof = True self._condition.notify_all() def feed(self, pcm: bytes) -> None: """Write signed 16-bit little-endian PCM into the encoder. A no-op once `close()` has run: the stream was torn down under the chunk being encoded, and the run is over. """ if self._closed: return if self._stdin_closed: raise RuntimeError("encoder stdin is already closed") stdin = self.process.stdin assert stdin is not None # noqa: S101 try: stdin.write(pcm) stdin.flush() except (OSError, ValueError) as e: # ValueError is what a write to an already-closed pipe raises, which # is reachable because `close()` can land between the check above # and here. if self._closed: return raise RuntimeError( f"The audio encoder exited with code {self.process.poll()}." ) from e def take(self, timeout: float | None = None) -> list[bytes]: """Pop every whole frame the encoder has emitted so far. Waits `STARTUP_WAIT` for the encoder to come up, once. A later chunk that completes no frame waits `STEADY_WAIT` and leaves its audio for the next chunk or `flush()`; paying the startup wait per short chunk held 20 ms chunks at a quarter of real time. Do not wait for a predicted frame count instead: the prediction is sometimes one too high, and then every chunk it is wrong about pays the whole timeout. An explicit `timeout` overrides both, for the caller that cannot take no frames for an answer. """ with self._condition: if timeout is not None: wait_for = timeout else: wait_for = STEADY_WAIT if self._waited_for_startup else STARTUP_WAIT self._waited_for_startup = True deadline = time.monotonic() + wait_for while not self._ready and not self._at_eof: remaining = deadline - time.monotonic() if remaining <= 0: break self._condition.wait(remaining) frames = list(self._ready) self._ready.clear() self._raise_if_encoder_died() return frames def _raise_if_encoder_died(self) -> None: """A stream that stops growing silently is worse than a loud failure.""" if self._closed or not self._at_eof: return code = self.process.poll() if code is not None and code != 0: raise RuntimeError(f"The audio encoder exited with code {code}.") def flush(self, timeout: float = 5.0) -> list[bytes]: """Close the input, return what the encoder had left, and release it. Terminal: the encoder cannot be fed again afterwards. """ if not self._stdin_closed: self._stdin_closed = True if self.process.stdin is not None: try: self.process.stdin.close() except OSError: pass killed = False try: self.process.wait(timeout=timeout) except subprocess.TimeoutExpired: self.process.kill() self.process.wait() killed = True self._reader.join(timeout=1.0) with self._condition: frames = list(self._ready) self._ready.clear() code = self.process.returncode was_closed = self._closed # The process is reaped by now, so this is just the pipes, which a # caller that flushes and drops the encoder would otherwise leave to gc. self.close() if killed: logger.warning( "The audio encoder was still running %.0f s after its input " "ended and was killed; the last frames of the stream may be " "missing.", timeout, ) elif code and not was_closed: raise RuntimeError(f"The audio encoder exited with code {code}.") return frames def close(self) -> None: """Give up on the process without waiting for its remaining output.""" self._closed = True self._stdin_closed = True if self.process.poll() is None: self.process.kill() try: self.process.wait(timeout=2.0) except subprocess.TimeoutExpired: pass # The reader has to be done with the pipe's fd before the fd is # closed, since a closed fd number can be reused by another thread. self._reader.join(timeout=1.0) for pipe in (self.process.stdin, self.process.stdout): if pipe is not None: try: pipe.close() except OSError: pass class EncoderSlot: """A stream registry's entry, made before its encoder exists. The encoder is created on a worker thread, and the coroutine waiting for it can be cancelled without the thread being stopped, so the thread can go on to publish an encoder after the coroutine is gone. The two hand over under a lock: the thread attaches unless the slot has been ended, and ending the slot closes whatever is attached, whichever comes first. """ def __init__(self) -> None: self._lock = threading.Lock() self._ended = False self.encoder: AacStreamEncoder | None = None def attach(self, encoder: AacStreamEncoder) -> bool: with self._lock: if self._ended: return False self.encoder = encoder return True def detach(self) -> AacStreamEncoder | None: """Take the encoder out and refuse any that arrives later.""" with self._lock: self._ended = True encoder, self.encoder = self.encoder, None return encoder def end(self) -> None: encoder = self.detach() if encoder is not None: encoder.close() def segment_from_frames( encoder: AacStreamEncoder, frames: list[bytes] ) -> MediaStreamChunk | None: if not frames: return None return { "data": b"".join(frames), # Derived from the frame count rather than from the source chunk's # length, so the playlist's #EXTINF matches what the segment decodes to. "duration": len(frames) * encoder.frame_duration, "extension": ".aac", }