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