Skip to main content

runifold_model/
stream.rs

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