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#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
13#[serde(tag = "type", rename_all = "snake_case")]
14#[non_exhaustive]
15pub enum ContentBlockKind {
16 Text,
18 Reasoning {
20 signature: Option<String>,
22 redacted: bool,
24 },
25 ToolCall {
27 id: String,
29 name: String,
31 },
32 Refusal,
34}
35
36#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
38pub struct ProviderEvent {
39 pub provider: String,
41 pub name: String,
43 pub payload: Value,
45}
46
47#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
49#[serde(tag = "type", rename_all = "snake_case")]
50#[non_exhaustive]
51pub enum ModelStreamEvent {
52 ResponseStarted {
54 id: Option<String>,
56 model: ModelRef,
58 },
59 ContentBlockStarted {
61 index: u32,
63 kind: ContentBlockKind,
65 },
66 TextDelta {
68 index: u32,
70 text: String,
72 },
73 ReasoningDelta {
75 index: u32,
77 text: String,
79 },
80 ReasoningSignatureDelta {
82 index: u32,
84 signature: String,
86 },
87 ToolArgumentsDelta {
89 index: u32,
91 json: String,
93 },
94 RefusalDelta {
96 index: u32,
98 text: String,
100 },
101 ContentBlockCompleted {
103 index: u32,
105 },
106 ContentPartCompleted {
108 index: u32,
110 part: ContentPart,
112 },
113 UsageUpdated {
115 usage: ModelUsage,
117 },
118 Warning {
120 warning: ModelWarning,
122 },
123 Heartbeat,
125 Provider {
127 event: ProviderEvent,
129 },
130 ResponseCompleted {
132 finish_reason: FinishReason,
134 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#[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 pub fn new() -> Self {
234 Self::default()
235 }
236
237 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}