Skip to main content

ferrin_google/
stream.rs

1//! Streaming state machine for `streamGenerateContent?alt=sse`.
2
3use std::collections::HashSet;
4
5use ferrin_provider_util::http::ParseResult;
6use ferrin_provider_util::stream_driver::StreamMachine;
7use ferrin_spec::FinishReason;
8use ferrin_spec::FinishReasonKind;
9use ferrin_spec::JsonObject;
10use ferrin_spec::JsonValue;
11use ferrin_spec::PartId;
12use ferrin_spec::ProviderMetadata;
13use ferrin_spec::ToolCall;
14use ferrin_spec::ToolCallId;
15use ferrin_spec::language_model::Source;
16use ferrin_spec::language_model::StreamPart;
17
18use crate::api_types::FunctionCall;
19use crate::api_types::GenerateContentResponse;
20use crate::api_types::Part;
21use crate::api_types::UsageMetadata;
22use crate::json_accumulator::JsonAccumulator;
23use crate::output::OutputMapper;
24use crate::output::convert_usage;
25use crate::output::map_finish_reason;
26
27#[derive(Debug)]
28struct ActiveToolCall {
29    id: ToolCallId,
30    tool_name: String,
31    accumulator: JsonAccumulator,
32    provider_metadata: Option<ProviderMetadata>,
33}
34
35/// Stream state: open text/reasoning blocks, streamed function calls, usage
36/// and the metadata of the final `finish` part.
37#[derive(Debug)]
38pub struct GoogleStreamState {
39    mapper: OutputMapper,
40    finish_reason: FinishReason,
41    usage: Option<UsageMetadata>,
42    raw_usage: Option<JsonObject>,
43    provider_metadata: Option<ProviderMetadata>,
44    last_grounding_metadata: Option<JsonValue>,
45    last_url_context_metadata: Option<JsonValue>,
46    has_tool_calls: bool,
47    emitted_response_metadata: bool,
48    text_block: Option<PartId>,
49    reasoning_block: Option<PartId>,
50    block_counter: u64,
51    emitted_source_urls: HashSet<String>,
52    active_calls: Vec<ActiveToolCall>,
53}
54
55impl GoogleStreamState {
56    /// Creates the state for one stream.
57    #[must_use]
58    pub fn new(mapper: OutputMapper) -> Self {
59        Self {
60            mapper,
61            finish_reason: FinishReason::new(FinishReasonKind::Other),
62            usage: None,
63            raw_usage: None,
64            provider_metadata: None,
65            last_grounding_metadata: None,
66            last_url_context_metadata: None,
67            has_tool_calls: false,
68            emitted_response_metadata: false,
69            text_block: None,
70            reasoning_block: None,
71            block_counter: 0,
72            emitted_source_urls: HashSet::new(),
73            active_calls: Vec::new(),
74        }
75    }
76
77    fn next_block_id(&mut self) -> PartId {
78        let id = PartId::new(self.block_counter.to_string());
79        self.block_counter += 1;
80        id
81    }
82
83    fn end_text(&mut self, parts: &mut Vec<StreamPart>) {
84        if let Some(id) = self.text_block.take() {
85            parts.push(StreamPart::TextEnd {
86                id,
87                provider_metadata: None,
88            });
89        }
90    }
91
92    fn end_reasoning(&mut self, parts: &mut Vec<StreamPart>) {
93        if let Some(id) = self.reasoning_block.take() {
94            parts.push(StreamPart::ReasoningEnd {
95                id,
96                provider_metadata: None,
97            });
98        }
99    }
100
101    fn finish_active_call(&mut self, parts: &mut Vec<StreamPart>) {
102        let Some(active) = self.active_calls.pop() else {
103            return;
104        };
105        let (final_json, closing_delta) = active.accumulator.finalize();
106        if !closing_delta.is_empty() {
107            parts.push(StreamPart::ToolInputDelta {
108                id: active.id.clone(),
109                delta: closing_delta,
110                provider_metadata: active.provider_metadata.clone(),
111            });
112        }
113        parts.push(StreamPart::ToolInputEnd {
114            id: active.id.clone(),
115            provider_metadata: active.provider_metadata.clone(),
116        });
117        let mut call = ToolCall::new(active.id, active.tool_name, final_json);
118        call.provider_metadata = active.provider_metadata;
119        parts.push(StreamPart::ToolCall(call));
120        self.has_tool_calls = true;
121    }
122
123    fn finish_metadata(
124        &self,
125        response: &GenerateContentResponse,
126        candidate_safety: Option<&JsonValue>,
127        finish_message: Option<&str>,
128    ) -> ProviderMetadata {
129        let mut object = JsonObject::new();
130        object.insert(
131            "promptFeedback".to_owned(),
132            response.prompt_feedback.clone().unwrap_or(JsonValue::Null),
133        );
134        object.insert(
135            "groundingMetadata".to_owned(),
136            self.last_grounding_metadata
137                .clone()
138                .unwrap_or(JsonValue::Null),
139        );
140        object.insert(
141            "urlContextMetadata".to_owned(),
142            self.last_url_context_metadata
143                .clone()
144                .unwrap_or(JsonValue::Null),
145        );
146        object.insert(
147            "safetyRatings".to_owned(),
148            candidate_safety.cloned().unwrap_or(JsonValue::Null),
149        );
150        object.insert(
151            "usageMetadata".to_owned(),
152            self.raw_usage
153                .clone()
154                .map_or(JsonValue::Null, JsonValue::Object),
155        );
156        object.insert(
157            "finishMessage".to_owned(),
158            finish_message.map_or(JsonValue::Null, JsonValue::from),
159        );
160        object.insert(
161            "serviceTier".to_owned(),
162            self.usage
163                .as_ref()
164                .and_then(|usage| usage.service_tier.clone())
165                .map_or(JsonValue::Null, JsonValue::from),
166        );
167        self.mapper.metadata(object)
168    }
169
170    fn text_part(
171        &mut self,
172        text: &str,
173        thought: bool,
174        signature: Option<&str>,
175        parts: &mut Vec<StreamPart>,
176    ) {
177        let metadata = self.mapper.signature_metadata(signature);
178        if text.is_empty() {
179            if let (Some(_), Some(id)) = (&metadata, &self.text_block) {
180                parts.push(StreamPart::TextDelta {
181                    id: id.clone(),
182                    delta: String::new(),
183                    provider_metadata: metadata,
184                });
185            }
186            return;
187        }
188        if thought {
189            self.end_text(parts);
190            if self.reasoning_block.is_none() {
191                let id = self.next_block_id();
192                self.reasoning_block = Some(id.clone());
193                parts.push(StreamPart::ReasoningStart {
194                    id,
195                    provider_metadata: metadata.clone(),
196                });
197            }
198            if let Some(id) = &self.reasoning_block {
199                parts.push(StreamPart::ReasoningDelta {
200                    id: id.clone(),
201                    delta: text.to_owned(),
202                    provider_metadata: metadata,
203                });
204            }
205        } else {
206            self.end_reasoning(parts);
207            if self.text_block.is_none() {
208                let id = self.next_block_id();
209                self.text_block = Some(id.clone());
210                parts.push(StreamPart::TextStart {
211                    id,
212                    provider_metadata: metadata.clone(),
213                });
214            }
215            if let Some(id) = &self.text_block {
216                parts.push(StreamPart::TextDelta {
217                    id: id.clone(),
218                    delta: text.to_owned(),
219                    provider_metadata: metadata,
220                });
221            }
222        }
223    }
224
225    fn content_parts(&mut self, part: &Part, parts: &mut Vec<StreamPart>) {
226        let signature = part.thought_signature.as_deref();
227        if let Some(code) = &part.executable_code
228            && code.code.is_some()
229        {
230            parts.push(StreamPart::ToolCall(self.mapper.code_execution_call(
231                code.language.as_deref(),
232                code.code.as_deref(),
233            )));
234        } else if let Some(result) = &part.code_execution_result {
235            parts.push(StreamPart::ToolResult(self.mapper.code_execution_result(
236                result.outcome.as_deref(),
237                result.output.as_deref(),
238            )));
239        } else if let Some(text) = &part.text {
240            self.text_part(text, part.thought == Some(true), signature, parts);
241        } else if let Some(inline) = &part.inline_data {
242            self.end_text(parts);
243            self.end_reasoning(parts);
244            match self.mapper.inline_file(
245                &inline.mime_type,
246                &inline.data,
247                part.thought == Some(true),
248                signature,
249            ) {
250                Ok(ferrin_spec::Content::ReasoningFile {
251                    data,
252                    media_type,
253                    provider_metadata,
254                }) => parts.push(StreamPart::ReasoningFile {
255                    data,
256                    media_type,
257                    provider_metadata,
258                }),
259                Ok(ferrin_spec::Content::File {
260                    data,
261                    media_type,
262                    filename,
263                    provider_metadata,
264                }) => parts.push(StreamPart::File {
265                    data,
266                    media_type,
267                    filename,
268                    provider_metadata,
269                }),
270                Ok(_) => {}
271                Err(error) => parts.push(StreamPart::error(&error)),
272            }
273        } else if let Some(call) = &part.tool_call {
274            parts.push(StreamPart::ToolCall(self.mapper.server_tool_call(
275                call.tool_type.as_deref(),
276                call.args.as_ref(),
277                call.id.as_deref(),
278                signature,
279            )));
280        } else if let Some(response) = &part.tool_response {
281            parts.push(StreamPart::ToolResult(
282                self.mapper.server_tool_response(response),
283            ));
284        }
285    }
286
287    fn complete_call(
288        &mut self,
289        call: &FunctionCall,
290        name: &str,
291        metadata: Option<ProviderMetadata>,
292        parts: &mut Vec<StreamPart>,
293    ) {
294        let mapped = self
295            .mapper
296            .function_call(call.id.as_deref(), name, call.args.as_ref(), None);
297        let id = mapped.tool_call_id.clone();
298        let tool_name = mapped.tool_name.clone();
299        parts.push(StreamPart::ToolInputStart {
300            id: id.clone(),
301            tool_name: tool_name.clone(),
302            provider_executed: false,
303            dynamic: false,
304            title: None,
305            provider_metadata: metadata.clone(),
306        });
307        if call.args.is_some() {
308            parts.push(StreamPart::ToolInputDelta {
309                id: id.clone(),
310                delta: mapped.input.clone(),
311                provider_metadata: metadata.clone(),
312            });
313        }
314        parts.push(StreamPart::ToolInputEnd {
315            id: id.clone(),
316            provider_metadata: metadata.clone(),
317        });
318        let mut tool_call = ToolCall::new(id, tool_name, mapped.input);
319        tool_call.provider_metadata = metadata;
320        parts.push(StreamPart::ToolCall(tool_call));
321        self.has_tool_calls = true;
322    }
323
324    fn function_call_part(&mut self, part: &Part, parts: &mut Vec<StreamPart>) {
325        let Some(call) = &part.function_call else {
326            return;
327        };
328        let metadata = self
329            .mapper
330            .signature_metadata(part.thought_signature.as_deref());
331        if call.is_streaming_fragment() {
332            if let Some(name) = &call.name {
333                let id = call
334                    .id
335                    .clone()
336                    .filter(|id| !id.is_empty())
337                    .unwrap_or_else(|| self.mapper.generate_id());
338                let id = ToolCallId::new(id);
339                let tool_name = self.mapper.custom_tool_name(name);
340                self.active_calls.push(ActiveToolCall {
341                    id: id.clone(),
342                    tool_name: tool_name.clone(),
343                    accumulator: JsonAccumulator::new(),
344                    provider_metadata: metadata.clone(),
345                });
346                parts.push(StreamPart::ToolInputStart {
347                    id,
348                    tool_name: tool_name.into(),
349                    provider_executed: false,
350                    dynamic: false,
351                    title: None,
352                    provider_metadata: metadata.clone(),
353                });
354            }
355            if let Some(partial_args) = &call.partial_args {
356                if let Some(active) = self.active_calls.last_mut() {
357                    let delta = active.accumulator.process(partial_args);
358                    if !delta.is_empty() {
359                        parts.push(StreamPart::ToolInputDelta {
360                            id: active.id.clone(),
361                            delta,
362                            provider_metadata: metadata,
363                        });
364                    }
365                }
366                if call.completes_stream() {
367                    self.finish_active_call(parts);
368                }
369            }
370        } else if call.is_terminal() {
371            if !self.active_calls.is_empty() {
372                self.finish_active_call(parts);
373            }
374        } else if let Some(name) = &call.name {
375            // A complete call (`name` + `args`) or a single-chunk call without
376            // arguments (`{name}`).
377            self.complete_call(call, name, metadata, parts);
378        }
379    }
380
381    fn handle_chunk(
382        &mut self,
383        value: &GenerateContentResponse,
384        raw: &JsonValue,
385    ) -> Vec<StreamPart> {
386        let mut parts = Vec::new();
387        if !self.emitted_response_metadata
388            && let Some(id) = &value.response_id
389        {
390            self.emitted_response_metadata = true;
391            parts.push(StreamPart::ResponseMetadata {
392                id: Some(id.clone()),
393                timestamp: None,
394                model_id: None,
395            });
396        }
397        if let Some(usage) = &value.usage_metadata {
398            self.usage = Some(usage.clone());
399            self.raw_usage = raw
400                .get("usageMetadata")
401                .and_then(JsonValue::as_object)
402                .cloned();
403        }
404        let Some(candidate) = value.candidate() else {
405            if let Some(reason) = value.block_reason() {
406                self.finish_reason =
407                    FinishReason::with_raw(FinishReasonKind::ContentFilter, reason);
408                self.provider_metadata = Some(self.finish_metadata(value, None, None));
409            }
410            return parts;
411        };
412        if candidate.grounding_metadata.is_some() {
413            self.last_grounding_metadata = candidate.grounding_metadata.clone();
414        }
415        if candidate.url_context_metadata.is_some() {
416            self.last_url_context_metadata = candidate.url_context_metadata.clone();
417        }
418        for source in self.mapper.sources(&candidate.grounding_chunks()) {
419            if let Source::Url { url, .. } = &source {
420                if self.emitted_source_urls.insert(url.clone()) {
421                    parts.push(StreamPart::Source(source));
422                }
423            } else {
424                parts.push(StreamPart::Source(source));
425            }
426        }
427        let candidate_parts = candidate.parts().to_vec();
428        for part in &candidate_parts {
429            self.content_parts(part, &mut parts);
430        }
431        for part in &candidate_parts {
432            self.function_call_part(part, &mut parts);
433        }
434        let block_reason = value.block_reason();
435        let prompt_blocked = candidate.finish_reason.is_none() && block_reason.is_some();
436        if let Some(raw_reason) = candidate.finish_reason.as_deref().or(block_reason) {
437            self.finish_reason = if prompt_blocked {
438                FinishReason::with_raw(FinishReasonKind::ContentFilter, raw_reason)
439            } else {
440                map_finish_reason(Some(raw_reason), self.has_tool_calls)
441            };
442            self.provider_metadata = Some(self.finish_metadata(
443                value,
444                candidate.safety_ratings.as_ref(),
445                candidate.finish_message.as_deref(),
446            ));
447        }
448        parts
449    }
450}
451
452impl StreamMachine for GoogleStreamState {
453    type Chunk = GenerateContentResponse;
454
455    fn handle(
456        &mut self,
457        chunk: ParseResult<GenerateContentResponse>,
458        include_raw: bool,
459    ) -> Vec<StreamPart> {
460        let mut parts = Vec::new();
461        match chunk {
462            ParseResult::Ok { value, raw } => {
463                if include_raw {
464                    parts.push(StreamPart::Raw {
465                        raw_value: raw.clone(),
466                    });
467                }
468                parts.extend(self.handle_chunk(&value, &raw));
469            }
470            ParseResult::Err { error, raw } => {
471                if include_raw {
472                    parts.push(StreamPart::Raw {
473                        raw_value: raw.map_or(JsonValue::Null, JsonValue::from),
474                    });
475                }
476                parts.push(StreamPart::error(&error));
477            }
478        }
479        parts
480    }
481
482    fn finish(mut self) -> Vec<StreamPart> {
483        let mut parts = Vec::new();
484        self.end_text(&mut parts);
485        self.end_reasoning(&mut parts);
486        parts.push(StreamPart::Finish {
487            finish_reason: self.finish_reason,
488            usage: convert_usage(self.usage.as_ref(), self.raw_usage),
489            provider_metadata: self.provider_metadata,
490        });
491        parts
492    }
493}