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 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 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 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 if let Some(effort) = level.as_str() {
107 thinking_config["thinkingLevel"] = serde_json::json!(effort);
108 }
109 }
110
111 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 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 let model_name = model.id.clone();
146
147 let stream = response
148 .bytes_stream()
149 .scan(
150 Vec::new(), 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}