oxicode_ai/providers/
google.rs1use 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#[derive(Clone)]
19pub struct GoogleProvider {
20 client: &'static Client,
21 api_key: Option<String>,
22}
23
24impl GoogleProvider {
25 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 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 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 let contents = convert_messages(context)?;
68
69 let tools_json = convert_tools(&context.tools, false);
71
72 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 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 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 if let Some(effort) = level.as_str() {
105 thinking_config["thinkingLevel"] = serde_json::json!(effort);
106 }
107 }
108
109 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 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 let model_name = model.id.clone();
144
145 let stream = response
146 .bytes_stream()
147 .scan(
148 Vec::new(), 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}