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 body = build_request_body(
162 &contents,
163 context.system_prompt.as_deref(),
164 tools_json.as_ref(),
165 options.temperature,
166 options.max_tokens,
167 );
168 let response = self
169 .client
170 .post(&url)
171 .header("Authorization", format!("Bearer {}", access_token))
172 .header("Content-Type", "application/json")
173 .json(&body)
174 .send()
175 .await
176 .map_err(ProviderError::RequestFailed)?;
177 if !response.status().is_success() {
178 let status = response.status();
179 let body: String = response.text().await.unwrap_or_default();
180 return Err(ProviderError::HttpError(
181 crate::error::HttpErrorDetail::new(status.as_u16(), body),
182 ));
183 }
184 let model_name = model.id.clone();
185 let stream = response
187 .bytes_stream()
188 .scan(
189 Vec::new(), move |pending_bytes, chunk: Result<bytes::Bytes, reqwest::Error>| {
191 let events = match chunk {
192 Ok(bytes) => {
193 let mut combined =
194 Vec::with_capacity(pending_bytes.len() + bytes.len());
195 combined.extend_from_slice(pending_bytes);
196 combined.extend_from_slice(&bytes);
197 let (text, trailing) = split_complete_lines(&combined);
198 *pending_bytes = trailing;
199 parse_google_events(&text, Api::GoogleVertex, "vertex", &model_name)
200 }
201 Err(e) => vec![ProviderEvent::Error {
202 reason: StopReason::Error,
203 error: create_error_message(
204 Api::GoogleVertex,
205 "vertex",
206 &e.to_string(),
207 ),
208 }],
209 };
210 async move { Some(futures::stream::iter(events)) }
211 },
212 )
213 .flatten();
214 Ok(Box::pin(stream) as Pin<Box<dyn Stream<Item = ProviderEvent> + Send>>)
215 })
216 }
217}
218
219fn base64_url_encode(value: &serde_json::Value) -> String {
220 use base64::Engine as _;
221 #[allow(clippy::expect_used)]
224 let json = serde_json::to_string(value).expect("serializing json value");
225 base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(json.as_bytes())
226}
227
228fn sign_rs256(
229 header_b64: &str,
230 claims_b64: &str,
231 private_key_pem: &str,
232) -> Result<String, ProviderError> {
233 use base64::Engine as _;
234 use pkcs8::DecodePrivateKey;
235 use rsa::RsaPrivateKey;
236 use rsa::pkcs1v15::SigningKey;
237 use sha2::Sha256;
238 use signature::{SignatureEncoding, Signer};
239 let message = format!("{}.{}", header_b64, claims_b64);
240 let key =
241 RsaPrivateKey::from_pkcs8_pem(private_key_pem).map_err(|_| ProviderError::InvalidApiKey)?;
242 let signing_key = SigningKey::<Sha256>::new_unprefixed(key);
243 let signature = signing_key.sign(message.as_bytes());
244 let sig_bytes = signature.to_bytes();
245 let sig_b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(&sig_bytes);
246 Ok(format!("{}.{}", message, sig_b64))
247}
248
249#[derive(Debug, serde::Deserialize)]
250struct TokenResponse {
252 access_token: String,
253 _expires_in: usize,
254 _token_type: String,
255}
256
257#[derive(Debug, serde::Deserialize)]
258struct ServiceAccountCreds {
260 #[serde(rename = "type")]
261 _type: String,
262 _project_id: String,
263 _private_key_id: String,
264 private_key: String,
265 client_email: String,
266 _client_id: String,
267 _auth_uri: String,
268 _token_uri: String,
269}
270
271#[cfg(test)]
272mod tests {
273 use super::*;
274 use crate::{Context, Message};
275
276 #[test]
277 fn test_build_vertex_contents_with_text() {
278 let mut ctx = Context::new();
279 ctx.add_message(Message::user("Hello, world!"));
280 let contents = convert_messages(&ctx).unwrap();
281 assert_eq!(contents.len(), 1);
282 assert_eq!(contents[0]["role"], "user");
283 assert_eq!(contents[0]["parts"][0]["text"], "Hello, world!");
284 }
285
286 #[test]
287 fn test_build_vertex_tools() {
288 let tools = vec![crate::Tool::new(
289 "get_weather",
290 "Get weather for a location",
291 serde_json::json!({
292 "type": "object",
293 "properties": {
294 "location": { "type": "string", "description": "The city name" }
295 },
296 "required": ["location"]
297 }),
298 )];
299 let tools_json = convert_tools(&tools, false).unwrap();
300 let declarations = tools_json[0]["functionDeclarations"].as_array().unwrap();
301 assert_eq!(declarations.len(), 1);
302 assert_eq!(declarations[0]["name"], "get_weather");
303 }
304
305 #[test]
306 fn test_parse_vertex_events_basic_text() {
307 let sse_data = r#"data: {"candidates":[{"content":{"parts":[{"text":"Hello"}]}}]}"#;
308 let events = parse_google_events(sse_data, Api::GoogleVertex, "vertex", "gemini-1.5-pro");
309 assert!(!events.is_empty());
310 if let ProviderEvent::TextDelta { delta, .. } = &events[0] {
311 assert_eq!(delta, "Hello");
312 } else {
313 panic!("Expected TextDelta event");
314 }
315 }
316
317 #[test]
318 fn test_get_region_default() {
319 unsafe { std::env::remove_var("GOOGLE_CLOUD_REGION") };
321 assert_eq!(VertexProvider::get_region(), "us-central1");
322 }
323
324 #[test]
325 fn test_create_error_message() {
326 let msg = create_error_message(Api::GoogleVertex, "vertex", "Something went wrong");
327 assert_eq!(msg.provider, "vertex");
328 assert_eq!(msg.api, Api::GoogleVertex);
329 assert_eq!(msg.stop_reason, StopReason::Error);
330 }
331}