Skip to main content

tokio_prompt_orchestrator/
stream_agg.rs

1//! # Streaming Response Aggregator
2//!
3//! Collects streaming token chunks from LLM inference sessions into complete
4//! responses, with real-time per-session broadcast for subscribers.
5//!
6//! ## Overview
7//!
8//! [`StreamAggregator`] buffers [`StreamChunk`]s keyed by `session_id`.  When
9//! the final chunk arrives, [`StreamAggregator::complete`] flushes the buffer
10//! and returns the assembled text.  Subscribers registered via
11//! [`StreamAggregator::subscribe`] receive every chunk in real-time via a
12//! [`tokio::sync::broadcast`] channel.
13//!
14//! ## Example
15//!
16//! ```rust
17//! use tokio_prompt_orchestrator::stream_agg::{StreamAggregator, StreamChunk};
18//! use std::collections::HashMap;
19//!
20//! # #[tokio::main]
21//! # async fn main() {
22//! let agg = StreamAggregator::new(64);
23//!
24//! agg.feed(StreamChunk { session_id: 1, token: "Hello".into(), is_final: false, metadata: HashMap::new() });
25//! agg.feed(StreamChunk { session_id: 1, token: " world".into(), is_final: true, metadata: HashMap::new() });
26//!
27//! let text = agg.complete(1);
28//! assert_eq!(text.as_deref(), Some("Hello world"));
29//! # }
30//! ```
31
32use dashmap::DashMap;
33use futures::Stream;
34use std::collections::HashMap;
35use std::pin::Pin;
36use std::sync::atomic::{AtomicU64, Ordering};
37use std::sync::Arc;
38use tokio::sync::broadcast;
39
40// ============================================================================
41// Domain types
42// ============================================================================
43
44/// A single token chunk produced by a streaming LLM inference response.
45#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
46pub struct StreamChunk {
47    /// The session this token belongs to.
48    pub session_id: u64,
49    /// The token text (may be a sub-word piece).
50    pub token: String,
51    /// When `true`, this is the last token in the stream.
52    pub is_final: bool,
53    /// Arbitrary key-value metadata (e.g. model, finish_reason).
54    pub metadata: HashMap<String, String>,
55}
56
57/// Aggregate statistics for a [`StreamAggregator`].
58#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
59pub struct AggStats {
60    /// Number of sessions currently buffering tokens.
61    pub active_streams: usize,
62    /// Cumulative number of sessions that have been completed via [`StreamAggregator::complete`].
63    pub completed_streams: u64,
64    /// Total tokens received across all sessions (including completed ones).
65    pub total_tokens_streamed: u64,
66}
67
68// ============================================================================
69// Internal state per session
70// ============================================================================
71
72struct SessionStream {
73    buffer: String,
74    tx: broadcast::Sender<StreamChunk>,
75}
76
77// ============================================================================
78// StreamAggregator
79// ============================================================================
80
81/// Concurrent, lock-free streaming response aggregator.
82///
83/// `channel_capacity` controls the bounded broadcast channel depth.  Lagging
84/// subscribers that cannot keep up will have their oldest messages dropped.
85#[derive(Clone)]
86pub struct StreamAggregator {
87    sessions: Arc<DashMap<u64, SessionStream>>,
88    channel_capacity: usize,
89    completed_streams: Arc<AtomicU64>,
90    total_tokens: Arc<AtomicU64>,
91}
92
93impl StreamAggregator {
94    /// Create a new aggregator with the given per-session broadcast capacity.
95    pub fn new(channel_capacity: usize) -> Self {
96        Self {
97            sessions: Arc::new(DashMap::new()),
98            channel_capacity: channel_capacity.max(1),
99            completed_streams: Arc::new(AtomicU64::new(0)),
100            total_tokens: Arc::new(AtomicU64::new(0)),
101        }
102    }
103
104    /// Feed a token chunk into the aggregator.
105    ///
106    /// Internally:
107    /// - A per-session buffer is created on first access.
108    /// - The token is appended to the buffer.
109    /// - The chunk is broadcast to any active subscribers.
110    pub fn feed(&self, chunk: StreamChunk) {
111        self.total_tokens.fetch_add(1, Ordering::Relaxed);
112        let sid = chunk.session_id;
113
114        let mut entry = self.sessions.entry(sid).or_insert_with(|| {
115            let (tx, _) = broadcast::channel(self.channel_capacity);
116            SessionStream {
117                buffer: String::new(),
118                tx,
119            }
120        });
121
122        entry.buffer.push_str(&chunk.token);
123        // Broadcast — ignore errors (no active subscribers is fine)
124        let _ = entry.tx.send(chunk);
125    }
126
127    /// Flush the session buffer and return the complete assembled text.
128    ///
129    /// Returns `None` if the session has never received any chunks.
130    /// After calling `complete`, the session entry is removed from the map.
131    pub fn complete(&self, session_id: u64) -> Option<String> {
132        if let Some((_, stream)) = self.sessions.remove(&session_id) {
133            self.completed_streams.fetch_add(1, Ordering::Relaxed);
134            Some(stream.buffer)
135        } else {
136            None
137        }
138    }
139
140    /// Subscribe to real-time token chunks for a session.
141    ///
142    /// Returns a [`Stream`] that yields every [`StreamChunk`] fed for
143    /// `session_id` after the subscription is registered.  Chunks that
144    /// arrived before the call are not replayed.
145    ///
146    /// If the session does not yet exist, an empty buffer entry is created so
147    /// that future [`feed`][Self::feed] calls reach the subscriber.
148    pub fn subscribe(
149        &self,
150        session_id: u64,
151    ) -> Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'static>> {
152        let rx = {
153            let entry = self.sessions.entry(session_id).or_insert_with(|| {
154                let (tx, _) = broadcast::channel(self.channel_capacity);
155                SessionStream {
156                    buffer: String::new(),
157                    tx,
158                }
159            });
160            entry.tx.subscribe()
161        };
162
163        // Convert broadcast::Receiver into a Stream using async_stream.
164        // Filter out broadcast errors (lagged / sender dropped).
165        let stream = async_stream::stream! {
166            let mut rx = rx;
167            loop {
168                match rx.recv().await {
169                    Ok(chunk) => yield chunk,
170                    Err(broadcast::error::RecvError::Lagged(_)) => continue,
171                    Err(broadcast::error::RecvError::Closed) => break,
172                }
173            }
174        };
175        Box::pin(stream)
176    }
177
178    /// Return aggregate statistics.
179    pub fn stats(&self) -> AggStats {
180        AggStats {
181            active_streams: self.sessions.len(),
182            completed_streams: self.completed_streams.load(Ordering::Relaxed),
183            total_tokens_streamed: self.total_tokens.load(Ordering::Relaxed),
184        }
185    }
186}
187
188// ============================================================================
189// Unit tests
190// ============================================================================
191
192#[cfg(test)]
193mod tests {
194    use super::*;
195    use futures::StreamExt as _;
196
197    fn chunk(session_id: u64, token: &str, is_final: bool) -> StreamChunk {
198        StreamChunk {
199            session_id,
200            token: token.into(),
201            is_final,
202            metadata: HashMap::new(),
203        }
204    }
205
206    fn agg() -> StreamAggregator {
207        StreamAggregator::new(64)
208    }
209
210    // --- feed + complete ---
211
212    #[test]
213    fn test_complete_empty_session() {
214        let a = agg();
215        assert!(a.complete(99).is_none());
216    }
217
218    #[test]
219    fn test_single_chunk_complete() {
220        let a = agg();
221        a.feed(chunk(1, "Hello", true));
222        assert_eq!(a.complete(1).as_deref(), Some("Hello"));
223    }
224
225    #[test]
226    fn test_multiple_chunks_assembled() {
227        let a = agg();
228        a.feed(chunk(1, "The", false));
229        a.feed(chunk(1, " quick", false));
230        a.feed(chunk(1, " fox", true));
231        assert_eq!(a.complete(1).as_deref(), Some("The quick fox"));
232    }
233
234    #[test]
235    fn test_complete_removes_session() {
236        let a = agg();
237        a.feed(chunk(2, "x", true));
238        a.complete(2);
239        assert!(a.complete(2).is_none());
240    }
241
242    #[test]
243    fn test_independent_sessions() {
244        let a = agg();
245        a.feed(chunk(1, "A", false));
246        a.feed(chunk(2, "B", false));
247        a.feed(chunk(1, "A2", true));
248        a.feed(chunk(2, "B2", true));
249        assert_eq!(a.complete(1).as_deref(), Some("AA2"));
250        assert_eq!(a.complete(2).as_deref(), Some("BB2"));
251    }
252
253    #[test]
254    fn test_complete_after_complete_returns_none() {
255        let a = agg();
256        a.feed(chunk(5, "z", true));
257        a.complete(5);
258        assert!(a.complete(5).is_none());
259    }
260
261    // --- stats ---
262
263    #[test]
264    fn test_stats_initial() {
265        let a = agg();
266        let s = a.stats();
267        assert_eq!(s.active_streams, 0);
268        assert_eq!(s.completed_streams, 0);
269        assert_eq!(s.total_tokens_streamed, 0);
270    }
271
272    #[test]
273    fn test_stats_active_count() {
274        let a = agg();
275        a.feed(chunk(1, "a", false));
276        a.feed(chunk(2, "b", false));
277        assert_eq!(a.stats().active_streams, 2);
278    }
279
280    #[test]
281    fn test_stats_completed_increments() {
282        let a = agg();
283        a.feed(chunk(1, "a", true));
284        a.complete(1);
285        assert_eq!(a.stats().completed_streams, 1);
286    }
287
288    #[test]
289    fn test_stats_total_tokens() {
290        let a = agg();
291        a.feed(chunk(1, "a", false));
292        a.feed(chunk(1, "b", false));
293        a.feed(chunk(1, "c", true));
294        assert_eq!(a.stats().total_tokens_streamed, 3);
295    }
296
297    #[test]
298    fn test_stats_active_drops_after_complete() {
299        let a = agg();
300        a.feed(chunk(1, "x", true));
301        assert_eq!(a.stats().active_streams, 1);
302        a.complete(1);
303        assert_eq!(a.stats().active_streams, 0);
304    }
305
306    // --- subscribe ---
307
308    #[tokio::test]
309    async fn test_subscribe_receives_chunks() {
310        let a = agg();
311        let mut sub = a.subscribe(10);
312
313        a.feed(chunk(10, "tok1", false));
314        a.feed(chunk(10, "tok2", true));
315
316        let first = tokio::time::timeout(std::time::Duration::from_millis(100), sub.next())
317            .await
318            .expect("timeout")
319            .expect("stream ended");
320        assert_eq!(first.token, "tok1");
321    }
322
323    #[tokio::test]
324    async fn test_subscribe_multiple_sessions_isolated() {
325        let a = agg();
326        let mut sub1 = a.subscribe(1);
327        let _sub2 = a.subscribe(2);
328
329        a.feed(chunk(2, "not-for-1", false));
330        a.feed(chunk(1, "for-1", true));
331
332        let received = tokio::time::timeout(std::time::Duration::from_millis(200), sub1.next())
333            .await
334            .expect("timeout")
335            .expect("no chunk");
336        assert_eq!(received.token, "for-1");
337    }
338
339    #[tokio::test]
340    async fn test_subscribe_then_complete_consistent() {
341        let a = agg();
342        let mut sub = a.subscribe(7);
343
344        a.feed(chunk(7, "part1", false));
345        a.feed(chunk(7, "part2", true));
346
347        let _c1 = sub.next().await;
348        let _c2 = sub.next().await;
349
350        let full = a.complete(7);
351        assert_eq!(full.as_deref(), Some("part1part2"));
352    }
353
354    #[tokio::test]
355    async fn test_subscribe_before_any_feed() {
356        let a = agg();
357        let mut sub = a.subscribe(99);
358        a.feed(chunk(99, "hello", true));
359        let got = tokio::time::timeout(std::time::Duration::from_millis(100), sub.next())
360            .await
361            .expect("timeout")
362            .expect("no chunk");
363        assert_eq!(got.token, "hello");
364    }
365
366    // --- clone shares state ---
367
368    #[test]
369    fn test_clone_shares_state() {
370        let a1 = agg();
371        let a2 = a1.clone();
372        a1.feed(chunk(1, "x", true));
373        assert!(a2.complete(1).is_some());
374    }
375
376    #[test]
377    fn test_metadata_preserved() {
378        let a = agg();
379        let mut meta = HashMap::new();
380        meta.insert("model".into(), "claude-3".into());
381        let c = StreamChunk {
382            session_id: 1,
383            token: "tok".into(),
384            is_final: false,
385            metadata: meta,
386        };
387        a.feed(c);
388        // Just verifying feed doesn't panic; metadata lives in the broadcast
389        assert_eq!(a.stats().total_tokens_streamed, 1);
390    }
391}