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 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            // Use split_complete_lines for safe UTF-8 boundary handling
186            let stream = response
187                .bytes_stream()
188                .scan(
189                    Vec::new(), // pending_bytes
190                    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    // SAFETY: `serde_json::to_string` on a `Value` cannot fail — the Value type
222    // has no custom serializer that returns Err. Infallible by construction.
223    #[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)]
250// serde deserialization structs
251struct TokenResponse {
252    access_token: String,
253    _expires_in: usize,
254    _token_type: String,
255}
256
257#[derive(Debug, serde::Deserialize)]
258// serde deserialization structs
259struct 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        // SAFETY: test-only, single-threaded test binary.
320        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}