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 tool_config = super::google_shared::build_tool_config(options.tool_choice.as_ref());
74            let mut body = build_request_body(
75                &contents,
76                context.system_prompt.as_deref(),
77                tools_json.as_ref(),
78                options.temperature,
79                options.max_tokens,
80                tool_config.as_ref(),
81            );
82
83            // ── Google thinking config (via ProviderOptions) ────────────────
84            // When the model supports reasoning, apply thinkingConfig from
85            // provider_options.google. Mirrors opencode's Gemini thinking support.
86            if model.reasoning {
87                let google_opts = options
88                    .provider_options
89                    .as_ref()
90                    .and_then(|po| po.google.as_ref());
91
92                let mut thinking_config = serde_json::json!({});
93
94                // Include thoughts (always true for reasoning models)
95                thinking_config["includeThoughts"] = serde_json::json!(true);
96
97                if let Some(opts) = google_opts {
98                    if let Some(ref level) = opts.thinking_level {
99                        thinking_config["thinkingLevel"] = serde_json::json!(level);
100                    }
101                    if let Some(budget) = opts.thinking_budget {
102                        thinking_config["thinkingBudget"] = serde_json::json!(budget);
103                    }
104                } else if let Some(ref level) = options.thinking_level {
105                    // Fallback: derive from thinking_level
106                    if let Some(effort) = level.as_str() {
107                        thinking_config["thinkingLevel"] = serde_json::json!(effort);
108                    }
109                }
110
111                // Merge into generationConfig
112                if let Some(gc) = body.get_mut("generationConfig") {
113                    if let serde_json::Value::Object(map) = gc {
114                        map.insert("thinkingConfig".to_string(), thinking_config);
115                    }
116                } else {
117                    body["generationConfig"] = serde_json::json!({
118                        "thinkingConfig": thinking_config,
119                    });
120                }
121            }
122
123            // Make request with API key in header (not URL query param)
124            let response = self
125                .client
126                .post(&url)
127                .header("x-goog-api-key", api_key)
128                .header("Content-Type", "application/json")
129                .json(&body)
130                .send()
131                .await
132                .map_err(ProviderError::RequestFailed)?;
133
134            if !response.status().is_success() {
135                let status = response.status();
136                let body: String = response.text().await.unwrap_or_default();
137                return Err(ProviderError::HttpError(
138                    crate::error::HttpErrorDetail::new(status.as_u16(), body),
139                ));
140            }
141
142            // Create event stream — use split_complete_lines (like OpenAI provider)
143            // to handle UTF-8 boundaries safely.  Google SSE lines can be split
144            // across HTTP chunks at arbitrary byte boundaries.
145            let model_name = model.id.clone();
146
147            let stream = response
148                .bytes_stream()
149                .scan(
150                    Vec::new(), // pending_bytes
151                    move |pending_bytes, chunk: Result<bytes::Bytes, reqwest::Error>| {
152                        let events = match chunk {
153                            Ok(bytes) => {
154                                let mut combined =
155                                    Vec::with_capacity(pending_bytes.len() + bytes.len());
156                                combined.extend_from_slice(pending_bytes);
157                                combined.extend_from_slice(&bytes);
158                                let (text, trailing) = split_complete_lines(&combined);
159                                *pending_bytes = trailing;
160                                parse_google_events(
161                                    &text,
162                                    Api::GoogleGenerativeAi,
163                                    "google",
164                                    &model_name,
165                                )
166                            }
167                            Err(e) => vec![ProviderEvent::Error {
168                                reason: StopReason::Error,
169                                error: create_error_message(
170                                    Api::GoogleGenerativeAi,
171                                    "google",
172                                    &e.to_string(),
173                                ),
174                            }],
175                        };
176                        async move { Some(futures::stream::iter(events)) }
177                    },
178                )
179                .flatten();
180
181            Ok(Box::pin(stream) as Pin<Box<dyn Stream<Item = ProviderEvent> + Send>>)
182        })
183    }
184}
185
186#[cfg(test)]
187mod tests {
188    use super::*;
189    use crate::{Context, Message};
190
191    #[test]
192    fn test_build_google_contents_with_text() {
193        let mut ctx = Context::new();
194        ctx.add_message(Message::user("Hello, world!"));
195
196        let contents = convert_messages(&ctx).unwrap();
197        assert_eq!(contents.len(), 1);
198        assert_eq!(contents[0]["role"], "user");
199        assert_eq!(contents[0]["parts"][0]["text"], "Hello, world!");
200    }
201
202    #[test]
203    fn test_build_google_tools() {
204        let tools = vec![crate::Tool::new(
205            "get_weather",
206            "Get weather for a location",
207            serde_json::json!({
208                "type": "object",
209                "properties": {
210                    "location": {
211                        "type": "string",
212                        "description": "The city name"
213                    }
214                },
215                "required": ["location"]
216            }),
217        )];
218
219        let tools_json = convert_tools(&tools, false).unwrap();
220        let declarations = tools_json[0]["functionDeclarations"].as_array().unwrap();
221        assert_eq!(declarations.len(), 1);
222        assert_eq!(declarations[0]["name"], "get_weather");
223    }
224
225    #[test]
226    fn test_parse_google_events_basic_text() {
227        let sse_data = r#"data: {"candidates":[{"content":{"parts":[{"text":"Hello"}]}}]}"#;
228        let events = parse_google_events(
229            sse_data,
230            Api::GoogleGenerativeAi,
231            "google",
232            "gemini-1.5-pro",
233        );
234        assert!(!events.is_empty());
235    }
236
237    #[test]
238    fn test_create_error_message() {
239        let msg = create_error_message(Api::GoogleGenerativeAi, "google", "Something went wrong");
240        assert_eq!(msg.provider, "google");
241        assert_eq!(msg.api, Api::GoogleGenerativeAi);
242        assert_eq!(msg.stop_reason, StopReason::Error);
243    }
244}