Skip to main content

ferrin_core/middleware/builtin/
simulate_streaming.rs

1//! Streaming simulated from a complete response.
2
3use ferrin_spec::BoxFuture;
4use ferrin_spec::CallOptions;
5use ferrin_spec::Content;
6use ferrin_spec::PartId;
7use ferrin_spec::StreamPart;
8use ferrin_spec::StreamResult;
9use ferrin_spec::error::ProviderError;
10use ferrin_spec::language_model::GenerateResult;
11use futures_util::stream;
12
13use crate::middleware::LanguageModelMiddleware;
14use crate::middleware::MiddlewareContext;
15use crate::middleware::StreamNext;
16
17/// Middleware created by [`simulate_streaming`].
18#[derive(Debug, Clone, Copy, Default)]
19pub struct SimulateStreaming;
20
21/// Serves `do_stream` by calling `do_generate` on the wrapped model and
22/// expanding the result into `stream-start`, `response-metadata`, one
23/// start/delta/end triple per text or reasoning part (empty text parts are
24/// skipped), the other parts as is, and `finish`.
25#[must_use]
26pub fn simulate_streaming() -> SimulateStreaming {
27    SimulateStreaming
28}
29
30/// Expands a complete result into stream parts.
31#[must_use]
32pub fn simulate_parts(result: &GenerateResult) -> Vec<StreamPart> {
33    let mut parts = vec![
34        StreamPart::StreamStart {
35            warnings: result.warnings.clone(),
36        },
37        StreamPart::ResponseMetadata {
38            id: result.response.id.clone(),
39            timestamp: result.response.timestamp,
40            model_id: result.response.model_id.clone(),
41        },
42    ];
43    let mut next_id = 0u32;
44    for part in &result.content {
45        match part {
46            Content::Text { text, .. } => {
47                if text.is_empty() {
48                    continue;
49                }
50                let id = PartId::new(next_id.to_string());
51                next_id += 1;
52                parts.push(StreamPart::TextStart {
53                    id: id.clone(),
54                    provider_metadata: None,
55                });
56                parts.push(StreamPart::TextDelta {
57                    id: id.clone(),
58                    delta: text.clone(),
59                    provider_metadata: None,
60                });
61                parts.push(StreamPart::TextEnd {
62                    id,
63                    provider_metadata: None,
64                });
65            }
66            Content::Reasoning {
67                text,
68                provider_metadata,
69            } => {
70                let id = PartId::new(next_id.to_string());
71                next_id += 1;
72                parts.push(StreamPart::ReasoningStart {
73                    id: id.clone(),
74                    provider_metadata: provider_metadata.clone(),
75                });
76                parts.push(StreamPart::ReasoningDelta {
77                    id: id.clone(),
78                    delta: text.clone(),
79                    provider_metadata: None,
80                });
81                parts.push(StreamPart::ReasoningEnd {
82                    id,
83                    provider_metadata: None,
84                });
85            }
86            Content::ReasoningFile {
87                data,
88                media_type,
89                provider_metadata,
90            } => parts.push(StreamPart::ReasoningFile {
91                data: data.clone(),
92                media_type: media_type.clone(),
93                provider_metadata: provider_metadata.clone(),
94            }),
95            Content::File {
96                data,
97                media_type,
98                filename,
99                provider_metadata,
100            } => parts.push(StreamPart::File {
101                data: data.clone(),
102                media_type: media_type.clone(),
103                filename: filename.clone(),
104                provider_metadata: provider_metadata.clone(),
105            }),
106            Content::Custom {
107                kind,
108                provider_metadata,
109            } => parts.push(StreamPart::Custom {
110                kind: kind.clone(),
111                provider_metadata: provider_metadata.clone(),
112            }),
113            Content::Source(source) => parts.push(StreamPart::Source(source.clone())),
114            Content::ToolCall(call) => parts.push(StreamPart::ToolCall(call.clone())),
115            Content::ToolResult(result) => parts.push(StreamPart::ToolResult(result.clone())),
116            Content::ToolApprovalRequest {
117                approval_id,
118                tool_call_id,
119                provider_metadata,
120            } => parts.push(StreamPart::ToolApprovalRequest {
121                approval_id: approval_id.clone(),
122                tool_call_id: tool_call_id.clone(),
123                provider_metadata: provider_metadata.clone(),
124            }),
125            #[allow(unreachable_patterns, reason = "Content is non-exhaustive")]
126            _ => {}
127        }
128    }
129    parts.push(StreamPart::Finish {
130        finish_reason: result.finish_reason.clone(),
131        usage: result.usage.clone(),
132        provider_metadata: result.provider_metadata.clone(),
133    });
134    parts
135}
136
137impl LanguageModelMiddleware for SimulateStreaming {
138    fn wrap_stream<'a>(
139        &'a self,
140        options: CallOptions,
141        _next: StreamNext<'a>,
142        ctx: MiddlewareContext<'a>,
143    ) -> BoxFuture<'a, Result<StreamResult, ProviderError>> {
144        Box::pin(async move {
145            let result = ctx.model.do_generate(options).await?;
146            let parts = simulate_parts(&result);
147            Ok(StreamResult {
148                stream: Box::pin(stream::iter(parts)),
149                request: result.request,
150                response: result.response,
151            })
152        })
153    }
154}