1use anyhow::{Result, anyhow};
2
3use serde_json::{Value, json};
4
5use crate::api::common::{
6 apply_on_payload, build_http_client_for_target, finish_stream_error, invoke_on_response_from_reqwest,
7 is_request_aborted, merge_model_headers,
8};
9use crate::api::google_shared::{
10 convert_messages, convert_tools, is_thinking_part, map_stop_reason_finish, map_tool_choice,
11 retain_thought_signature,
12};
13use crate::api::simple_options::build_base_options;
14use crate::models::{calculate_cost, clamp_thinking_level};
15use crate::types::{
16 AssistantContentBlock, AssistantMessage, AssistantMessageEvent, Context, Model, ProviderStreams,
17 SimpleStreamOptions, StopReason, StreamOptions,
18};
19use crate::utils::event_stream::AssistantMessageEventStream;
20use crate::utils::sanitize_unicode::sanitize_surrogates;
21
22use super::sse::for_each_sse_json_event;
23
24#[derive(Clone, Default)]
25pub struct GoogleOptions {
26 pub base: StreamOptions,
27 pub tool_choice: Option<String>,
28 pub thinking: Option<GoogleThinkingConfig>,
29}
30
31#[derive(Debug, Clone)]
32pub struct GoogleThinkingConfig {
33 pub enabled: bool,
34 pub budget_tokens: Option<i32>,
35 pub level: Option<String>,
36}
37
38pub struct GoogleGenerativeAIApi;
39static TOOL_CALL_COUNTER: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
40
41impl ProviderStreams for GoogleGenerativeAIApi {
42 fn stream(&self, model: &Model, context: &Context, options: Option<StreamOptions>) -> AssistantMessageEventStream {
43 self.stream_with_options(
44 model,
45 context,
46 GoogleOptions {
47 base: options.unwrap_or_default(),
48 ..Default::default()
49 },
50 )
51 }
52
53 fn stream_simple(
54 &self,
55 model: &Model,
56 context: &Context,
57 options: Option<SimpleStreamOptions>,
58 ) -> AssistantMessageEventStream {
59 let opts = options.as_ref();
60 let base = build_base_options(model, context, opts, opts.and_then(|o| o.base.api_key.clone()));
61 if opts.and_then(|o| o.reasoning).is_none() {
62 return self.stream_with_options(
63 model,
64 context,
65 GoogleOptions {
66 base,
67 thinking: Some(GoogleThinkingConfig {
68 enabled: false,
69 budget_tokens: None,
70 level: None,
71 }),
72 ..Default::default()
73 },
74 );
75 }
76 let reasoning = clamp_thinking_level(model, opts.unwrap().reasoning.unwrap());
77 self.stream_with_options(
78 model,
79 context,
80 GoogleOptions {
81 base,
82 thinking: Some(GoogleThinkingConfig {
83 enabled: true,
84 budget_tokens: Some(get_google_budget(model, reasoning)),
85 level: None,
86 }),
87 ..Default::default()
88 },
89 )
90 }
91}
92
93impl GoogleGenerativeAIApi {
94 pub fn stream_with_options(
95 &self,
96 model: &Model,
97 context: &Context,
98 options: GoogleOptions,
99 ) -> AssistantMessageEventStream {
100 let stream = AssistantMessageEventStream::new();
101 let model = model.clone();
102 let context = context.clone();
103 let s = stream.clone();
104 tokio::spawn(async move {
105 let mut output = AssistantMessage::empty(&model);
106 if let Err(e) = run_google(&model, &context, &options, &s, &mut output).await {
107 let aborted = crate::api::common::is_abort_error(&e);
108 finish_stream_error(&s, &mut output, e, aborted);
109 }
110 });
111 stream
112 }
113}
114
115async fn run_google(
116 model: &Model,
117 context: &Context,
118 options: &GoogleOptions,
119 stream: &AssistantMessageEventStream,
120 output: &mut AssistantMessage,
121) -> Result<()> {
122 let api_key = options
123 .base
124 .api_key
125 .as_deref()
126 .ok_or_else(|| anyhow!("No API key for provider: {}", model.provider))?;
127 let mut params = build_params(model, context, options)?;
128 params = apply_on_payload(options.base.on_payload.as_ref(), params, model).await;
129 let headers = merge_model_headers(model, Some(&options.base));
130
131 let url = format!(
132 "{}/v1beta/models/{}:streamGenerateContent?alt=sse&key={}",
133 model.base_url.trim_end_matches('/'),
134 model.id,
135 api_key
136 );
137 let client = build_http_client_for_target(options.base.timeout_ms, Some(&url), options.base.env.as_ref())?;
138 let mut req = client.post(&url).json(¶ms);
139 for (k, v) in &headers {
140 req = req.header(k, v);
141 }
142 let response = crate::api::common::send_with_abort(&options.base.signal, req).await?;
143 invoke_on_response_from_reqwest(options.base.on_response.as_ref(), &response, model).await;
144 let response = crate::api::common::check_response_ok(response).await?;
145
146 stream.push(AssistantMessageEvent::Start {
147 partial: output.clone(),
148 });
149 let mut current_block: Option<usize> = None;
150 for_each_sse_json_event(response, &options.base.signal, |chunk| {
151 output.response_id = output
152 .response_id
153 .clone()
154 .or_else(|| chunk.get("responseId").and_then(|v| v.as_str()).map(|s| s.to_string()));
155 if let Some(candidate) = chunk.get("candidates").and_then(|c| c.get(0)) {
156 if let Some(parts) = candidate.pointer("/content/parts").and_then(|v| v.as_array()) {
157 for part in parts {
158 if let Some(text) = part.get("text").and_then(|v| v.as_str()) {
159 let is_thinking = is_thinking_part(part);
160 let idx = ensure_block(output, stream, &mut current_block, is_thinking);
161 match &mut output.content[idx] {
162 AssistantContentBlock::Thinking(t) => {
163 t.thinking.push_str(text);
164 t.thinking_signature = retain_thought_signature(
165 t.thinking_signature.as_deref(),
166 part.get("thoughtSignature").and_then(|v| v.as_str()),
167 );
168 stream.push(AssistantMessageEvent::ThinkingDelta {
169 content_index: idx,
170 delta: text.to_string(),
171 partial: output.clone(),
172 });
173 }
174 AssistantContentBlock::Text(t) => {
175 t.text.push_str(text);
176 t.text_signature = retain_thought_signature(
177 t.text_signature.as_deref(),
178 part.get("thoughtSignature").and_then(|v| v.as_str()),
179 );
180 stream.push(AssistantMessageEvent::TextDelta {
181 content_index: idx,
182 delta: text.to_string(),
183 partial: output.clone(),
184 });
185 }
186 _ => {}
187 }
188 }
189 if let Some(fc) = part.get("functionCall") {
190 end_current_block(output, stream, &mut current_block);
191 let name = fc.get("name").and_then(|v| v.as_str()).unwrap_or("");
192 let id = fc
193 .get("id")
194 .and_then(|v| v.as_str())
195 .map(|s| s.to_string())
196 .unwrap_or_else(|| {
197 format!(
198 "{}_{}_{}",
199 name,
200 chrono::Utc::now().timestamp_millis(),
201 TOOL_CALL_COUNTER.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
202 )
203 });
204 let tc = crate::types::ToolCall::new(id, name, fc.get("args").cloned().unwrap_or(json!({})));
205 let idx = output.content.len();
206 output.content.push(AssistantContentBlock::ToolCall(tc.clone()));
207 stream.push(AssistantMessageEvent::ToolcallStart {
208 content_index: idx,
209 partial: output.clone(),
210 });
211 stream.push(AssistantMessageEvent::ToolcallDelta {
212 content_index: idx,
213 delta: tc.arguments.to_string(),
214 partial: output.clone(),
215 });
216 stream.push(AssistantMessageEvent::ToolcallEnd {
217 content_index: idx,
218 tool_call: tc,
219 partial: output.clone(),
220 });
221 }
222 }
223 }
224 if let Some(reason) = candidate.get("finishReason").and_then(|v| v.as_str()) {
225 output.stop_reason = map_stop_reason_finish(reason);
226 if output.content.iter().any(|b| b.is_tool_call()) {
227 output.stop_reason = StopReason::ToolUse;
228 }
229 }
230 }
231 if let Some(meta) = chunk.get("usageMetadata") {
232 let prompt = meta.get("promptTokenCount").and_then(|v| v.as_u64()).unwrap_or(0);
233 let cached = meta
234 .get("cachedContentTokenCount")
235 .and_then(|v| v.as_u64())
236 .unwrap_or(0);
237 output.usage.input = prompt.saturating_sub(cached);
238 output.usage.output = meta.get("candidatesTokenCount").and_then(|v| v.as_u64()).unwrap_or(0)
239 + meta.get("thoughtsTokenCount").and_then(|v| v.as_u64()).unwrap_or(0);
240 output.usage.cache_read = cached;
241 output.usage.reasoning = meta.get("thoughtsTokenCount").and_then(|v| v.as_u64());
242 output.usage.total_tokens = meta.get("totalTokenCount").and_then(|v| v.as_u64()).unwrap_or(0);
243 calculate_cost(model, &mut output.usage);
244 }
245 Ok(())
246 })
247 .await?;
248 end_current_block(output, stream, &mut current_block);
249 if is_request_aborted(&options.base.signal) {
250 output.stop_reason = StopReason::Aborted;
251 }
252 stream.push(AssistantMessageEvent::Done {
253 reason: output.stop_reason,
254 message: output.clone(),
255 });
256 stream.end();
257 Ok(())
258}
259
260fn ensure_block(
261 output: &mut AssistantMessage,
262 stream: &AssistantMessageEventStream,
263 current: &mut Option<usize>,
264 thinking: bool,
265) -> usize {
266 let need_new = current.is_none()
267 || !matches!(
268 (thinking, current.and_then(|i| output.content.get(i))),
269 (true, Some(AssistantContentBlock::Thinking(_))) | (false, Some(AssistantContentBlock::Text(_)))
270 );
271 if need_new {
272 end_current_block(output, stream, current);
273 let idx = output.content.len();
274 if thinking {
275 output
276 .content
277 .push(AssistantContentBlock::Thinking(crate::types::ThinkingContent::new("")));
278 stream.push(AssistantMessageEvent::ThinkingStart {
279 content_index: idx,
280 partial: output.clone(),
281 });
282 } else {
283 output
284 .content
285 .push(AssistantContentBlock::Text(crate::types::TextContent::new("")));
286 stream.push(AssistantMessageEvent::TextStart {
287 content_index: idx,
288 partial: output.clone(),
289 });
290 }
291 *current = Some(idx);
292 }
293 current.unwrap()
294}
295
296fn end_current_block(output: &mut AssistantMessage, stream: &AssistantMessageEventStream, current: &mut Option<usize>) {
297 if let Some(idx) = current.take() {
298 match &output.content[idx] {
299 AssistantContentBlock::Text(t) => stream.push(AssistantMessageEvent::TextEnd {
300 content_index: idx,
301 content: t.text.clone(),
302 partial: output.clone(),
303 }),
304 AssistantContentBlock::Thinking(t) => stream.push(AssistantMessageEvent::ThinkingEnd {
305 content_index: idx,
306 content: t.thinking.clone(),
307 partial: output.clone(),
308 }),
309 _ => {}
310 }
311 }
312}
313
314fn build_params(model: &Model, context: &Context, options: &GoogleOptions) -> Result<Value> {
315 let contents = convert_messages(model, context);
316 let mut generation_config = json!({});
317 if let Some(temp) = options.base.temperature {
318 generation_config["temperature"] = json!(temp);
319 }
320 if let Some(max) = options.base.max_tokens {
321 generation_config["maxOutputTokens"] = json!(max);
322 }
323 let mut body = json!({ "contents": contents });
324 if let Some(sp) = &context.system_prompt {
325 body["systemInstruction"] = json!({ "parts": [{ "text": sanitize_surrogates(sp) }] });
326 }
327 if let Some(tools) = &context.tools {
328 if let Some(t) = convert_tools(tools, false) {
329 body["tools"] = json!(t);
330 }
331 if let Some(choice) = &options.tool_choice {
332 body["toolConfig"] = json!({ "functionCallingConfig": { "mode": map_tool_choice(choice) } });
333 }
334 }
335 if let Some(thinking) = &options.thinking
336 && thinking.enabled
337 && model.reasoning
338 {
339 let mut tc = json!({ "includeThoughts": true });
340 if let Some(level) = &thinking.level {
341 tc["thinkingLevel"] = json!(level);
342 } else if let Some(budget) = thinking.budget_tokens {
343 tc["thinkingBudget"] = json!(budget);
344 }
345 generation_config["thinkingConfig"] = tc;
346 }
347 if generation_config.as_object().map(|o| !o.is_empty()).unwrap_or(false) {
348 body["generationConfig"] = generation_config;
349 }
350 Ok(body)
351}
352
353pub fn get_google_budget(model: &Model, effort: crate::types::ThinkingLevel) -> i32 {
354 let level = match effort {
355 crate::types::ThinkingLevel::Minimal => "minimal",
356 crate::types::ThinkingLevel::Low => "low",
357 crate::types::ThinkingLevel::Medium => "medium",
358 crate::types::ThinkingLevel::High | crate::types::ThinkingLevel::Xhigh => "high",
359 };
360 if model.id.contains("2.5-pro") {
361 return match level {
362 "minimal" => 128,
363 "low" => 2048,
364 "medium" => 8192,
365 _ => 32768,
366 };
367 }
368 if model.id.contains("2.5-flash-lite") {
369 return match level {
370 "minimal" => 512,
371 "low" => 2048,
372 "medium" => 8192,
373 _ => 24576,
374 };
375 }
376 if model.id.contains("2.5-flash") {
377 return match level {
378 "minimal" => 128,
379 "low" => 2048,
380 "medium" => 8192,
381 _ => 24576,
382 };
383 }
384 -1
385}