1use std::collections::VecDeque;
10use std::sync::Arc;
11use tokio::sync::RwLock;
12
13pub const DEFAULT_STREAM_MAX_SIZE: usize = 10 * 1024 * 1024;
15
16#[derive(Clone)]
36pub struct BoundedStream {
37 inner: Arc<RwLock<BoundedStreamInner>>,
38 notify: Arc<tokio::sync::Notify>,
43}
44
45struct BoundedStreamInner {
46 buffer: VecDeque<u8>,
48 max_size: usize,
50 total_written: u64,
52 bytes_evicted: u64,
54 closed: bool,
56}
57
58impl BoundedStream {
59 pub fn new(max_size: usize) -> Self {
61 Self {
62 inner: Arc::new(RwLock::new(BoundedStreamInner {
63 buffer: VecDeque::with_capacity(max_size.min(8192)), max_size,
65 total_written: 0,
66 bytes_evicted: 0,
67 closed: false,
68 })),
69 notify: Arc::new(tokio::sync::Notify::new()),
70 }
71 }
72
73 pub fn default_size() -> Self {
75 Self::new(DEFAULT_STREAM_MAX_SIZE)
76 }
77
78 pub async fn write(&self, data: &[u8]) {
83 {
84 let mut inner = self.inner.write().await;
85
86 if inner.closed {
87 return;
88 }
89
90 inner.total_written += data.len() as u64;
91
92 if data.len() >= inner.max_size {
94 let start = data.len() - inner.max_size;
95 inner.bytes_evicted += inner.buffer.len() as u64 + start as u64;
96 inner.buffer.clear();
97 inner.buffer.extend(&data[start..]);
98 } else {
99 let needed = data.len();
101 let available = inner.max_size.saturating_sub(inner.buffer.len());
102
103 if needed > available {
104 let to_evict = needed - available;
105 let actual_evict = to_evict.min(inner.buffer.len());
106 inner.buffer.drain(..actual_evict);
107 inner.bytes_evicted += actual_evict as u64;
108 }
109
110 inner.buffer.extend(data);
112 }
113 }
114 self.notify.notify_waiters();
116 }
117
118 pub async fn read(&self) -> Vec<u8> {
123 let inner = self.inner.read().await;
124 inner.buffer.iter().copied().collect()
125 }
126
127 pub async fn read_string(&self) -> String {
129 let data = self.read().await;
130 String::from_utf8_lossy(&data).into_owned()
131 }
132
133 pub async fn close(&self) {
137 {
138 let mut inner = self.inner.write().await;
139 inner.closed = true;
140 }
141 self.notify.notify_waiters();
144 }
145
146 pub async fn is_closed(&self) -> bool {
148 let inner = self.inner.read().await;
149 inner.closed
150 }
151
152 pub async fn len(&self) -> usize {
154 let inner = self.inner.read().await;
155 inner.buffer.len()
156 }
157
158 pub async fn is_empty(&self) -> bool {
160 self.len().await == 0
161 }
162
163 pub async fn has_overflowed(&self) -> bool {
171 let inner = self.inner.read().await;
172 inner.bytes_evicted > 0
173 }
174
175 pub async fn changed_since(&self, seen_total_written: u64) -> StreamStats {
189 loop {
190 let notified = self.notify.notified();
191 tokio::pin!(notified);
192 notified.as_mut().enable();
193
194 let stats = self.stats().await;
195 if stats.closed || stats.total_written > seen_total_written {
196 return stats;
197 }
198
199 notified.await;
200 }
201 }
202
203 pub async fn stats(&self) -> StreamStats {
205 let inner = self.inner.read().await;
206 StreamStats {
207 current_size: inner.buffer.len(),
208 max_size: inner.max_size,
209 total_written: inner.total_written,
210 bytes_evicted: inner.bytes_evicted,
211 closed: inner.closed,
212 }
213 }
214}
215
216impl std::fmt::Debug for BoundedStream {
217 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
218 f.debug_struct("BoundedStream")
219 .field("inner", &"<locked>")
220 .finish()
221 }
222}
223
224#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
226#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
227pub struct StreamStats {
228 pub current_size: usize,
230 pub max_size: usize,
232 pub total_written: u64,
234 pub bytes_evicted: u64,
236 pub closed: bool,
238}
239
240impl StreamStats {
241 pub fn overflow_marker(&self, label: &str) -> String {
252 let max_mb = self.max_size as f64 / (1024.0 * 1024.0);
253 format!(
254 "[{label} truncated: output exceeded the {max_mb:.0}MB capture buffer \
255 — first {} bytes lost ({} bytes total written); enable output-limit \
256 to spill to disk]\n",
257 self.bytes_evicted, self.total_written,
258 )
259 }
260}
261
262pub async fn drain_to_stream<R>(reader: R, stream: Arc<BoundedStream>)
267where
268 R: tokio::io::AsyncRead + Unpin,
269{
270 drain_to_stream_teed(reader, stream, None).await
271}
272
273pub async fn drain_to_stream_teed<R>(
281 mut reader: R,
282 stream: Arc<BoundedStream>,
283 tee: Option<Arc<BoundedStream>>,
284) where
285 R: tokio::io::AsyncRead + Unpin,
286{
287 use tokio::io::AsyncReadExt;
288
289 let mut buf = [0u8; 8192];
290 loop {
291 match reader.read(&mut buf).await {
292 Ok(0) => break, Ok(n) => {
294 stream.write(&buf[..n]).await;
295 if let Some(tee) = &tee {
296 tee.write(&buf[..n]).await;
297 }
298 }
299 Err(e) => {
300 tracing::warn!("drain_to_stream read error: {}", e);
301 break;
302 }
303 }
304 }
305 stream.close().await;
306}
307
308#[cfg(test)]
309mod tests {
310 use super::*;
311
312 #[tokio::test]
315 async fn changed_since_wakes_on_a_later_write() {
316 let stream = Arc::new(BoundedStream::new(1024));
317 let writer = stream.clone();
318 tokio::spawn(async move {
319 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
320 writer.write(b"late").await;
321 });
322
323 let stats = tokio::time::timeout(
324 std::time::Duration::from_secs(5),
325 stream.changed_since(0),
326 )
327 .await
328 .expect("changed_since parked instead of waking on the write");
329 assert_eq!(stats.total_written, 4);
330 assert!(!stats.closed);
331 }
332
333 #[tokio::test]
336 async fn changed_since_returns_at_once_when_data_already_arrived() {
337 let stream = BoundedStream::new(1024);
338 stream.write(b"early").await;
339 let stats = tokio::time::timeout(
340 std::time::Duration::from_millis(500),
341 stream.changed_since(0),
342 )
343 .await
344 .expect("already-written data must not block");
345 assert_eq!(stats.total_written, 5);
346 }
347
348 #[tokio::test]
351 async fn changed_since_returns_on_close() {
352 let stream = Arc::new(BoundedStream::new(1024));
353 let closer = stream.clone();
354 tokio::spawn(async move {
355 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
356 closer.close().await;
357 });
358
359 let stats = tokio::time::timeout(
360 std::time::Duration::from_secs(5),
361 stream.changed_since(0),
362 )
363 .await
364 .expect("close must wake a waiter");
365 assert!(stats.closed, "the caller stops on this flag, not on a timeout");
366 assert_eq!(stats.total_written, 0);
367 }
368
369 #[tokio::test]
372 async fn drain_to_stream_teed_copies_to_both_and_closes_only_the_primary() {
373 let primary = Arc::new(BoundedStream::new(1024));
374 let tee = Arc::new(BoundedStream::new(1024));
375
376 let reader = std::io::Cursor::new(b"hello tee".to_vec());
377 drain_to_stream_teed(reader, primary.clone(), Some(tee.clone())).await;
378
379 assert_eq!(primary.read().await, b"hello tee");
380 assert_eq!(tee.read().await, b"hello tee");
381 assert!(primary.is_closed().await, "the drained command's own stream is done");
382 assert!(
383 !tee.is_closed().await,
384 "the job's stream must stay open for the next command in the job"
385 );
386 }
387
388 #[tokio::test]
389 async fn test_basic_write_read() {
390 let stream = BoundedStream::new(100);
391 stream.write(b"hello").await;
392 assert_eq!(stream.read().await, b"hello");
393 }
394
395 #[tokio::test]
396 async fn test_multiple_writes() {
397 let stream = BoundedStream::new(100);
398 stream.write(b"hello ").await;
399 stream.write(b"world").await;
400 assert_eq!(stream.read().await, b"hello world");
401 }
402
403 #[tokio::test]
404 async fn test_eviction_on_overflow() {
405 let stream = BoundedStream::new(10);
406 stream.write(b"12345").await;
407 stream.write(b"67890").await;
408 assert_eq!(stream.len().await, 10);
409
410 stream.write(b"ABCDE").await;
412 assert_eq!(stream.read().await, b"67890ABCDE");
413
414 let stats = stream.stats().await;
415 assert_eq!(stats.bytes_evicted, 5);
416 assert_eq!(stats.total_written, 15);
417 }
418
419 #[tokio::test]
420 async fn test_large_write_exceeds_buffer() {
421 let stream = BoundedStream::new(10);
422 stream.write(b"0123456789ABCDEFGHIJ").await;
424 assert_eq!(stream.read().await, b"ABCDEFGHIJ");
425 }
426
427 #[tokio::test]
428 async fn test_close_prevents_writes() {
429 let stream = BoundedStream::new(100);
430 stream.write(b"before").await;
431 stream.close().await;
432 stream.write(b"after").await;
433 assert_eq!(stream.read().await, b"before");
434 }
435
436 #[tokio::test]
437 async fn test_read_string() {
438 let stream = BoundedStream::new(100);
439 stream.write(b"hello world").await;
440 assert_eq!(stream.read_string().await, "hello world");
441 }
442
443 #[tokio::test]
444 async fn test_concurrent_writes() {
445 use std::sync::Arc;
446
447 let stream = Arc::new(BoundedStream::new(1000));
448
449 let handles: Vec<_> = (0..10)
450 .map(|i| {
451 let s = stream.clone();
452 tokio::spawn(async move {
453 for j in 0..10 {
454 s.write(format!("[{}-{}]", i, j).as_bytes()).await;
455 }
456 })
457 })
458 .collect();
459
460 for h in handles {
461 h.await.expect("task should not panic");
462 }
463
464 let data = stream.read().await;
467 assert!(!data.is_empty());
468 }
469
470 #[tokio::test]
471 async fn test_stats() {
472 let stream = BoundedStream::new(10);
473 stream.write(b"1234567890").await;
474
475 let stats = stream.stats().await;
476 assert_eq!(stats.current_size, 10);
477 assert_eq!(stats.max_size, 10);
478 assert_eq!(stats.total_written, 10);
479 assert_eq!(stats.bytes_evicted, 0);
480 assert!(!stats.closed);
481 }
482
483 #[tokio::test]
484 async fn test_empty_stream() {
485 let stream = BoundedStream::new(100);
486 assert!(stream.is_empty().await);
487 assert_eq!(stream.len().await, 0);
488 assert_eq!(stream.read().await, Vec::<u8>::new());
489 }
490
491 #[tokio::test]
492 async fn test_drain_to_stream() {
493 use std::io::Cursor;
494
495 let data = b"test data from reader";
496 let cursor = Cursor::new(data.to_vec());
497 let stream = Arc::new(BoundedStream::new(100));
498
499 drain_to_stream(cursor, stream.clone()).await;
500
501 assert_eq!(stream.read().await, data);
502 assert!(stream.is_closed().await);
503 }
504
505 #[tokio::test]
506 async fn test_default_size() {
507 let stream = BoundedStream::default_size();
508 let stats = stream.stats().await;
509 assert_eq!(stats.max_size, DEFAULT_STREAM_MAX_SIZE);
510 }
511
512 #[tokio::test]
513 async fn test_has_overflowed() {
514 let stream = BoundedStream::new(10);
515 assert!(!stream.has_overflowed().await, "empty stream has not overflowed");
516
517 stream.write(b"1234567890").await;
518 assert!(
519 !stream.has_overflowed().await,
520 "exactly filling the buffer is not an overflow"
521 );
522
523 stream.write(b"more").await; assert!(
525 stream.has_overflowed().await,
526 "writing past capacity must flip has_overflowed"
527 );
528 }
529
530 #[test]
531 fn stream_stats_round_trips_through_serde() {
532 let stats = StreamStats {
537 current_size: 42,
538 max_size: 100,
539 total_written: 142,
540 bytes_evicted: 100,
541 closed: true,
542 };
543 let json = serde_json::to_string(&stats).unwrap();
544 let back: StreamStats = serde_json::from_str(&json).unwrap();
545 assert_eq!(back.current_size, stats.current_size);
546 assert_eq!(back.max_size, stats.max_size);
547 assert_eq!(back.total_written, stats.total_written);
548 assert_eq!(back.bytes_evicted, stats.bytes_evicted);
549 assert_eq!(back.closed, stats.closed);
550 }
551
552 #[test]
553 fn test_overflow_marker_wording() {
554 let stats = StreamStats {
555 current_size: 10 * 1024 * 1024,
556 max_size: 10 * 1024 * 1024,
557 total_written: 15 * 1024 * 1024,
558 bytes_evicted: 5 * 1024 * 1024,
559 closed: true,
560 };
561 let marker = stats.overflow_marker("stdout");
562 assert!(marker.starts_with("[stdout truncated:"), "got: {marker}");
563 assert!(marker.contains("10MB"), "got: {marker}");
564 assert!(marker.contains(&(5 * 1024 * 1024).to_string()), "got: {marker}");
565 assert!(marker.contains(&(15 * 1024 * 1024).to_string()), "got: {marker}");
566 assert!(marker.contains("output-limit"), "got: {marker}");
567 }
568}