koan_core/audio/
streaming.rs1use std::io::{self, Read, Seek, SeekFrom};
9use std::sync::{Arc, Condvar, Mutex};
10use std::time::Duration;
11
12const READ_TIMEOUT: Duration = Duration::from_secs(30);
16
17struct Inner {
19 data: Vec<u8>,
20 total_len: Option<u64>,
22 done: bool,
25 failed: bool,
28 read_timeout: Duration,
30}
31
32#[derive(Clone)]
36pub struct StreamBuffer {
37 inner: Arc<(Mutex<Inner>, Condvar)>,
38}
39
40impl StreamBuffer {
41 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 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 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 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 pub fn is_abandoned(&self) -> bool {
90 Arc::strong_count(&self.inner) == 1
91 }
92
93 pub fn bytes_downloaded(&self) -> u64 {
95 let (lock, _) = &*self.inner;
96 lock.lock().unwrap().data.len() as u64
97 }
98
99 pub fn total_len(&self) -> Option<u64> {
101 let (lock, _) = &*self.inner;
102 lock.lock().unwrap().total_len
103 }
104
105 pub fn reader(&self) -> StreamingSource {
107 StreamingSource {
108 inner: self.inner.clone(),
109 pos: 0,
110 }
111 }
112}
113
114pub 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 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 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 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
194impl symphonia::core::io::MediaSource for StreamingSource {
196 fn is_seekable(&self) -> bool {
197 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 let pos = src.seek(SeekFrom::End(0)).unwrap();
299 assert_eq!(pos, 10);
300
301 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 assert!(buf.bytes_downloaded() != 0); buf.finish();
327 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 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 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 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 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}