gemini_genai_rs/transport/auth/
mod.rs1pub mod google_ai;
10pub(crate) mod url_builders;
11pub mod vertex;
12
13pub use google_ai::*;
14pub use vertex::*;
15
16use async_trait::async_trait;
17
18use crate::protocol::types::ModelId;
19use crate::session::AuthError;
20
21#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
25pub enum ServiceEndpoint {
26 LiveWs,
28 GenerateContent,
30 StreamGenerateContent,
32 EmbedContent,
34 CountTokens,
36 ComputeTokens,
38 ListModels,
40 GetModel,
42 Files,
44 CachedContents,
46 TuningJobs,
48 BatchJobs,
50}
51
52impl ServiceEndpoint {
53 pub fn model_method(&self) -> Option<&'static str> {
56 match self {
57 Self::GenerateContent => Some("generateContent"),
58 Self::StreamGenerateContent => Some("streamGenerateContent"),
59 Self::EmbedContent => Some("embedContent"),
60 Self::CountTokens => Some("countTokens"),
61 Self::ComputeTokens => Some("computeTokens"),
62 _ => None,
63 }
64 }
65
66 pub fn requires_model(&self) -> bool {
68 matches!(
69 self,
70 Self::GenerateContent
71 | Self::StreamGenerateContent
72 | Self::EmbedContent
73 | Self::CountTokens
74 | Self::ComputeTokens
75 | Self::GetModel
76 )
77 }
78}
79
80#[async_trait]
82pub trait AuthProvider: Send + Sync + 'static {
83 fn ws_url(&self, model: &ModelId) -> String;
85
86 async fn auth_headers(&self) -> Result<Vec<(String, String)>, AuthError>;
88
89 fn query_params(&self) -> Vec<(String, String)> {
91 vec![]
92 }
93
94 async fn refresh(&self) -> Result<(), AuthError> {
96 Ok(())
97 }
98}
99
100pub trait RestAuth: AuthProvider {
107 fn rest_url(&self, endpoint: ServiceEndpoint, model: Option<&ModelId>) -> String;
109}
110
111#[cfg(test)]
116mod tests {
117 use super::*;
118 use crate::protocol::types::ModelId;
119
120 #[test]
121 fn google_ai_auth_url() {
122 let auth = GoogleAIAuth::new("test-key-123");
123 let url = auth.ws_url(&ModelId::FLASH_LATEST);
124 assert!(url.contains("generativelanguage.googleapis.com"));
125 assert!(url.contains("v1beta"));
126 assert!(url.contains("key=test-key-123"));
127 }
128
129 #[test]
130 fn google_ai_auth_query_params() {
131 let auth = GoogleAIAuth::new("my-api-key");
132 let params = auth.query_params();
133 assert_eq!(params.len(), 1);
134 assert_eq!(params[0].0, "key");
135 assert_eq!(params[0].1, "my-api-key");
136 }
137
138 #[tokio::test]
139 async fn google_ai_rest_key_travels_in_a_header_not_the_url() {
140 let auth = GoogleAIAuth::new("test-key");
141 let headers = auth.auth_headers().await.unwrap();
142 assert_eq!(
143 headers,
144 vec![("x-goog-api-key".to_string(), "test-key".to_string())]
145 );
146 let url = auth.rest_url(
147 ServiceEndpoint::GenerateContent,
148 Some(&ModelId::FLASH_LATEST),
149 );
150 assert!(!url.contains("test-key"), "{url}");
151 }
152
153 #[test]
154 fn google_ai_token_auth_url() {
155 let auth = GoogleAITokenAuth::new("oauth2-token-abc");
156 let url = auth.ws_url(&ModelId::FLASH_LATEST);
157 assert!(url.contains("generativelanguage.googleapis.com"));
158 assert!(url.contains("access_token=oauth2-token-abc"));
159 assert!(url.contains("v1alpha"));
160 }
161
162 #[test]
163 fn vertex_ai_auth_url_regional() {
164 let auth = VertexAIAuth::new("my-project", "us-central1", "token");
165 let url = auth.ws_url(&ModelId::FLASH_LATEST);
166 assert!(url.contains("us-central1-aiplatform.googleapis.com"));
167 assert!(url.contains("v1beta1"));
168 assert!(url.contains("x-goog-project-id=my-project"));
169 }
170
171 #[test]
172 fn vertex_ai_auth_url_global() {
173 let auth = VertexAIAuth::new("my-project", "global", "token");
174 let url = auth.ws_url(&ModelId::FLASH_LATEST);
175 assert!(url.starts_with("wss://aiplatform.googleapis.com/"));
177 assert!(!url.contains("global-aiplatform"));
178 }
179
180 #[tokio::test]
181 async fn vertex_ai_auth_headers() {
182 let auth = VertexAIAuth::new("proj", "us-central1", "my-bearer-token");
183 let headers = auth.auth_headers().await.unwrap();
184 assert_eq!(headers.len(), 1);
185 assert_eq!(headers[0].0, "Authorization");
186 assert_eq!(headers[0].1, "Bearer my-bearer-token");
187 }
188
189 #[test]
190 fn vertex_ai_auth_url_contains_model() {
191 let auth = VertexAIAuth::new("proj", "us-central1", "tok");
192 let url = auth.ws_url(&ModelId::from_static("models/gemini-2.0-flash-live-001"));
193 assert!(url.contains("model=gemini-2.0-flash-live-001"));
194 }
195
196 #[test]
197 fn auth_provider_is_object_safe() {
198 fn _assert(_: &dyn AuthProvider) {}
199 }
200
201 #[tokio::test]
202 async fn default_refresh_is_noop() {
203 let auth = GoogleAIAuth::new("key");
204 auth.refresh().await.unwrap();
206 }
207
208 #[tokio::test]
209 async fn default_query_params_empty_for_vertex() {
210 let auth = VertexAIAuth::new("proj", "loc", "tok");
211 let params = auth.query_params();
212 assert!(params.is_empty());
213 }
214
215 #[test]
220 fn google_ai_rest_url_generate_content() {
221 let auth = GoogleAIAuth::new("test-key");
222 let model = ModelId::from_static("models/gemini-2.0-flash-live-001");
223 let url = auth.rest_url(ServiceEndpoint::GenerateContent, Some(&model));
224 assert!(url.starts_with("https://generativelanguage.googleapis.com/v1beta/"));
225 assert!(url.contains(":generateContent"));
226 assert!(
227 !url.contains("test-key"),
228 "the key rides in a header: {url}"
229 );
230 }
231
232 #[test]
233 fn google_ai_rest_url_list_models() {
234 let auth = GoogleAIAuth::new("key123");
235 let url = auth.rest_url(ServiceEndpoint::ListModels, None);
236 assert!(url.ends_with("/models"), "{url}");
237 }
238
239 #[test]
240 fn google_ai_rest_url_files() {
241 let auth = GoogleAIAuth::new("key");
242 let url = auth.rest_url(ServiceEndpoint::Files, None);
243 assert!(url.ends_with("/files"), "{url}");
244 }
245
246 #[test]
247 fn google_ai_token_rest_url_no_key_in_url() {
248 let auth = GoogleAITokenAuth::new("oauth-token");
249 let url = auth.rest_url(ServiceEndpoint::CountTokens, Some(&ModelId::FLASH_LATEST));
250 assert!(url.contains(":countTokens"));
251 assert!(!url.contains("key="));
252 assert!(!url.contains("access_token="));
253 }
254
255 #[test]
256 fn vertex_rest_url_generate_content() {
257 let auth = VertexAIAuth::new("my-project", "us-central1", "token");
258 let model = ModelId::from_static("models/gemini-2.0-flash-live-001");
259 let url = auth.rest_url(ServiceEndpoint::GenerateContent, Some(&model));
260 assert!(url.starts_with("https://us-central1-aiplatform.googleapis.com/v1beta1/"));
261 assert!(url.contains("projects/my-project/locations/us-central1"));
262 assert!(url.contains(":generateContent"));
263 }
264
265 #[test]
266 fn vertex_rest_url_list_models() {
267 let auth = VertexAIAuth::new("proj", "us-east1", "tok");
268 let url = auth.rest_url(ServiceEndpoint::ListModels, None);
269 assert!(url.contains("publishers/google/models"));
270 }
271
272 #[test]
273 fn vertex_rest_url_global() {
274 let auth = VertexAIAuth::new("proj", "global", "tok");
275 let model = ModelId::FLASH_LATEST;
276 let url = auth.rest_url(ServiceEndpoint::EmbedContent, Some(&model));
277 assert!(url.starts_with("https://aiplatform.googleapis.com/"));
278 assert!(!url.contains("global-aiplatform"));
279 assert!(url.contains(":embedContent"));
280 }
281
282 #[test]
283 fn service_endpoint_model_method() {
284 assert_eq!(
285 ServiceEndpoint::GenerateContent.model_method(),
286 Some("generateContent")
287 );
288 assert_eq!(
289 ServiceEndpoint::StreamGenerateContent.model_method(),
290 Some("streamGenerateContent")
291 );
292 assert_eq!(ServiceEndpoint::ListModels.model_method(), None);
293 assert_eq!(ServiceEndpoint::Files.model_method(), None);
294 }
295
296 #[tokio::test]
297 async fn vertex_ai_refreshable_token() {
298 use std::sync::atomic::{AtomicU32, Ordering};
299 let counter = std::sync::Arc::new(AtomicU32::new(0));
300 let c = counter.clone();
301 let auth = VertexAIAuth::with_token_refresher("proj", "us-central1", move || {
302 c.fetch_add(1, Ordering::SeqCst);
303 format!("token-{}", c.load(Ordering::SeqCst))
304 });
305 let h1 = auth.auth_headers().await.unwrap();
306 assert!(h1[0].1.starts_with("Bearer token-"));
307 let h2 = auth.auth_headers().await.unwrap();
308 assert!(h2[0].1.starts_with("Bearer token-"));
309 assert_eq!(counter.load(Ordering::SeqCst), 2);
311 }
312
313 #[test]
314 fn service_endpoint_requires_model() {
315 assert!(ServiceEndpoint::GenerateContent.requires_model());
316 assert!(ServiceEndpoint::CountTokens.requires_model());
317 assert!(!ServiceEndpoint::ListModels.requires_model());
318 assert!(!ServiceEndpoint::Files.requires_model());
319 }
320}