Skip to main content

elph_ai/api/
google_generative_ai.rs

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(&params);
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}