Skip to main content

koan_core/audio/
streaming.rs

1//! StreamingSource — a Read+Seek adapter over a shared, incrementally-filled byte buffer.
2//!
3//! The download thread writes chunks into a `StreamBuffer` while the Symphonia decoder
4//! reads from a `StreamingSource` backed by the same buffer. The source blocks briefly
5//! when the read position catches up to the write head, enabling true streaming decode
6//! without waiting for the full download to complete.
7
8use std::io::{self, Read, Seek, SeekFrom};
9use std::sync::{Arc, Condvar, Mutex};
10use std::time::Duration;
11
12/// Longest a read may block waiting for bytes before the stream counts as dead.
13/// A download that stops advancing must surface as an error, not park the decode
14/// thread forever holding the ring buffer producer and the whole buffered track.
15const READ_TIMEOUT: Duration = Duration::from_secs(30);
16
17/// State shared between the download writer and the decoder reader.
18struct Inner {
19    data: Vec<u8>,
20    /// Total expected byte length. `None` if not yet known (no Content-Length).
21    total_len: Option<u64>,
22    /// Set to true when the download thread has finished cleanly — the buffer
23    /// holds the whole source and readers see EOF past its end.
24    done: bool,
25    /// Set to true when the download died before delivering everything.
26    /// Reads past the buffered bytes fail rather than reporting EOF.
27    failed: bool,
28    /// How long a read blocks for new bytes before giving up.
29    read_timeout: Duration,
30}
31
32/// A shared, growable byte buffer that the download thread writes into.
33///
34/// Clone it to get additional handles; all clones share the same underlying data.
35#[derive(Clone)]
36pub struct StreamBuffer {
37    inner: Arc<(Mutex<Inner>, Condvar)>,
38}
39
40impl StreamBuffer {
41    /// Create a new empty buffer. `total_len` may be provided once Content-Length is known.
42    pub fn new(total_len: Option<u64>) -> Self {
43        Self::with_read_timeout(total_len, READ_TIMEOUT)
44    }
45
46    fn with_read_timeout(total_len: Option<u64>, read_timeout: Duration) -> Self {
47        Self {
48            inner: Arc::new((
49                Mutex::new(Inner {
50                    data: Vec::new(),
51                    total_len,
52                    done: false,
53                    failed: false,
54                    read_timeout,
55                }),
56                Condvar::new(),
57            )),
58        }
59    }
60
61    /// Append downloaded bytes. Called by the download thread.
62    pub fn push(&self, chunk: &[u8]) {
63        let (lock, cvar) = &*self.inner;
64        let mut inner = lock.lock().unwrap();
65        inner.data.extend_from_slice(chunk);
66        cvar.notify_all();
67    }
68
69    /// Signal that the download delivered everything. Readers see EOF past the
70    /// buffered bytes.
71    pub fn finish(&self) {
72        let (lock, cvar) = &*self.inner;
73        let mut inner = lock.lock().unwrap();
74        inner.done = true;
75        cvar.notify_all();
76    }
77
78    /// Signal that the download died. Readers past the buffered bytes get a
79    /// broken-pipe error — reporting EOF here would silently truncate the track.
80    pub fn fail(&self) {
81        let (lock, cvar) = &*self.inner;
82        let mut inner = lock.lock().unwrap();
83        inner.failed = true;
84        cvar.notify_all();
85    }
86
87    /// True when this is the last handle: nothing can read what is written from
88    /// here on, so a writer holding it has no reason to keep going.
89    pub fn is_abandoned(&self) -> bool {
90        Arc::strong_count(&self.inner) == 1
91    }
92
93    /// Total bytes received so far.
94    pub fn bytes_downloaded(&self) -> u64 {
95        let (lock, _) = &*self.inner;
96        lock.lock().unwrap().data.len() as u64
97    }
98
99    /// Total expected length (from Content-Length), if known.
100    pub fn total_len(&self) -> Option<u64> {
101        let (lock, _) = &*self.inner;
102        lock.lock().unwrap().total_len
103    }
104
105    /// Create a `StreamingSource` that reads from this buffer starting at offset 0.
106    pub fn reader(&self) -> StreamingSource {
107        StreamingSource {
108            inner: self.inner.clone(),
109            pos: 0,
110        }
111    }
112}
113
114/// A `Read + Seek` view into a `StreamBuffer`.
115///
116/// Blocks on `read` when the read position is at or beyond the write head,
117/// until more bytes arrive or the download finishes.
118pub struct StreamingSource {
119    inner: Arc<(Mutex<Inner>, Condvar)>,
120    pos: u64,
121}
122
123impl Read for StreamingSource {
124    fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
125        if buf.is_empty() {
126            return Ok(0);
127        }
128
129        let (lock, cvar) = &*self.inner;
130
131        // Wait until there is data at `pos`, or the download ends one way or another.
132        let guard = lock.lock().unwrap();
133        let timeout = guard.read_timeout;
134        let (inner, wait) = cvar
135            .wait_timeout_while(guard, timeout, |s| {
136                s.data.len() as u64 <= self.pos && !s.done && !s.failed
137            })
138            .unwrap();
139
140        let available = inner.data.len() as u64;
141        if available <= self.pos {
142            if inner.failed {
143                return Err(io::Error::new(
144                    io::ErrorKind::BrokenPipe,
145                    "stream download failed before delivering the whole track",
146                ));
147            }
148            if wait.timed_out() {
149                return Err(io::Error::new(
150                    io::ErrorKind::TimedOut,
151                    "stream download stalled",
152                ));
153            }
154            // Done and no more data — EOF.
155            return Ok(0);
156        }
157
158        let start = self.pos as usize;
159        let end = (start + buf.len()).min(inner.data.len());
160        let n = end - start;
161        buf[..n].copy_from_slice(&inner.data[start..end]);
162        self.pos += n as u64;
163        Ok(n)
164    }
165}
166
167impl Seek for StreamingSource {
168    fn seek(&mut self, pos: SeekFrom) -> io::Result<u64> {
169        let (lock, _) = &*self.inner;
170        let inner = lock.lock().unwrap();
171
172        let new_pos: i64 = match pos {
173            SeekFrom::Start(n) => n as i64,
174            SeekFrom::Current(n) => self.pos as i64 + n,
175            SeekFrom::End(n) => {
176                // For End seeks we need total_len. If not known yet, use current data len.
177                let len = inner.total_len.unwrap_or(inner.data.len() as u64) as i64;
178                len + n
179            }
180        };
181
182        if new_pos < 0 {
183            return Err(io::Error::new(
184                io::ErrorKind::InvalidInput,
185                "seek before beginning of stream",
186            ));
187        }
188
189        self.pos = new_pos as u64;
190        Ok(self.pos)
191    }
192}
193
194// Symphonia requires MediaSource: Read + Seek + Send + Any
195impl symphonia::core::io::MediaSource for StreamingSource {
196    fn is_seekable(&self) -> bool {
197        // Seekable only if the total length is known (needed for seek-to-end math).
198        // Forward seeks always work; backward seeks require buffered data already present.
199        // We advertise seekable=true and handle backward seeks via the buffered Vec.
200        true
201    }
202
203    fn byte_len(&self) -> Option<u64> {
204        let (lock, _) = &*self.inner;
205        lock.lock().unwrap().total_len
206    }
207}
208
209#[cfg(test)]
210mod tests {
211    use std::io::{self, Read, Seek, SeekFrom};
212
213    use super::*;
214
215    fn filled_buffer(data: &[u8]) -> StreamBuffer {
216        let buf = StreamBuffer::new(Some(data.len() as u64));
217        buf.push(data);
218        buf.finish();
219        buf
220    }
221
222    #[test]
223    fn new_buffer_starts_empty() {
224        let buf = StreamBuffer::new(Some(1024));
225        assert_eq!(buf.bytes_downloaded(), 0);
226        assert_eq!(buf.total_len(), Some(1024));
227    }
228
229    #[test]
230    fn new_buffer_unknown_total() {
231        let buf = StreamBuffer::new(None);
232        assert_eq!(buf.total_len(), None);
233    }
234
235    #[test]
236    fn read_all_data_available() {
237        let data = b"hello streaming world";
238        let buf = filled_buffer(data);
239        let mut src = buf.reader();
240
241        let mut out = Vec::new();
242        src.read_to_end(&mut out).unwrap();
243        assert_eq!(out, data);
244    }
245
246    #[test]
247    fn read_partial_then_rest() {
248        let data = b"abcdefghij";
249        let buf = filled_buffer(data);
250        let mut src = buf.reader();
251
252        let mut first = [0u8; 4];
253        let n = src.read(&mut first).unwrap();
254        assert_eq!(n, 4);
255        assert_eq!(&first, b"abcd");
256
257        let mut rest = Vec::new();
258        src.read_to_end(&mut rest).unwrap();
259        assert_eq!(rest, b"efghij");
260    }
261
262    #[test]
263    fn seek_from_start() {
264        let data = b"0123456789";
265        let buf = filled_buffer(data);
266        let mut src = buf.reader();
267
268        let pos = src.seek(SeekFrom::Start(5)).unwrap();
269        assert_eq!(pos, 5);
270
271        let mut out = [0u8; 3];
272        src.read_exact(&mut out).unwrap();
273        assert_eq!(&out, b"567");
274    }
275
276    #[test]
277    fn seek_from_current() {
278        let data = b"0123456789";
279        let buf = filled_buffer(data);
280        let mut src = buf.reader();
281
282        src.seek(SeekFrom::Start(2)).unwrap();
283        let pos = src.seek(SeekFrom::Current(3)).unwrap();
284        assert_eq!(pos, 5);
285
286        let mut out = [0u8; 2];
287        src.read_exact(&mut out).unwrap();
288        assert_eq!(&out, b"56");
289    }
290
291    #[test]
292    fn seek_from_end() {
293        let data = b"0123456789";
294        let buf = filled_buffer(data);
295        let mut src = buf.reader();
296
297        // SeekFrom::End(0) should position at total_len (EOF).
298        let pos = src.seek(SeekFrom::End(0)).unwrap();
299        assert_eq!(pos, 10);
300
301        // SeekFrom::End(-3) should position at offset 7.
302        let pos = src.seek(SeekFrom::End(-3)).unwrap();
303        assert_eq!(pos, 7);
304
305        let mut out = [0u8; 3];
306        src.read_exact(&mut out).unwrap();
307        assert_eq!(&out, b"789");
308    }
309
310    #[test]
311    fn seek_before_start_errors() {
312        let data = b"hello";
313        let buf = filled_buffer(data);
314        let mut src = buf.reader();
315
316        let result = src.seek(SeekFrom::Current(-1));
317        assert!(result.is_err());
318    }
319
320    #[test]
321    fn is_complete_when_done() {
322        let buf = StreamBuffer::new(Some(5));
323        buf.push(b"hello");
324        // Not yet finished.
325        assert!(buf.bytes_downloaded() != 0); // just check bytes_downloaded works
326        buf.finish();
327        // After finish, a reader should see EOF immediately.
328        let mut src = buf.reader();
329        let mut out = Vec::new();
330        src.read_to_end(&mut out).unwrap();
331        assert_eq!(out, b"hello");
332    }
333
334    #[test]
335    fn byte_len_returns_total() {
336        use symphonia::core::io::MediaSource;
337        let buf = StreamBuffer::new(Some(42));
338        let src = buf.reader();
339        assert_eq!(src.byte_len(), Some(42));
340    }
341
342    #[test]
343    fn is_seekable_true() {
344        use symphonia::core::io::MediaSource;
345        let buf = StreamBuffer::new(None);
346        let src = buf.reader();
347        assert!(src.is_seekable());
348    }
349
350    #[test]
351    fn push_increments_bytes_downloaded() {
352        let buf = StreamBuffer::new(Some(10));
353        buf.push(b"hello");
354        assert_eq!(buf.bytes_downloaded(), 5);
355        buf.push(b"world");
356        assert_eq!(buf.bytes_downloaded(), 10);
357    }
358
359    #[test]
360    fn multiple_readers_independent_positions() {
361        let data = b"0123456789";
362        let buf = filled_buffer(data);
363
364        let mut r1 = buf.reader();
365        let mut r2 = buf.reader();
366
367        r1.seek(SeekFrom::Start(7)).unwrap();
368
369        let mut out1 = [0u8; 3];
370        r1.read_exact(&mut out1).unwrap();
371        assert_eq!(&out1, b"789");
372
373        let mut out2 = [0u8; 3];
374        r2.read_exact(&mut out2).unwrap();
375        assert_eq!(&out2, b"012");
376    }
377
378    #[test]
379    fn failed_download_errors_instead_of_reporting_eof() {
380        let buf = StreamBuffer::new(Some(1000));
381        buf.push(b"partial");
382        buf.fail();
383
384        let mut src = buf.reader();
385        let mut out = [0u8; 7];
386        src.read_exact(&mut out).unwrap();
387        assert_eq!(&out, b"partial");
388
389        // Past the buffered bytes: an error, never a clean EOF — Ok(0) here
390        // would end the track early and look like a short file.
391        let err = src.read(&mut out).unwrap_err();
392        assert_eq!(err.kind(), io::ErrorKind::BrokenPipe);
393    }
394
395    #[test]
396    fn failed_download_wakes_a_blocked_reader() {
397        let buf = StreamBuffer::new(Some(1000));
398        let mut src = buf.reader();
399
400        let writer = buf.clone();
401        let waiter = std::thread::spawn(move || {
402            let mut out = [0u8; 8];
403            src.read(&mut out).map(|_| ()).map_err(|e| e.kind())
404        });
405
406        std::thread::sleep(std::time::Duration::from_millis(20));
407        writer.fail();
408
409        assert_eq!(waiter.join().unwrap(), Err(io::ErrorKind::BrokenPipe));
410    }
411
412    #[test]
413    fn stalled_download_times_out() {
414        // A download with a Content-Length that never arrives: the read must
415        // give up rather than park the decode thread forever.
416        let buf = StreamBuffer::with_read_timeout(Some(1000), std::time::Duration::from_millis(20));
417        let mut src = buf.reader();
418        let mut out = [0u8; 8];
419        let err = src.read(&mut out).unwrap_err();
420        assert_eq!(err.kind(), io::ErrorKind::TimedOut);
421    }
422
423    #[test]
424    fn abandoned_when_no_other_handles_remain() {
425        let buf = StreamBuffer::new(None);
426        assert!(buf.is_abandoned());
427
428        let reader = buf.reader();
429        assert!(!buf.is_abandoned());
430
431        drop(reader);
432        assert!(buf.is_abandoned());
433    }
434
435    #[test]
436    fn test_partial_availability() {
437        // Push only 500 bytes without finish() — simulates in-progress download.
438        let data: Vec<u8> = (0u8..=255).cycle().take(1000).collect();
439        let buf = StreamBuffer::new(Some(1000));
440        buf.push(&data[..500]);
441
442        assert_eq!(buf.bytes_downloaded(), 500);
443        assert_eq!(buf.total_len(), Some(1000));
444
445        // Spawn thread to call finish() after brief delay so reader doesn't block forever.
446        let buf2 = buf.clone();
447        std::thread::spawn(move || {
448            std::thread::sleep(std::time::Duration::from_millis(10));
449            buf2.finish();
450        });
451
452        let mut src = buf.reader();
453        let mut out = Vec::new();
454        src.read_to_end(&mut out).unwrap();
455        assert_eq!(out.len(), 500);
456        assert_eq!(out, &data[..500]);
457    }
458}