tokio_prompt_orchestrator/
stream_agg.rs1use 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#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
46pub struct StreamChunk {
47 pub session_id: u64,
49 pub token: String,
51 pub is_final: bool,
53 pub metadata: HashMap<String, String>,
55}
56
57#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
59pub struct AggStats {
60 pub active_streams: usize,
62 pub completed_streams: u64,
64 pub total_tokens_streamed: u64,
66}
67
68struct SessionStream {
73 buffer: String,
74 tx: broadcast::Sender<StreamChunk>,
75}
76
77#[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 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 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 let _ = entry.tx.send(chunk);
125 }
126
127 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 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 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 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#[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 #[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 #[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 #[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 #[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 assert_eq!(a.stats().total_tokens_streamed, 1);
390 }
391}