Skip to main content

oxicode_ai/providers/
vertex.rs

1//! Google Vertex AI provider
2//!
3//! This provider uses Google Cloud authentication (service account or gcloud CLI)
4//! to access Vertex AI models via the Gemini API.
5
6use 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/// Google Vertex AI provider
21///
22/// Uses Bearer token authentication via:
23/// - Service account JSON file (GOOGLE_APPLICATION_CREDENTIALS)
24/// - gcloud CLI access token (from `gcloud auth print-access-token`)
25#[derive(Clone)]
26pub struct VertexProvider {
27    client: &'static Client,
28}
29
30impl VertexProvider {
31    /// Create a new Vertex provider with default settings.
32    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            // Use split_complete_lines for safe UTF-8 boundary handling
188            let stream = response
189                .bytes_stream()
190                .scan(
191                    Vec::new(), // pending_bytes
192                    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    // SAFETY: `serde_json::to_string` on a `Value` cannot fail — the Value type
224    // has no custom serializer that returns Err. Infallible by construction.
225    #[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)]
252// serde deserialization structs
253struct TokenResponse {
254    access_token: String,
255    _expires_in: usize,
256    _token_type: String,
257}
258
259#[derive(Debug, serde::Deserialize)]
260// serde deserialization structs
261struct 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        // SAFETY: test-only, single-threaded test binary.
322        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}