Skip to main content

tokio_prompt_orchestrator/
streaming_processor.rs

1//! Streaming token processor for SSE/streaming LLM responses.
2
3use std::collections::VecDeque;
4use std::time::{Duration, Instant};
5
6/// Events emitted by the streaming processor.
7#[derive(Debug, Clone)]
8pub enum StreamEvent {
9    /// A text token fragment.
10    Token(String),
11    /// A tool/function call request.
12    ToolCall { name: String, arguments: String },
13    /// Stream completed successfully.
14    Done { total_tokens: usize },
15    /// An error occurred during streaming.
16    Error(String),
17    /// Keepalive heartbeat (empty SSE comment).
18    Heartbeat,
19}
20
21/// Bounded queue of [`StreamEvent`]s with backpressure support.
22pub struct StreamBuffer {
23    inner: VecDeque<StreamEvent>,
24    max_size: usize,
25    total_received: usize,
26}
27
28impl StreamBuffer {
29    /// Create a new buffer with the given capacity.
30    pub fn new(max_size: usize) -> Self {
31        Self {
32            inner: VecDeque::with_capacity(max_size),
33            max_size,
34            total_received: 0,
35        }
36    }
37
38    /// Push an event into the buffer.
39    ///
40    /// Returns `false` (backpressure) if the buffer is already full.
41    pub fn push(&mut self, event: StreamEvent) -> bool {
42        if self.inner.len() >= self.max_size {
43            return false;
44        }
45        self.total_received += 1;
46        self.inner.push_back(event);
47        true
48    }
49
50    /// Pop the oldest event from the buffer.
51    pub fn pop(&mut self) -> Option<StreamEvent> {
52        self.inner.pop_front()
53    }
54
55    /// Concatenate and remove all [`StreamEvent::Token`] events, returning the result.
56    pub fn drain_tokens(&mut self) -> String {
57        let mut out = String::new();
58        self.inner.retain(|e| {
59            if let StreamEvent::Token(t) = e {
60                out.push_str(t);
61                false
62            } else {
63                true
64            }
65        });
66        out
67    }
68
69    /// Returns `true` if the buffer has reached its maximum size.
70    pub fn is_full(&self) -> bool {
71        self.inner.len() >= self.max_size
72    }
73
74    /// Number of events currently in the buffer.
75    pub fn len(&self) -> usize {
76        self.inner.len()
77    }
78
79    /// Returns `true` if the buffer contains no events.
80    pub fn is_empty(&self) -> bool {
81        self.inner.is_empty()
82    }
83
84    /// Total number of events ever pushed (including dropped ones).
85    pub fn total_received(&self) -> usize {
86        self.total_received
87    }
88}
89
90/// Result produced by [`TokenAccumulator::finish`].
91#[derive(Debug, Clone)]
92pub struct AccumulatorResult {
93    /// Full accumulated text.
94    pub text: String,
95    /// Number of tokens accumulated.
96    pub token_count: usize,
97    /// Wall-clock duration from first to last token, in milliseconds.
98    pub duration_ms: u64,
99    /// Tokens per second (0 if no duration elapsed).
100    pub tokens_per_second: f64,
101}
102
103/// Accumulates token fragments and tracks timing statistics.
104pub struct TokenAccumulator {
105    accumulated: String,
106    token_count: usize,
107    char_count: usize,
108    first_token_at: Option<Instant>,
109    last_token_at: Option<Instant>,
110}
111
112impl TokenAccumulator {
113    /// Create a new, empty accumulator.
114    pub fn new() -> Self {
115        Self {
116            accumulated: String::new(),
117            token_count: 0,
118            char_count: 0,
119            first_token_at: None,
120            last_token_at: None,
121        }
122    }
123
124    /// Append a token fragment.
125    pub fn push_token(&mut self, token: &str) {
126        let now = Instant::now();
127        if self.first_token_at.is_none() {
128            self.first_token_at = Some(now);
129        }
130        self.last_token_at = Some(now);
131        self.accumulated.push_str(token);
132        self.token_count += 1;
133        self.char_count += token.len();
134    }
135
136    /// Duration from the very first token to the very first token (i.e. always `Some(0)` after
137    /// the first push, but semantically "time to first token" measured at call site).
138    ///
139    /// Returns `None` if no tokens have been pushed yet.
140    pub fn time_to_first_token(&self) -> Option<Duration> {
141        self.first_token_at.map(|t| t.elapsed())
142    }
143
144    /// Tokens per second calculated over the full accumulation window.
145    ///
146    /// Returns `0.0` if fewer than two tokens were pushed or no time has elapsed.
147    pub fn tokens_per_second(&self) -> f64 {
148        match (self.first_token_at, self.last_token_at) {
149            (Some(first), Some(last)) => {
150                let secs = last.duration_since(first).as_secs_f64();
151                if secs > 0.0 {
152                    self.token_count as f64 / secs
153                } else {
154                    0.0
155                }
156            }
157            _ => 0.0,
158        }
159    }
160
161    /// Snapshot the current accumulation state as an [`AccumulatorResult`].
162    pub fn finish(&self) -> AccumulatorResult {
163        let duration_ms = match (self.first_token_at, self.last_token_at) {
164            (Some(first), Some(last)) => last.duration_since(first).as_millis() as u64,
165            _ => 0,
166        };
167        AccumulatorResult {
168            text: self.accumulated.clone(),
169            token_count: self.token_count,
170            duration_ms,
171            tokens_per_second: self.tokens_per_second(),
172        }
173    }
174
175    /// Reference to the accumulated text so far.
176    pub fn text(&self) -> &str {
177        &self.accumulated
178    }
179
180    /// Number of tokens pushed so far.
181    pub fn token_count(&self) -> usize {
182        self.token_count
183    }
184
185    /// Total characters accumulated.
186    pub fn char_count(&self) -> usize {
187        self.char_count
188    }
189}
190
191impl Default for TokenAccumulator {
192    fn default() -> Self {
193        Self::new()
194    }
195}
196
197/// Running statistics for a [`StreamingProcessor`].
198#[derive(Debug, Clone, Default)]
199pub struct StreamStats {
200    /// Total SSE events processed.
201    pub events_processed: usize,
202    /// Total token fragments received.
203    pub tokens_received: usize,
204    /// Total tool-call events seen.
205    pub tool_calls_seen: usize,
206    /// Total error events seen.
207    pub errors_seen: usize,
208    /// Number of times the buffer was full (backpressure activations).
209    pub buffer_full_count: usize,
210}
211
212/// Parses SSE chunks from LLM streaming endpoints, accumulates tokens, and
213/// maintains a bounded [`StreamBuffer`] with backpressure.
214pub struct StreamingProcessor {
215    buffer: StreamBuffer,
216    accumulator: TokenAccumulator,
217    stats: StreamStats,
218}
219
220impl StreamingProcessor {
221    /// Create a new processor with the given buffer capacity.
222    pub fn new(buffer_size: usize) -> Self {
223        Self {
224            buffer: StreamBuffer::new(buffer_size),
225            accumulator: TokenAccumulator::new(),
226            stats: StreamStats::default(),
227        }
228    }
229
230    /// Parse a raw SSE chunk and return the extracted events.
231    ///
232    /// Handles:
233    /// - Lines prefixed with `"data: "`
234    /// - The `"[DONE]"` sentinel
235    /// - JSON payloads with `choices[0].delta.content` (OpenAI-style)
236    pub fn process_chunk(&mut self, raw: &str) -> Vec<StreamEvent> {
237        let mut events = Vec::new();
238
239        for line in raw.lines() {
240            let line = line.trim();
241            if line.is_empty() {
242                continue;
243            }
244            // SSE comment → heartbeat
245            if line.starts_with(':') {
246                events.push(StreamEvent::Heartbeat);
247                continue;
248            }
249            // Strip "data: " prefix
250            let data = if let Some(rest) = line.strip_prefix("data: ") {
251                rest.trim()
252            } else {
253                continue;
254            };
255
256            // [DONE] sentinel
257            if data == "[DONE]" {
258                let total = self.accumulator.token_count();
259                events.push(StreamEvent::Done { total_tokens: total });
260                continue;
261            }
262
263            // Try to parse as JSON
264            match serde_json::from_str::<serde_json::Value>(data) {
265                Ok(json) => {
266                    // Extract delta content (OpenAI-style)
267                    if let Some(content) = json
268                        .pointer("/choices/0/delta/content")
269                        .and_then(|v| v.as_str())
270                    {
271                        if !content.is_empty() {
272                            events.push(StreamEvent::Token(content.to_string()));
273                        }
274                    }
275
276                    // Tool calls embedded in delta
277                    let tool_events = Self::extract_tool_calls(data);
278                    for (name, arguments) in tool_events {
279                        events.push(StreamEvent::ToolCall { name, arguments });
280                    }
281
282                    // finish_reason → Done
283                    if let Some(reason) = Self::detect_finish_reason(data) {
284                        if reason == "stop" || reason == "length" {
285                            let total = self.accumulator.token_count();
286                            events.push(StreamEvent::Done { total_tokens: total });
287                        }
288                    }
289                }
290                Err(e) => {
291                    events.push(StreamEvent::Error(format!("JSON parse error: {}", e)));
292                }
293            }
294        }
295
296        // Feed events into buffer and accumulator
297        for event in &events {
298            self.stats.events_processed += 1;
299            match event {
300                StreamEvent::Token(t) => {
301                    self.accumulator.push_token(t);
302                    self.stats.tokens_received += 1;
303                }
304                StreamEvent::ToolCall { .. } => self.stats.tool_calls_seen += 1,
305                StreamEvent::Error(_) => self.stats.errors_seen += 1,
306                _ => {}
307            }
308            if !self.buffer.push(event.clone()) {
309                self.stats.buffer_full_count += 1;
310            }
311        }
312
313        events
314    }
315
316    /// Extract tool/function call blocks from a raw JSON string.
317    ///
318    /// Returns a list of `(name, arguments)` pairs.
319    pub fn extract_tool_calls(raw: &str) -> Vec<(String, String)> {
320        let mut results = Vec::new();
321        let Ok(json) = serde_json::from_str::<serde_json::Value>(raw) else {
322            return results;
323        };
324
325        // OpenAI style: choices[0].delta.tool_calls[]
326        if let Some(tool_calls) = json.pointer("/choices/0/delta/tool_calls").and_then(|v| v.as_array()) {
327            for tc in tool_calls {
328                let name = tc
329                    .pointer("/function/name")
330                    .and_then(|v| v.as_str())
331                    .unwrap_or("")
332                    .to_string();
333                let arguments = tc
334                    .pointer("/function/arguments")
335                    .and_then(|v| v.as_str())
336                    .unwrap_or("{}")
337                    .to_string();
338                if !name.is_empty() {
339                    results.push((name, arguments));
340                }
341            }
342        }
343
344        // Anthropic style: type == "tool_use"
345        if json.get("type").and_then(|v| v.as_str()) == Some("tool_use") {
346            let name = json.get("name").and_then(|v| v.as_str()).unwrap_or("").to_string();
347            let arguments = json
348                .get("input")
349                .map(|v| v.to_string())
350                .unwrap_or_else(|| "{}".to_string());
351            if !name.is_empty() {
352                results.push((name, arguments));
353            }
354        }
355
356        results
357    }
358
359    /// Detect the finish reason from a raw JSON SSE payload.
360    ///
361    /// Recognises `"stop"`, `"length"`, and `"tool_calls"`.
362    pub fn detect_finish_reason(raw: &str) -> Option<String> {
363        let json = serde_json::from_str::<serde_json::Value>(raw).ok()?;
364        let reason = json
365            .pointer("/choices/0/finish_reason")
366            .and_then(|v| v.as_str())?;
367        match reason {
368            "stop" | "length" | "tool_calls" => Some(reason.to_string()),
369            _ => None,
370        }
371    }
372
373    /// Feed a batch of pre-parsed events into the accumulator and return a reference to it.
374    pub fn accumulate(&mut self, events: &[StreamEvent]) -> &TokenAccumulator {
375        for event in events {
376            if let StreamEvent::Token(t) = event {
377                self.accumulator.push_token(t);
378            }
379        }
380        &self.accumulator
381    }
382
383    /// Current processor statistics.
384    pub fn stats(&self) -> StreamStats {
385        self.stats.clone()
386    }
387
388    /// Reference to the internal buffer.
389    pub fn buffer(&self) -> &StreamBuffer {
390        &self.buffer
391    }
392
393    /// Mutable reference to the internal buffer.
394    pub fn buffer_mut(&mut self) -> &mut StreamBuffer {
395        &mut self.buffer
396    }
397
398    /// Reference to the token accumulator.
399    pub fn accumulator(&self) -> &TokenAccumulator {
400        &self.accumulator
401    }
402}