Skip to main content

oxicode_ai/providers/
google.rs

1//! Google Generative AI provider (Gemini API)
2
3use futures::Stream;
4use futures::stream::StreamExt;
5use reqwest::Client;
6use std::future::Future;
7use std::pin::Pin;
8
9use super::google_shared::{
10    build_request_body, convert_messages, convert_tools, create_error_message, parse_google_events,
11};
12use super::shared_client;
13use super::sse::split_complete_lines;
14use super::{Provider, ProviderError, ProviderEvent, StreamOptions, StreamResult};
15use crate::{Api, Context, Model, StopReason};
16
17/// Google Generative AI provider
18#[derive(Clone)]
19pub struct GoogleProvider {
20    client: &'static Client,
21    api_key: Option<String>,
22}
23
24impl GoogleProvider {
25    /// Create a new Google provider without an API key.
26    ///
27    /// API keys are resolved at request time via auth.json or StreamOptions.
28    pub fn new() -> Self {
29        Self {
30            client: shared_client(),
31            api_key: None,
32        }
33    }
34}
35
36impl Default for GoogleProvider {
37    fn default() -> Self {
38        Self::new()
39    }
40}
41
42impl Provider for GoogleProvider {
43    fn stream<'a>(
44        &'a self,
45        model: &'a Model,
46        context: &'a Context,
47        options: Option<StreamOptions>,
48    ) -> Pin<Box<dyn Future<Output = StreamResult> + Send + 'a>> {
49        Box::pin(async move {
50            let options = options.unwrap_or_default();
51
52            // Get API key
53            let api_key = options
54                .api_key
55                .as_ref()
56                .or(self.api_key.as_ref())
57                .ok_or_else(|| ProviderError::MissingApiKey)?;
58
59            // Build the request URL (without key - uses header instead for security)
60            let model_id = &model.id;
61            let url = format!(
62                "https://generativelanguage.googleapis.com/v1beta/models/{}:streamGenerateContent?alt=sse",
63                model_id
64            );
65
66            // Build contents using shared conversion
67            let contents = convert_messages(context)?;
68
69            // Build tools using shared conversion
70            let tools_json = convert_tools(&context.tools, false);
71
72            // Build request body using shared helper
73            let mut body = build_request_body(
74                &contents,
75                context.system_prompt.as_deref(),
76                tools_json.as_ref(),
77                options.temperature,
78                options.max_tokens,
79            );
80
81            // ── Google thinking config (via ProviderOptions) ────────────────
82            // When the model supports reasoning, apply thinkingConfig from
83            // provider_options.google. Mirrors opencode's Gemini thinking support.
84            if model.reasoning {
85                let google_opts = options
86                    .provider_options
87                    .as_ref()
88                    .and_then(|po| po.google.as_ref());
89
90                let mut thinking_config = serde_json::json!({});
91
92                // Include thoughts (always true for reasoning models)
93                thinking_config["includeThoughts"] = serde_json::json!(true);
94
95                if let Some(opts) = google_opts {
96                    if let Some(ref level) = opts.thinking_level {
97                        thinking_config["thinkingLevel"] = serde_json::json!(level);
98                    }
99                    if let Some(budget) = opts.thinking_budget {
100                        thinking_config["thinkingBudget"] = serde_json::json!(budget);
101                    }
102                } else if let Some(ref level) = options.thinking_level {
103                    // Fallback: derive from thinking_level
104                    if let Some(effort) = level.as_str() {
105                        thinking_config["thinkingLevel"] = serde_json::json!(effort);
106                    }
107                }
108
109                // Merge into generationConfig
110                if let Some(gc) = body.get_mut("generationConfig") {
111                    if let serde_json::Value::Object(map) = gc {
112                        map.insert("thinkingConfig".to_string(), thinking_config);
113                    }
114                } else {
115                    body["generationConfig"] = serde_json::json!({
116                        "thinkingConfig": thinking_config,
117                    });
118                }
119            }
120
121            // Make request with API key in header (not URL query param)
122            let response = self
123                .client
124                .post(&url)
125                .header("x-goog-api-key", api_key)
126                .header("Content-Type", "application/json")
127                .json(&body)
128                .send()
129                .await
130                .map_err(ProviderError::RequestFailed)?;
131
132            if !response.status().is_success() {
133                let status = response.status();
134                let body: String = response.text().await.unwrap_or_default();
135                return Err(ProviderError::HttpError(
136                    crate::error::HttpErrorDetail::new(status.as_u16(), body),
137                ));
138            }
139
140            // Create event stream — use split_complete_lines (like OpenAI provider)
141            // to handle UTF-8 boundaries safely.  Google SSE lines can be split
142            // across HTTP chunks at arbitrary byte boundaries.
143            let model_name = model.id.clone();
144
145            let stream = response
146                .bytes_stream()
147                .scan(
148                    Vec::new(), // pending_bytes
149                    move |pending_bytes, chunk: Result<bytes::Bytes, reqwest::Error>| {
150                        let events = match chunk {
151                            Ok(bytes) => {
152                                let mut combined =
153                                    Vec::with_capacity(pending_bytes.len() + bytes.len());
154                                combined.extend_from_slice(pending_bytes);
155                                combined.extend_from_slice(&bytes);
156                                let (text, trailing) = split_complete_lines(&combined);
157                                *pending_bytes = trailing;
158                                parse_google_events(
159                                    &text,
160                                    Api::GoogleGenerativeAi,
161                                    "google",
162                                    &model_name,
163                                )
164                            }
165                            Err(e) => vec![ProviderEvent::Error {
166                                reason: StopReason::Error,
167                                error: create_error_message(
168                                    Api::GoogleGenerativeAi,
169                                    "google",
170                                    &e.to_string(),
171                                ),
172                            }],
173                        };
174                        async move { Some(futures::stream::iter(events)) }
175                    },
176                )
177                .flatten();
178
179            Ok(Box::pin(stream) as Pin<Box<dyn Stream<Item = ProviderEvent> + Send>>)
180        })
181    }
182}
183
184#[cfg(test)]
185mod tests {
186    use super::*;
187    use crate::{Context, Message};
188
189    #[test]
190    fn test_build_google_contents_with_text() {
191        let mut ctx = Context::new();
192        ctx.add_message(Message::user("Hello, world!"));
193
194        let contents = convert_messages(&ctx).unwrap();
195        assert_eq!(contents.len(), 1);
196        assert_eq!(contents[0]["role"], "user");
197        assert_eq!(contents[0]["parts"][0]["text"], "Hello, world!");
198    }
199
200    #[test]
201    fn test_build_google_tools() {
202        let tools = vec![crate::Tool::new(
203            "get_weather",
204            "Get weather for a location",
205            serde_json::json!({
206                "type": "object",
207                "properties": {
208                    "location": {
209                        "type": "string",
210                        "description": "The city name"
211                    }
212                },
213                "required": ["location"]
214            }),
215        )];
216
217        let tools_json = convert_tools(&tools, false).unwrap();
218        let declarations = tools_json[0]["functionDeclarations"].as_array().unwrap();
219        assert_eq!(declarations.len(), 1);
220        assert_eq!(declarations[0]["name"], "get_weather");
221    }
222
223    #[test]
224    fn test_parse_google_events_basic_text() {
225        let sse_data = r#"data: {"candidates":[{"content":{"parts":[{"text":"Hello"}]}}]}"#;
226        let events = parse_google_events(
227            sse_data,
228            Api::GoogleGenerativeAi,
229            "google",
230            "gemini-1.5-pro",
231        );
232        assert!(!events.is_empty());
233    }
234
235    #[test]
236    fn test_create_error_message() {
237        let msg = create_error_message(Api::GoogleGenerativeAi, "google", "Something went wrong");
238        assert_eq!(msg.provider, "google");
239        assert_eq!(msg.api, Api::GoogleGenerativeAi);
240        assert_eq!(msg.stop_reason, StopReason::Error);
241    }
242}