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#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
14#[serde(tag = "type", rename_all = "snake_case")]
15#[non_exhaustive]
16pub enum ContentBlockKind {
17 Text,
19 Reasoning {
21 signature: Option<String>,
23 redacted: bool,
25 },
26 ToolCall {
28 id: String,
30 name: String,
32 },
33 Refusal,
35 Image {
37 media_type: String,
39 },
40 Audio {
42 media_type: String,
44 },
45 Document {
47 media_type: String,
49 name: Option<String>,
51 },
52}
53
54#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
56pub struct ProviderEvent {
57 pub provider: String,
59 pub name: String,
61 pub payload: Value,
63}
64
65#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
67#[serde(tag = "type", rename_all = "snake_case")]
68#[non_exhaustive]
69pub enum ModelStreamEvent {
70 ResponseStarted {
72 id: Option<String>,
74 model: ModelRef,
76 },
77 ContentBlockStarted {
79 index: u32,
81 kind: ContentBlockKind,
83 },
84 TextDelta {
86 index: u32,
88 text: String,
90 },
91 ReasoningDelta {
93 index: u32,
95 text: String,
97 },
98 ReasoningSignatureDelta {
100 index: u32,
102 signature: String,
104 },
105 ToolArgumentsDelta {
107 index: u32,
109 json: String,
111 },
112 RefusalDelta {
114 index: u32,
116 text: String,
118 },
119 BinaryDelta {
121 index: u32,
123 data: String,
125 },
126 ContentBlockCompleted {
128 index: u32,
130 },
131 ContentPartCompleted {
133 index: u32,
135 part: ContentPart,
137 },
138 UsageUpdated {
140 usage: ModelUsage,
142 },
143 Warning {
145 warning: ModelWarning,
147 },
148 Heartbeat,
150 Provider {
152 event: ProviderEvent,
154 },
155 ResponseCompleted {
157 finish_reason: FinishReason,
159 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#[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 pub fn new() -> Self {
304 Self::default()
305 }
306
307 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}