1use futures::Stream;
7use futures::stream::StreamExt;
8use reqwest::Client;
9use std::future::Future;
10use std::pin::Pin;
11
12use super::google_shared::{
13 build_request_body, convert_messages, convert_tools, create_error_message, parse_google_events,
14};
15use super::shared_client;
16use super::sse::split_complete_lines;
17use super::{Provider, ProviderError, ProviderEvent, StreamOptions, StreamResult};
18use crate::{Api, Context, Model, StopReason};
19
20#[derive(Clone)]
26pub struct VertexProvider {
27 client: &'static Client,
28}
29
30impl VertexProvider {
31 pub fn new() -> Self {
33 Self {
34 client: shared_client(),
35 }
36 }
37
38 async fn get_access_token(&self) -> Result<String, ProviderError> {
39 if let Ok(token) = std::env::var("GOOGLE_ACCESS_TOKEN")
40 && !token.is_empty()
41 {
42 return Ok(token);
43 }
44 if let Ok(token) = Self::get_gcloud_token().await {
45 return Ok(token);
46 }
47 if let Ok(creds) = std::env::var("GOOGLE_APPLICATION_CREDENTIALS")
48 && !creds.is_empty()
49 {
50 return Self::get_token_from_service_account(&creds).await;
51 }
52 Err(ProviderError::MissingApiKey)
53 }
54
55 async fn get_gcloud_token() -> Result<String, ProviderError> {
56 use std::io;
57 use tokio::process::Command;
58 let output = Command::new("gcloud")
59 .args(["auth", "print-access-token"])
60 .output()
61 .await
62 .map_err(ProviderError::IoError)?;
63 if output.status.success() {
64 let token = String::from_utf8_lossy(&output.stdout).trim().to_string();
65 if !token.is_empty() {
66 return Ok(token);
67 }
68 }
69 Err(ProviderError::IoError(io::Error::new(
70 io::ErrorKind::NotFound,
71 "gcloud token not available",
72 )))
73 }
74
75 async fn get_token_from_service_account(
76 credentials_path: &str,
77 ) -> Result<String, ProviderError> {
78 use std::fs;
79 use tokio::time::{Duration, sleep};
80 let creds_json = fs::read_to_string(credentials_path).map_err(ProviderError::IoError)?;
81 let creds: ServiceAccountCreds =
82 serde_json::from_str(&creds_json).map_err(|_| ProviderError::InvalidApiKey)?;
83 let now = std::time::SystemTime::now()
84 .duration_since(std::time::UNIX_EPOCH)
85 .map_err(|_| ProviderError::InvalidResponse("system time before Unix epoch".into()))?
86 .as_secs();
87 let header = base64_url_encode(&serde_json::json!({"alg": "RS256", "typ": "JWT"}));
88 let claims = serde_json::json!({
89 "iss": creds.client_email,
90 "sub": creds.client_email,
91 "aud": "https://oauth2.googleapis.com/token",
92 "iat": now,
93 "exp": now + 3600,
94 "scope": "https://www.googleapis.com/auth/cloud-platform"
95 });
96 let claims_b64 = base64_url_encode(&claims);
97 let signature = sign_rs256(&header, &claims_b64, &creds.private_key)?;
98 let jwt = signature;
99 let client = shared_client();
100 let response = client
101 .post("https://oauth2.googleapis.com/token")
102 .form(&[
103 ("grant_type", "urn:ietf:params:oauth:grant-type:jwt-bearer"),
104 ("assertion", &jwt),
105 ])
106 .send()
107 .await
108 .map_err(ProviderError::RequestFailed)?;
109 if !response.status().is_success() {
110 return Err(ProviderError::HttpError(
111 crate::error::HttpErrorDetail::new(
112 response.status().as_u16(),
113 response.text().await.unwrap_or_default(),
114 ),
115 ));
116 }
117 let token_response: TokenResponse = response
118 .json()
119 .await
120 .map_err(ProviderError::RequestFailed)?;
121 sleep(Duration::from_secs(60 * 55)).await;
122 Ok(token_response.access_token)
123 }
124
125 fn get_project_id() -> Result<String, ProviderError> {
126 std::env::var("GOOGLE_CLOUD_PROJECT")
127 .or_else(|_| std::env::var("GOOGLE_PROJECT"))
128 .map_err(|_| ProviderError::MissingApiKey)
129 }
130
131 fn get_region() -> String {
132 std::env::var("GOOGLE_CLOUD_REGION").unwrap_or_else(|_| "us-central1".to_string())
133 }
134}
135
136impl Default for VertexProvider {
137 fn default() -> Self {
138 Self::new()
139 }
140}
141
142impl Provider for VertexProvider {
143 fn stream<'a>(
144 &'a self,
145 model: &'a Model,
146 context: &'a Context,
147 options: Option<StreamOptions>,
148 ) -> Pin<Box<dyn Future<Output = StreamResult> + Send + 'a>> {
149 Box::pin(async move {
150 let options = options.unwrap_or_default();
151 let access_token = self.get_access_token().await?;
152 let project_id = Self::get_project_id()?;
153 let region = Self::get_region();
154 let model_id = &model.id;
155 let url = format!(
156 "https://{}-aiplatform.googleapis.com/v1/projects/{}/locations/{}/publishers/google/models/{}:streamGenerateContent",
157 region, project_id, region, model_id
158 );
159 let contents = convert_messages(context)?;
160 let tools_json = convert_tools(&context.tools, false);
161 let tool_config = super::google_shared::build_tool_config(options.tool_choice.as_ref());
162 let body = build_request_body(
163 &contents,
164 context.system_prompt.as_deref(),
165 tools_json.as_ref(),
166 options.temperature,
167 options.max_tokens,
168 tool_config.as_ref(),
169 );
170 let response = self
171 .client
172 .post(&url)
173 .header("Authorization", format!("Bearer {}", access_token))
174 .header("Content-Type", "application/json")
175 .json(&body)
176 .send()
177 .await
178 .map_err(ProviderError::RequestFailed)?;
179 if !response.status().is_success() {
180 let status = response.status();
181 let body: String = response.text().await.unwrap_or_default();
182 return Err(ProviderError::HttpError(
183 crate::error::HttpErrorDetail::new(status.as_u16(), body),
184 ));
185 }
186 let model_name = model.id.clone();
187 let stream = response
189 .bytes_stream()
190 .scan(
191 Vec::new(), move |pending_bytes, chunk: Result<bytes::Bytes, reqwest::Error>| {
193 let events = match chunk {
194 Ok(bytes) => {
195 let mut combined =
196 Vec::with_capacity(pending_bytes.len() + bytes.len());
197 combined.extend_from_slice(pending_bytes);
198 combined.extend_from_slice(&bytes);
199 let (text, trailing) = split_complete_lines(&combined);
200 *pending_bytes = trailing;
201 parse_google_events(&text, Api::GoogleVertex, "vertex", &model_name)
202 }
203 Err(e) => vec![ProviderEvent::Error {
204 reason: StopReason::Error,
205 error: create_error_message(
206 Api::GoogleVertex,
207 "vertex",
208 &e.to_string(),
209 ),
210 }],
211 };
212 async move { Some(futures::stream::iter(events)) }
213 },
214 )
215 .flatten();
216 Ok(Box::pin(stream) as Pin<Box<dyn Stream<Item = ProviderEvent> + Send>>)
217 })
218 }
219}
220
221fn base64_url_encode(value: &serde_json::Value) -> String {
222 use base64::Engine as _;
223 #[allow(clippy::expect_used)]
226 let json = serde_json::to_string(value).expect("serializing json value");
227 base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(json.as_bytes())
228}
229
230fn sign_rs256(
231 header_b64: &str,
232 claims_b64: &str,
233 private_key_pem: &str,
234) -> Result<String, ProviderError> {
235 use base64::Engine as _;
236 use pkcs8::DecodePrivateKey;
237 use rsa::RsaPrivateKey;
238 use rsa::pkcs1v15::SigningKey;
239 use sha2::Sha256;
240 use signature::{SignatureEncoding, Signer};
241 let message = format!("{}.{}", header_b64, claims_b64);
242 let key =
243 RsaPrivateKey::from_pkcs8_pem(private_key_pem).map_err(|_| ProviderError::InvalidApiKey)?;
244 let signing_key = SigningKey::<Sha256>::new_unprefixed(key);
245 let signature = signing_key.sign(message.as_bytes());
246 let sig_bytes = signature.to_bytes();
247 let sig_b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(&sig_bytes);
248 Ok(format!("{}.{}", message, sig_b64))
249}
250
251#[derive(Debug, serde::Deserialize)]
252struct TokenResponse {
254 access_token: String,
255 _expires_in: usize,
256 _token_type: String,
257}
258
259#[derive(Debug, serde::Deserialize)]
260struct ServiceAccountCreds {
262 #[serde(rename = "type")]
263 _type: String,
264 _project_id: String,
265 _private_key_id: String,
266 private_key: String,
267 client_email: String,
268 _client_id: String,
269 _auth_uri: String,
270 _token_uri: String,
271}
272
273#[cfg(test)]
274mod tests {
275 use super::*;
276 use crate::{Context, Message};
277
278 #[test]
279 fn test_build_vertex_contents_with_text() {
280 let mut ctx = Context::new();
281 ctx.add_message(Message::user("Hello, world!"));
282 let contents = convert_messages(&ctx).unwrap();
283 assert_eq!(contents.len(), 1);
284 assert_eq!(contents[0]["role"], "user");
285 assert_eq!(contents[0]["parts"][0]["text"], "Hello, world!");
286 }
287
288 #[test]
289 fn test_build_vertex_tools() {
290 let tools = vec![crate::Tool::new(
291 "get_weather",
292 "Get weather for a location",
293 serde_json::json!({
294 "type": "object",
295 "properties": {
296 "location": { "type": "string", "description": "The city name" }
297 },
298 "required": ["location"]
299 }),
300 )];
301 let tools_json = convert_tools(&tools, false).unwrap();
302 let declarations = tools_json[0]["functionDeclarations"].as_array().unwrap();
303 assert_eq!(declarations.len(), 1);
304 assert_eq!(declarations[0]["name"], "get_weather");
305 }
306
307 #[test]
308 fn test_parse_vertex_events_basic_text() {
309 let sse_data = r#"data: {"candidates":[{"content":{"parts":[{"text":"Hello"}]}}]}"#;
310 let events = parse_google_events(sse_data, Api::GoogleVertex, "vertex", "gemini-1.5-pro");
311 assert!(!events.is_empty());
312 if let ProviderEvent::TextDelta { delta, .. } = &events[0] {
313 assert_eq!(delta, "Hello");
314 } else {
315 panic!("Expected TextDelta event");
316 }
317 }
318
319 #[test]
320 fn test_get_region_default() {
321 unsafe { std::env::remove_var("GOOGLE_CLOUD_REGION") };
323 assert_eq!(VertexProvider::get_region(), "us-central1");
324 }
325
326 #[test]
327 fn test_create_error_message() {
328 let msg = create_error_message(Api::GoogleVertex, "vertex", "Something went wrong");
329 assert_eq!(msg.provider, "vertex");
330 assert_eq!(msg.api, Api::GoogleVertex);
331 assert_eq!(msg.stop_reason, StopReason::Error);
332 }
333}