1use std::collections::HashMap;
52use std::fmt;
53use std::pin::Pin;
54use std::task::{Context, Poll};
55use std::time::Duration;
56
57use futures::stream::{Stream, StreamExt};
58use serde::{Deserialize, Serialize};
59
60use crate::error::ProviderError;
61use crate::ids::{CallId, ModelKey, ProviderKey, RequestId};
62use crate::request::{ContentPart, ToolCall};
63use crate::response::{FinishReason, ModelResponse, ResponseWarning, TokenUsage};
64
65#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
67#[serde(tag = "type", rename_all = "snake_case")]
68#[non_exhaustive]
69pub enum StreamEvent {
70 TextDelta {
72 text: String,
74 },
75 ToolCallStart {
77 id: CallId,
79 name: String,
81 },
82 ToolCallDelta {
84 id: CallId,
86 arguments_fragment: String,
88 },
89 ToolCallEnd {
91 id: CallId,
93 },
94 Usage {
97 usage: TokenUsage,
99 },
100 ResponseId {
107 id: String,
109 },
110 Warning {
118 warning: ResponseWarning,
120 },
121 Finish {
123 reason: FinishReason,
125 },
126}
127
128impl StreamEvent {
129 #[must_use]
131 pub fn text(text: impl Into<String>) -> Self {
132 Self::TextDelta { text: text.into() }
133 }
134
135 #[must_use]
137 pub fn tool_call_start(id: impl Into<CallId>, name: impl Into<String>) -> Self {
138 Self::ToolCallStart {
139 id: id.into(),
140 name: name.into(),
141 }
142 }
143
144 #[must_use]
146 pub fn tool_call_delta(id: impl Into<CallId>, fragment: impl Into<String>) -> Self {
147 Self::ToolCallDelta {
148 id: id.into(),
149 arguments_fragment: fragment.into(),
150 }
151 }
152
153 #[must_use]
155 pub fn tool_call_end(id: impl Into<CallId>) -> Self {
156 Self::ToolCallEnd { id: id.into() }
157 }
158
159 #[must_use]
161 pub fn response_id(id: impl Into<String>) -> Self {
162 Self::ResponseId { id: id.into() }
163 }
164
165 #[must_use]
167 pub const fn warning(warning: ResponseWarning) -> Self {
168 Self::Warning { warning }
169 }
170
171 #[must_use]
173 pub const fn kind(&self) -> &'static str {
174 match self {
175 Self::TextDelta { .. } => "text_delta",
176 Self::ToolCallStart { .. } => "tool_call_start",
177 Self::ToolCallDelta { .. } => "tool_call_delta",
178 Self::ToolCallEnd { .. } => "tool_call_end",
179 Self::Usage { .. } => "usage",
180 Self::ResponseId { .. } => "response_id",
181 Self::Warning { .. } => "warning",
182 Self::Finish { .. } => "finish",
183 }
184 }
185
186 #[must_use]
189 pub const fn is_user_visible(&self) -> bool {
190 matches!(self, Self::TextDelta { .. })
191 }
192}
193
194pub type StreamItem = Result<StreamEvent, ProviderError>;
196
197pub struct ModelStream {
203 inner: Pin<Box<dyn Stream<Item = StreamItem> + Send>>,
204}
205
206impl ModelStream {
207 #[must_use]
209 pub fn new(stream: impl Stream<Item = StreamItem> + Send + 'static) -> Self {
210 Self {
211 inner: Box::pin(stream),
212 }
213 }
214
215 #[must_use]
217 pub fn from_events(events: Vec<StreamEvent>) -> Self {
218 Self::new(futures::stream::iter(events.into_iter().map(Ok)))
219 }
220
221 #[must_use]
223 pub fn from_items(items: Vec<StreamItem>) -> Self {
224 Self::new(futures::stream::iter(items))
225 }
226
227 #[must_use]
229 pub fn failed(error: ProviderError) -> Self {
230 Self::from_items(vec![Err(error)])
231 }
232
233 pub async fn collect_items(self) -> Vec<StreamItem> {
235 self.inner.collect().await
236 }
237}
238
239impl fmt::Debug for ModelStream {
240 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
241 f.write_str("ModelStream(..)")
242 }
243}
244
245impl Stream for ModelStream {
246 type Item = StreamItem;
247
248 fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
249 self.inner.as_mut().poll_next(cx)
250 }
251
252 fn size_hint(&self) -> (usize, Option<usize>) {
253 self.inner.size_hint()
254 }
255}
256
257#[derive(Debug, Clone)]
259struct PendingCall {
260 order: usize,
261 name: String,
262 fragments: String,
263 ended: bool,
264}
265
266#[derive(Debug, Clone)]
272pub struct StreamAccumulator {
273 request_id: RequestId,
274 provider: ProviderKey,
275 model: ModelKey,
276 raw_id: Option<String>,
277 latency: Duration,
278 text: String,
279 calls: HashMap<CallId, PendingCall>,
280 usage: TokenUsage,
281 finish: Option<FinishReason>,
282 warnings: Vec<ResponseWarning>,
283}
284
285impl StreamAccumulator {
286 #[must_use]
288 pub fn new(
289 request_id: RequestId,
290 provider: impl Into<ProviderKey>,
291 model: impl Into<ModelKey>,
292 ) -> Self {
293 Self {
294 request_id,
295 provider: provider.into(),
296 model: model.into(),
297 raw_id: None,
298 latency: Duration::ZERO,
299 text: String::new(),
300 calls: HashMap::new(),
301 usage: TokenUsage::none(),
302 finish: None,
303 warnings: Vec::new(),
304 }
305 }
306
307 #[must_use]
313 pub fn with_raw_id(mut self, raw_id: impl Into<String>) -> Self {
314 self.raw_id = Some(raw_id.into());
315 self
316 }
317
318 #[must_use]
320 pub const fn with_latency(mut self, latency: Duration) -> Self {
321 self.latency = latency;
322 self
323 }
324
325 #[must_use]
331 pub fn with_warning(mut self, warning: ResponseWarning) -> Self {
332 self.warnings.push(warning);
333 self
334 }
335
336 pub fn push(&mut self, event: StreamEvent) -> Result<(), ProviderError> {
345 match event {
346 StreamEvent::TextDelta { text } => self.text.push_str(&text),
347 StreamEvent::ToolCallStart { id, name } => {
348 let order = self.calls.len();
349 if self.calls.contains_key(&id) {
350 return Err(self.malformed("tool_call_started_twice"));
351 }
352 self.calls.insert(
353 id,
354 PendingCall {
355 order,
356 name,
357 fragments: String::new(),
358 ended: false,
359 },
360 );
361 }
362 StreamEvent::ToolCallDelta {
363 id,
364 arguments_fragment,
365 } => {
366 let Some(call) = self.calls.get_mut(&id) else {
367 return Err(self.malformed("tool_call_delta_without_start"));
368 };
369 if call.ended {
370 return Err(self.malformed("tool_call_delta_after_end"));
371 }
372 call.fragments.push_str(&arguments_fragment);
373 }
374 StreamEvent::ToolCallEnd { id } => {
375 let Some(call) = self.calls.get_mut(&id) else {
376 return Err(self.malformed("tool_call_end_without_start"));
377 };
378 call.ended = true;
379 }
380 StreamEvent::Usage { usage } => self.usage = usage,
381 StreamEvent::ResponseId { id } => self.raw_id = Some(id),
382 StreamEvent::Warning { warning } => {
383 if !self.warnings.contains(&warning) {
384 self.warnings.push(warning);
385 }
386 }
387 StreamEvent::Finish { reason } => {
388 if self.finish.is_some() {
389 return Err(self.malformed("duplicate_finish"));
390 }
391 self.finish = Some(reason);
392 }
393 }
394 Ok(())
395 }
396
397 pub fn finish(self) -> Result<ModelResponse, ProviderError> {
405 let Some(finish) = self.finish else {
406 return Err(self.malformed("stream_ended_without_finish"));
407 };
408
409 let mut content = Vec::with_capacity(self.calls.len() + 1);
410 if !self.text.is_empty() {
411 content.push(ContentPart::text(self.text.clone()));
412 }
413
414 let mut ordered: Vec<(CallId, PendingCall)> = self.calls.clone().into_iter().collect();
415 ordered.sort_by_key(|(_, call)| call.order);
416 for (id, call) in ordered {
417 let raw = call.fragments.trim();
418 let arguments: serde_json::Value = if raw.is_empty() {
419 serde_json::Value::Object(serde_json::Map::new())
420 } else {
421 serde_json::from_str(raw)
422 .map_err(|_| self.malformed("tool_call_arguments_not_json"))?
423 };
424 content.push(ContentPart::ToolCall(ToolCall::new(
425 id, call.name, arguments,
426 )));
427 }
428
429 let mut warnings = self.warnings.clone();
430 warnings.push(ResponseWarning::Reconstructed);
431 if self.usage.is_unreported() {
432 warnings.push(ResponseWarning::UsageUnreported);
433 }
434
435 Ok(ModelResponse {
436 request_id: self.request_id,
437 provider: self.provider,
438 model: self.model,
439 content,
440 finish,
441 usage: self.usage,
442 raw_id: self.raw_id,
443 latency: self.latency,
444 warnings,
445 })
446 }
447
448 fn malformed(&self, code: &str) -> ProviderError {
450 ProviderError::malformed(code).with_model(&crate::ids::ModelRef {
451 provider: self.provider.clone(),
452 model: self.model.clone(),
453 })
454 }
455}
456
457pub async fn reconstruct(
470 mut stream: ModelStream,
471 mut seed: StreamAccumulator,
472) -> Result<ModelResponse, ProviderError> {
473 while let Some(item) = stream.next().await {
474 seed.push(item?)?;
475 }
476 seed.finish()
477}
478
479#[cfg(test)]
480mod tests {
481 use super::*;
482 use serde_json::json;
483
484 fn seed() -> StreamAccumulator {
485 StreamAccumulator::new(RequestId::nil(), "openai", "gpt-4o")
486 }
487
488 async fn rebuild(events: Vec<StreamEvent>) -> Result<ModelResponse, ProviderError> {
489 reconstruct(ModelStream::from_events(events), seed()).await
490 }
491
492 #[tokio::test]
493 async fn text_deltas_concatenate_in_order() {
494 let response = rebuild(vec![
495 StreamEvent::text("Ho "),
496 StreamEvent::text("preparato "),
497 StreamEvent::text("la modifica."),
498 StreamEvent::Usage {
499 usage: TokenUsage::new(10, 4),
500 },
501 StreamEvent::Finish {
502 reason: FinishReason::Stop,
503 },
504 ])
505 .await
506 .unwrap();
507 assert_eq!(response.text(), "Ho preparato la modifica.");
508 assert_eq!(response.content.len(), 1, "one text part, not three");
509 assert_eq!(response.usage, TokenUsage::new(10, 4));
510 assert_eq!(response.finish, FinishReason::Stop);
511 assert!(response.warnings.contains(&ResponseWarning::Reconstructed));
512 }
513
514 #[tokio::test]
515 async fn tool_call_fragments_concatenate_per_id() {
516 let response = rebuild(vec![
517 StreamEvent::tool_call_start("call_a", "plan"),
518 StreamEvent::tool_call_start("call_b", "plan"),
519 StreamEvent::tool_call_delta("call_a", "{\"n\":"),
520 StreamEvent::tool_call_delta("call_b", "{\"n\":"),
521 StreamEvent::tool_call_delta("call_a", "1}"),
522 StreamEvent::tool_call_delta("call_b", "2}"),
523 StreamEvent::tool_call_end("call_a"),
524 StreamEvent::tool_call_end("call_b"),
525 StreamEvent::Finish {
526 reason: FinishReason::ToolCalls,
527 },
528 ])
529 .await
530 .unwrap();
531 let calls = response.tool_calls();
532 assert_eq!(calls.len(), 2);
533 assert_eq!(calls[0].id.as_str(), "call_a");
534 assert_eq!(calls[0].arguments, json!({"n": 1}));
535 assert_eq!(calls[1].id.as_str(), "call_b");
536 assert_eq!(calls[1].arguments, json!({"n": 2}));
537 }
538
539 #[tokio::test]
540 async fn text_comes_before_tool_calls() {
541 let response = rebuild(vec![
542 StreamEvent::tool_call_start("c", "plan"),
543 StreamEvent::text("preambolo"),
544 StreamEvent::tool_call_delta("c", "{}"),
545 StreamEvent::tool_call_end("c"),
546 StreamEvent::Finish {
547 reason: FinishReason::ToolCalls,
548 },
549 ])
550 .await
551 .unwrap();
552 assert_eq!(response.content.len(), 2);
553 assert_eq!(response.content[0].kind(), "text");
554 assert_eq!(response.content[1].kind(), "tool_call");
555 }
556
557 #[tokio::test]
558 async fn a_call_with_no_fragments_gets_an_empty_object() {
559 let response = rebuild(vec![
560 StreamEvent::tool_call_start("c", "ping"),
561 StreamEvent::tool_call_end("c"),
562 StreamEvent::Finish {
563 reason: FinishReason::ToolCalls,
564 },
565 ])
566 .await
567 .unwrap();
568 assert_eq!(response.tool_calls()[0].arguments, json!({}));
569 }
570
571 #[tokio::test]
572 async fn reassembly_failures_are_malformed_and_never_partial() {
573 let cases = [
574 (
575 vec![
576 StreamEvent::tool_call_delta("ghost", "{}"),
577 StreamEvent::Finish {
578 reason: FinishReason::ToolCalls,
579 },
580 ],
581 "tool_call_delta_without_start",
582 ),
583 (
584 vec![
585 StreamEvent::tool_call_end("ghost"),
586 StreamEvent::Finish {
587 reason: FinishReason::ToolCalls,
588 },
589 ],
590 "tool_call_end_without_start",
591 ),
592 (
593 vec![
594 StreamEvent::tool_call_start("c", "plan"),
595 StreamEvent::tool_call_start("c", "plan"),
596 StreamEvent::Finish {
597 reason: FinishReason::ToolCalls,
598 },
599 ],
600 "tool_call_started_twice",
601 ),
602 (
603 vec![
604 StreamEvent::tool_call_start("c", "plan"),
605 StreamEvent::tool_call_end("c"),
606 StreamEvent::tool_call_delta("c", "{}"),
607 StreamEvent::Finish {
608 reason: FinishReason::ToolCalls,
609 },
610 ],
611 "tool_call_delta_after_end",
612 ),
613 (
614 vec![
615 StreamEvent::Finish {
616 reason: FinishReason::Stop,
617 },
618 StreamEvent::Finish {
619 reason: FinishReason::Stop,
620 },
621 ],
622 "duplicate_finish",
623 ),
624 (
625 vec![StreamEvent::text("ciao")],
626 "stream_ended_without_finish",
627 ),
628 (
629 vec![
630 StreamEvent::tool_call_start("c", "plan"),
631 StreamEvent::tool_call_delta("c", "{\"n\":"),
632 StreamEvent::Finish {
633 reason: FinishReason::ToolCalls,
634 },
635 ],
636 "tool_call_arguments_not_json",
637 ),
638 ];
639 for (events, expected) in cases {
640 let error = rebuild(events).await.unwrap_err();
641 assert_eq!(
642 error.code().map(|code| code.as_str().to_owned()),
643 Some(expected.to_owned()),
644 "{error}"
645 );
646 assert!(matches!(
647 error.kind(),
648 crate::error::ProviderErrorKind::Malformed
649 ));
650 assert_eq!(error.provider().map(ProviderKey::as_str), Some("openai"));
651 }
652 }
653
654 #[tokio::test]
655 async fn an_error_item_stops_reassembly_immediately() {
656 let stream = ModelStream::from_items(vec![
657 Ok(StreamEvent::text("half")),
658 Err(ProviderError::transport("connection_reset")),
659 Ok(StreamEvent::Finish {
660 reason: FinishReason::Stop,
661 }),
662 ]);
663 let error = reconstruct(stream, seed()).await.unwrap_err();
664 assert!(matches!(
665 error.kind(),
666 crate::error::ProviderErrorKind::Transport
667 ));
668 }
669
670 #[tokio::test]
671 async fn a_failed_stream_surfaces_its_error() {
672 let stream = ModelStream::failed(ProviderError::rate_limited(None));
673 let error = reconstruct(stream, seed()).await.unwrap_err();
674 assert!(matches!(
675 error.kind(),
676 crate::error::ProviderErrorKind::RateLimited { .. }
677 ));
678 assert_eq!(
679 format!("{:?}", ModelStream::from_events(vec![])),
680 "ModelStream(..)"
681 );
682 }
683
684 #[tokio::test]
685 async fn the_seed_supplies_identity_and_timing() {
686 let seed = StreamAccumulator::new(RequestId::nil(), "anthropic", "claude")
687 .with_raw_id("msg_01")
688 .with_latency(Duration::from_millis(250))
689 .with_warning(ResponseWarning::SynthesizedCallIds);
690 let response = reconstruct(
691 ModelStream::from_events(vec![
692 StreamEvent::text("x"),
693 StreamEvent::Finish {
694 reason: FinishReason::Stop,
695 },
696 ]),
697 seed,
698 )
699 .await
700 .unwrap();
701 assert_eq!(response.provider.as_str(), "anthropic");
702 assert_eq!(response.raw_id.as_deref(), Some("msg_01"));
703 assert_eq!(response.latency, Duration::from_millis(250));
704 assert!(
705 response
706 .warnings
707 .contains(&ResponseWarning::SynthesizedCallIds)
708 );
709 assert!(
710 response
711 .warnings
712 .contains(&ResponseWarning::UsageUnreported)
713 );
714 }
715
716 #[tokio::test]
717 async fn the_stream_can_carry_the_identifier_and_the_warning_itself() {
718 let response = rebuild(vec![
722 StreamEvent::response_id("chatcmpl-1"),
723 StreamEvent::warning(ResponseWarning::FeatureDropped {
724 feature: "stop_sequences".to_owned(),
725 }),
726 StreamEvent::text("ok"),
727 StreamEvent::response_id("chatcmpl-final"),
728 StreamEvent::Finish {
729 reason: FinishReason::Stop,
730 },
731 ])
732 .await
733 .unwrap();
734 assert_eq!(
735 response.raw_id.as_deref(),
736 Some("chatcmpl-final"),
737 "a later identifier replaces an earlier one"
738 );
739 assert!(
740 response
741 .warnings
742 .contains(&ResponseWarning::FeatureDropped {
743 feature: "stop_sequences".to_owned(),
744 })
745 );
746 assert!(response.warnings.contains(&ResponseWarning::Reconstructed));
747 }
748
749 #[tokio::test]
750 async fn a_warning_the_seed_and_the_stream_both_carry_is_recorded_once() {
751 let dropped = ResponseWarning::FeatureDropped {
752 feature: "cache_hint".to_owned(),
753 };
754 let seed = StreamAccumulator::new(RequestId::nil(), "openai", "gpt-4o")
755 .with_warning(dropped.clone());
756 let response = reconstruct(
757 ModelStream::from_events(vec![
758 StreamEvent::warning(dropped.clone()),
759 StreamEvent::Finish {
760 reason: FinishReason::Stop,
761 },
762 ]),
763 seed,
764 )
765 .await
766 .unwrap();
767 assert_eq!(
768 response
769 .warnings
770 .iter()
771 .filter(|warning| **warning == dropped)
772 .count(),
773 1
774 );
775 }
776
777 #[test]
778 fn events_round_trip_and_only_text_is_user_visible() {
779 let events = [
780 StreamEvent::text("a"),
781 StreamEvent::tool_call_start("c", "plan"),
782 StreamEvent::tool_call_delta("c", "{}"),
783 StreamEvent::tool_call_end("c"),
784 StreamEvent::Usage {
785 usage: TokenUsage::new(1, 1),
786 },
787 StreamEvent::response_id("resp_1"),
788 StreamEvent::warning(ResponseWarning::UsageUnreported),
789 StreamEvent::Finish {
790 reason: FinishReason::Stop,
791 },
792 ];
793 let mut kinds: Vec<&str> = events.iter().map(StreamEvent::kind).collect();
794 kinds.sort_unstable();
795 kinds.dedup();
796 assert_eq!(kinds.len(), 8);
797 for event in &events {
798 let json = serde_json::to_string(event).unwrap();
799 let back: StreamEvent = serde_json::from_str(&json).unwrap();
800 assert_eq!(&back, event);
801 }
802 assert!(events[0].is_user_visible());
803 assert!(!events[1].is_user_visible());
804 }
805}