1use 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#[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 #[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 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}