Skip to main content

runifold_model/
stream.rs

1use std::collections::BTreeMap;
2
3use serde::{Deserialize, Serialize};
4use serde_json::Value;
5
6use crate::{
7    ContentPart, FinishReason, ModelError, ModelErrorKind, ModelRef, ModelResponse, ModelUsage,
8    ModelWarning, ProviderData, ReasoningPart, ToolCall,
9};
10
11/// The type and initial metadata of a streamed content block.
12#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
13#[serde(tag = "type", rename_all = "snake_case")]
14#[non_exhaustive]
15pub enum ContentBlockKind {
16    /// Text output.
17    Text,
18    /// Reasoning or thinking output.
19    Reasoning {
20        /// Initial signature or continuation token.
21        signature: Option<String>,
22        /// Whether the provider redacted the reasoning body.
23        redacted: bool,
24    },
25    /// A streamed tool call.
26    ToolCall {
27        /// Provider- or runtime-assigned call identity.
28        id: String,
29        /// Tool name.
30        name: String,
31    },
32    /// A streamed refusal.
33    Refusal,
34}
35
36/// A provider event retained without normalization.
37#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
38pub struct ProviderEvent {
39    /// Provider namespace.
40    pub provider: String,
41    /// Provider event name.
42    pub name: String,
43    /// Original structured payload.
44    pub payload: Value,
45}
46
47/// Canonical events emitted by a streaming model call.
48#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
49#[serde(tag = "type", rename_all = "snake_case")]
50#[non_exhaustive]
51pub enum ModelStreamEvent {
52    /// The provider accepted the request and started a response.
53    ResponseStarted {
54        /// Provider response identity.
55        id: Option<String>,
56        /// Actual model serving the request.
57        model: ModelRef,
58    },
59    /// A delta-capable content block started.
60    ContentBlockStarted {
61        /// Stable output ordering index.
62        index: u32,
63        /// Block type and initial metadata.
64        kind: ContentBlockKind,
65    },
66    /// A text delta.
67    TextDelta {
68        /// Target block index.
69        index: u32,
70        /// Appended text.
71        text: String,
72    },
73    /// A reasoning-text delta.
74    ReasoningDelta {
75        /// Target block index.
76        index: u32,
77        /// Appended reasoning text.
78        text: String,
79    },
80    /// A reasoning-signature delta.
81    ReasoningSignatureDelta {
82        /// Target block index.
83        index: u32,
84        /// Appended signature data.
85        signature: String,
86    },
87    /// A raw JSON fragment for tool arguments.
88    ToolArgumentsDelta {
89        /// Target block index.
90        index: u32,
91        /// Appended raw JSON text.
92        json: String,
93    },
94    /// A refusal-text delta.
95    RefusalDelta {
96        /// Target block index.
97        index: u32,
98        /// Appended refusal text.
99        text: String,
100    },
101    /// A delta-capable block completed.
102    ContentBlockCompleted {
103        /// Completed block index.
104        index: u32,
105    },
106    /// A complete non-delta content part arrived.
107    ContentPartCompleted {
108        /// Stable output ordering index.
109        index: u32,
110        /// Completed content.
111        part: ContentPart,
112    },
113    /// A cumulative usage snapshot.
114    UsageUpdated {
115        /// Latest cumulative model usage.
116        usage: ModelUsage,
117    },
118    /// A translation or feature-degradation warning.
119    Warning {
120        /// Visible warning.
121        warning: ModelWarning,
122    },
123    /// A provider heartbeat without model content.
124    Heartbeat,
125    /// An unknown or provider-specific event.
126    Provider {
127        /// Retained provider event.
128        event: ProviderEvent,
129    },
130    /// The response completed.
131    ResponseCompleted {
132        /// Normalized terminal reason.
133        finish_reason: FinishReason,
134        /// Namespaced terminal provider metadata.
135        provider_metadata: BTreeMap<String, Value>,
136    },
137}
138
139#[derive(Debug)]
140enum PartialBlock {
141    Text(String),
142    Reasoning {
143        text: String,
144        signature: Option<String>,
145        redacted: bool,
146    },
147    ToolCall {
148        id: String,
149        name: String,
150        arguments: String,
151    },
152    Refusal(String),
153}
154
155impl PartialBlock {
156    fn from_kind(kind: ContentBlockKind) -> Self {
157        match kind {
158            ContentBlockKind::Text => Self::Text(String::new()),
159            ContentBlockKind::Reasoning {
160                signature,
161                redacted,
162            } => Self::Reasoning {
163                text: String::new(),
164                signature,
165                redacted,
166            },
167            ContentBlockKind::ToolCall { id, name } => Self::ToolCall {
168                id,
169                name,
170                arguments: String::new(),
171            },
172            ContentBlockKind::Refusal => Self::Refusal(String::new()),
173        }
174    }
175
176    fn complete(self) -> Result<ContentPart, ModelError> {
177        match self {
178            Self::Text(text) => Ok(ContentPart::Text { text }),
179            Self::Reasoning {
180                text,
181                signature,
182                redacted,
183            } => Ok(ContentPart::Reasoning(ReasoningPart {
184                text: (!text.is_empty()).then_some(text),
185                signature,
186                redacted,
187                provider_data: Vec::new(),
188            })),
189            Self::ToolCall {
190                id,
191                name,
192                arguments,
193            } => {
194                let parsed = if arguments.trim().is_empty() {
195                    serde_json::json!({})
196                } else {
197                    serde_json::from_str(&arguments).map_err(|error| {
198                        ModelError::local(
199                            ModelErrorKind::MalformedToolArguments,
200                            format!("tool call {id} returned invalid JSON arguments: {error}"),
201                        )
202                    })?
203                };
204                Ok(ContentPart::ToolCall(ToolCall {
205                    id,
206                    name,
207                    arguments: parsed,
208                    raw_arguments: Some(arguments),
209                    metadata: BTreeMap::new(),
210                }))
211            }
212            Self::Refusal(text) => Ok(ContentPart::Refusal { text }),
213        }
214    }
215}
216
217/// Strictly reconstructs a canonical response from model stream events.
218#[derive(Debug, Default)]
219pub struct ModelStreamAccumulator {
220    started: bool,
221    completed: bool,
222    id: Option<String>,
223    model: Option<ModelRef>,
224    open_blocks: BTreeMap<u32, PartialBlock>,
225    content: BTreeMap<u32, ContentPart>,
226    usage: ModelUsage,
227    warnings: Vec<ModelWarning>,
228    provider_events: Vec<ProviderData>,
229}
230
231impl ModelStreamAccumulator {
232    /// Creates an empty accumulator.
233    pub fn new() -> Self {
234        Self::default()
235    }
236
237    /// Applies one event and returns the response when the terminal event arrives.
238    ///
239    /// # Errors
240    ///
241    /// Returns [`ModelError`] when event order is invalid, a delta targets the
242    /// wrong block type, indices collide, a response completes with open
243    /// blocks, or tool arguments contain malformed JSON.
244    pub fn push(&mut self, event: ModelStreamEvent) -> Result<Option<ModelResponse>, ModelError> {
245        if self.completed {
246            return Err(state_error("received an event after response completion"));
247        }
248
249        match event {
250            ModelStreamEvent::ResponseStarted { id, model } => self.start(id, model),
251            ModelStreamEvent::ContentBlockStarted { index, kind } => self.start_block(index, kind),
252            ModelStreamEvent::TextDelta { index, text } => {
253                match self.open_block_mut(index)? {
254                    PartialBlock::Text(current) => current.push_str(&text),
255                    _ => return Err(wrong_delta(index, "text")),
256                }
257                Ok(None)
258            }
259            ModelStreamEvent::ReasoningDelta { index, text } => {
260                match self.open_block_mut(index)? {
261                    PartialBlock::Reasoning { text: current, .. } => current.push_str(&text),
262                    _ => return Err(wrong_delta(index, "reasoning")),
263                }
264                Ok(None)
265            }
266            ModelStreamEvent::ReasoningSignatureDelta { index, signature } => {
267                match self.open_block_mut(index)? {
268                    PartialBlock::Reasoning {
269                        signature: current, ..
270                    } => current.get_or_insert_with(String::new).push_str(&signature),
271                    _ => return Err(wrong_delta(index, "reasoning signature")),
272                }
273                Ok(None)
274            }
275            ModelStreamEvent::ToolArgumentsDelta { index, json } => {
276                match self.open_block_mut(index)? {
277                    PartialBlock::ToolCall { arguments, .. } => arguments.push_str(&json),
278                    _ => return Err(wrong_delta(index, "tool arguments")),
279                }
280                Ok(None)
281            }
282            ModelStreamEvent::RefusalDelta { index, text } => {
283                match self.open_block_mut(index)? {
284                    PartialBlock::Refusal(current) => current.push_str(&text),
285                    _ => return Err(wrong_delta(index, "refusal")),
286                }
287                Ok(None)
288            }
289            ModelStreamEvent::ContentBlockCompleted { index } => self.complete_block(index),
290            ModelStreamEvent::ContentPartCompleted { index, part } => {
291                self.complete_part(index, part)
292            }
293            ModelStreamEvent::UsageUpdated { usage } => {
294                self.require_started()?;
295                self.usage = usage;
296                Ok(None)
297            }
298            ModelStreamEvent::Warning { warning } => {
299                self.require_started()?;
300                self.warnings.push(warning);
301                Ok(None)
302            }
303            ModelStreamEvent::Heartbeat => {
304                self.require_started()?;
305                Ok(None)
306            }
307            ModelStreamEvent::Provider { event } => {
308                self.require_started()?;
309                self.provider_events.push(ProviderData {
310                    provider: event.provider,
311                    kind: event.name,
312                    value: event.payload,
313                });
314                Ok(None)
315            }
316            ModelStreamEvent::ResponseCompleted {
317                finish_reason,
318                provider_metadata,
319            } => self.complete(finish_reason, provider_metadata),
320        }
321    }
322
323    fn start(
324        &mut self,
325        id: Option<String>,
326        model: ModelRef,
327    ) -> Result<Option<ModelResponse>, ModelError> {
328        if self.started {
329            return Err(state_error("received more than one response-start event"));
330        }
331        self.started = true;
332        self.id = id;
333        self.model = Some(model);
334        Ok(None)
335    }
336
337    fn start_block(
338        &mut self,
339        index: u32,
340        kind: ContentBlockKind,
341    ) -> Result<Option<ModelResponse>, ModelError> {
342        self.require_started()?;
343        self.require_unused_index(index)?;
344        self.open_blocks
345            .insert(index, PartialBlock::from_kind(kind));
346        Ok(None)
347    }
348
349    fn complete_block(&mut self, index: u32) -> Result<Option<ModelResponse>, ModelError> {
350        self.require_started()?;
351        let block = self
352            .open_blocks
353            .remove(&index)
354            .ok_or_else(|| state_error(format!("content block {index} is not open")))?;
355        self.content.insert(index, block.complete()?);
356        Ok(None)
357    }
358
359    fn complete_part(
360        &mut self,
361        index: u32,
362        part: ContentPart,
363    ) -> Result<Option<ModelResponse>, ModelError> {
364        self.require_started()?;
365        self.require_unused_index(index)?;
366        self.content.insert(index, part);
367        Ok(None)
368    }
369
370    fn complete(
371        &mut self,
372        finish_reason: FinishReason,
373        provider_metadata: BTreeMap<String, Value>,
374    ) -> Result<Option<ModelResponse>, ModelError> {
375        self.require_started()?;
376        if !self.open_blocks.is_empty() {
377            let open = self
378                .open_blocks
379                .keys()
380                .map(u32::to_string)
381                .collect::<Vec<_>>()
382                .join(", ");
383            return Err(state_error(format!(
384                "response completed with open content blocks: {open}"
385            )));
386        }
387        self.completed = true;
388        let model = self
389            .model
390            .clone()
391            .ok_or_else(|| state_error("response model is missing"))?;
392        Ok(Some(ModelResponse {
393            id: self.id.clone(),
394            model,
395            content: std::mem::take(&mut self.content).into_values().collect(),
396            finish_reason,
397            usage: self.usage,
398            warnings: std::mem::take(&mut self.warnings),
399            provider_metadata,
400            provider_events: std::mem::take(&mut self.provider_events),
401        }))
402    }
403
404    fn require_started(&self) -> Result<(), ModelError> {
405        if self.started {
406            Ok(())
407        } else {
408            Err(state_error("received content before response start"))
409        }
410    }
411
412    fn require_unused_index(&self, index: u32) -> Result<(), ModelError> {
413        if self.open_blocks.contains_key(&index) || self.content.contains_key(&index) {
414            Err(state_error(format!(
415                "content block index {index} was already used"
416            )))
417        } else {
418            Ok(())
419        }
420    }
421
422    fn open_block_mut(&mut self, index: u32) -> Result<&mut PartialBlock, ModelError> {
423        self.require_started()?;
424        self.open_blocks
425            .get_mut(&index)
426            .ok_or_else(|| state_error(format!("content block {index} is not open")))
427    }
428}
429
430fn wrong_delta(index: u32, delta: &str) -> ModelError {
431    state_error(format!(
432        "{delta} delta does not match content block {index}"
433    ))
434}
435
436fn state_error(message: impl Into<String>) -> ModelError {
437    ModelError::local(ModelErrorKind::StreamState, message)
438}
439
440#[cfg(test)]
441mod tests {
442    use std::collections::BTreeMap;
443
444    use super::{ContentBlockKind, ModelStreamAccumulator, ModelStreamEvent, ProviderEvent};
445    use crate::{
446        ContentPart, FinishReason, ModelErrorKind, ModelRef, ModelUsage, ModelWarning, ToolCall,
447    };
448
449    fn started() -> ModelStreamEvent {
450        ModelStreamEvent::ResponseStarted {
451            id: Some("response-1".into()),
452            model: ModelRef::new("test", "model"),
453        }
454    }
455
456    fn completed() -> ModelStreamEvent {
457        ModelStreamEvent::ResponseCompleted {
458            finish_reason: FinishReason::Stop,
459            provider_metadata: BTreeMap::new(),
460        }
461    }
462
463    #[test]
464    fn accumulates_ordered_text_and_tool_calls() {
465        let mut accumulator = ModelStreamAccumulator::new();
466        let events = [
467            started(),
468            ModelStreamEvent::ContentBlockStarted {
469                index: 1,
470                kind: ContentBlockKind::ToolCall {
471                    id: "call-1".into(),
472                    name: "search".into(),
473                },
474            },
475            ModelStreamEvent::ToolArgumentsDelta {
476                index: 1,
477                json: "{\"query\":".into(),
478            },
479            ModelStreamEvent::ContentBlockStarted {
480                index: 0,
481                kind: ContentBlockKind::Text,
482            },
483            ModelStreamEvent::TextDelta {
484                index: 0,
485                text: "I will search.".into(),
486            },
487            ModelStreamEvent::ToolArgumentsDelta {
488                index: 1,
489                json: "\"rust\"}".into(),
490            },
491            ModelStreamEvent::ContentBlockCompleted { index: 0 },
492            ModelStreamEvent::ContentBlockCompleted { index: 1 },
493            ModelStreamEvent::UsageUpdated {
494                usage: ModelUsage {
495                    input_tokens: 5,
496                    output_tokens: 3,
497                    ..ModelUsage::default()
498                },
499            },
500            completed(),
501        ];
502
503        let response = events
504            .into_iter()
505            .find_map(|event| accumulator.push(event).unwrap())
506            .unwrap();
507
508        assert_eq!(response.content[0], ContentPart::text("I will search."));
509        assert_eq!(
510            response.content[1],
511            ContentPart::ToolCall(ToolCall {
512                id: "call-1".into(),
513                name: "search".into(),
514                arguments: serde_json::json!({"query": "rust"}),
515                raw_arguments: Some("{\"query\":\"rust\"}".into()),
516                metadata: BTreeMap::new(),
517            })
518        );
519        assert_eq!(response.usage.input_tokens, 5);
520    }
521
522    #[test]
523    fn preserves_provider_events_and_warnings() {
524        let mut accumulator = ModelStreamAccumulator::new();
525        accumulator.push(started()).unwrap();
526        accumulator
527            .push(ModelStreamEvent::Provider {
528                event: ProviderEvent {
529                    provider: "test".into(),
530                    name: "ping".into(),
531                    payload: serde_json::json!({"alive": true}),
532                },
533            })
534            .unwrap();
535        accumulator
536            .push(ModelStreamEvent::Warning {
537                warning: ModelWarning {
538                    code: "emulated".into(),
539                    message: "structured output was emulated".into(),
540                    metadata: BTreeMap::new(),
541                },
542            })
543            .unwrap();
544        let response = accumulator.push(completed()).unwrap().unwrap();
545
546        assert_eq!(response.provider_events.len(), 1);
547        assert_eq!(response.provider_events[0].kind, "ping");
548        assert_eq!(response.warnings.len(), 1);
549    }
550
551    #[test]
552    fn rejects_delta_without_matching_open_block() {
553        let mut accumulator = ModelStreamAccumulator::new();
554        accumulator.push(started()).unwrap();
555
556        let error = accumulator
557            .push(ModelStreamEvent::TextDelta {
558                index: 4,
559                text: "orphan".into(),
560            })
561            .unwrap_err();
562
563        assert_eq!(error.kind, ModelErrorKind::StreamState);
564    }
565
566    #[test]
567    fn rejects_completion_with_open_blocks() {
568        let mut accumulator = ModelStreamAccumulator::new();
569        accumulator.push(started()).unwrap();
570        accumulator
571            .push(ModelStreamEvent::ContentBlockStarted {
572                index: 0,
573                kind: ContentBlockKind::Text,
574            })
575            .unwrap();
576
577        let error = accumulator.push(completed()).unwrap_err();
578
579        assert_eq!(error.kind, ModelErrorKind::StreamState);
580    }
581
582    #[test]
583    fn rejects_malformed_tool_arguments() {
584        let mut accumulator = ModelStreamAccumulator::new();
585        accumulator.push(started()).unwrap();
586        accumulator
587            .push(ModelStreamEvent::ContentBlockStarted {
588                index: 0,
589                kind: ContentBlockKind::ToolCall {
590                    id: "bad".into(),
591                    name: "tool".into(),
592                },
593            })
594            .unwrap();
595        accumulator
596            .push(ModelStreamEvent::ToolArgumentsDelta {
597                index: 0,
598                json: "{invalid".into(),
599            })
600            .unwrap();
601
602        let error = accumulator
603            .push(ModelStreamEvent::ContentBlockCompleted { index: 0 })
604            .unwrap_err();
605
606        assert_eq!(error.kind, ModelErrorKind::MalformedToolArguments);
607    }
608}