tokio_prompt_orchestrator/
streaming_processor.rs1use std::collections::VecDeque;
4use std::time::{Duration, Instant};
5
6#[derive(Debug, Clone)]
8pub enum StreamEvent {
9 Token(String),
11 ToolCall { name: String, arguments: String },
13 Done { total_tokens: usize },
15 Error(String),
17 Heartbeat,
19}
20
21pub struct StreamBuffer {
23 inner: VecDeque<StreamEvent>,
24 max_size: usize,
25 total_received: usize,
26}
27
28impl StreamBuffer {
29 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 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 pub fn pop(&mut self) -> Option<StreamEvent> {
52 self.inner.pop_front()
53 }
54
55 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 pub fn is_full(&self) -> bool {
71 self.inner.len() >= self.max_size
72 }
73
74 pub fn len(&self) -> usize {
76 self.inner.len()
77 }
78
79 pub fn is_empty(&self) -> bool {
81 self.inner.is_empty()
82 }
83
84 pub fn total_received(&self) -> usize {
86 self.total_received
87 }
88}
89
90#[derive(Debug, Clone)]
92pub struct AccumulatorResult {
93 pub text: String,
95 pub token_count: usize,
97 pub duration_ms: u64,
99 pub tokens_per_second: f64,
101}
102
103pub 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 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 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 pub fn time_to_first_token(&self) -> Option<Duration> {
141 self.first_token_at.map(|t| t.elapsed())
142 }
143
144 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 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 pub fn text(&self) -> &str {
177 &self.accumulated
178 }
179
180 pub fn token_count(&self) -> usize {
182 self.token_count
183 }
184
185 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#[derive(Debug, Clone, Default)]
199pub struct StreamStats {
200 pub events_processed: usize,
202 pub tokens_received: usize,
204 pub tool_calls_seen: usize,
206 pub errors_seen: usize,
208 pub buffer_full_count: usize,
210}
211
212pub struct StreamingProcessor {
215 buffer: StreamBuffer,
216 accumulator: TokenAccumulator,
217 stats: StreamStats,
218}
219
220impl StreamingProcessor {
221 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 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 if line.starts_with(':') {
246 events.push(StreamEvent::Heartbeat);
247 continue;
248 }
249 let data = if let Some(rest) = line.strip_prefix("data: ") {
251 rest.trim()
252 } else {
253 continue;
254 };
255
256 if data == "[DONE]" {
258 let total = self.accumulator.token_count();
259 events.push(StreamEvent::Done { total_tokens: total });
260 continue;
261 }
262
263 match serde_json::from_str::<serde_json::Value>(data) {
265 Ok(json) => {
266 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 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 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 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 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 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 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 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 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 pub fn stats(&self) -> StreamStats {
385 self.stats.clone()
386 }
387
388 pub fn buffer(&self) -> &StreamBuffer {
390 &self.buffer
391 }
392
393 pub fn buffer_mut(&mut self) -> &mut StreamBuffer {
395 &mut self.buffer
396 }
397
398 pub fn accumulator(&self) -> &TokenAccumulator {
400 &self.accumulator
401 }
402}