1use crate::token_source::auth_token_generator::GoogleAuthTokenGenerator;
2use futures::{Future, TryFutureExt};
3use hyper::header::{HeaderMap, HeaderName, HeaderValue, USER_AGENT};
4use jiff::Timestamp;
5use std::pin::Pin;
6use std::sync::Arc;
7use std::task::{Context, Poll};
8use tonic::client::GrpcService;
9use tower::Service;
10use tower_layer::Layer;
11use tracing::*;
12
13const X_GOOG_API_CLIENT: HeaderName = HeaderName::from_static("x-goog-api-client");
14const GOOGLE_CLOUD_RESOURCE_PREFIX: HeaderName =
15 HeaderName::from_static("google-cloud-resource-prefix");
16
17fn default_headers(cloud_resource_prefix: Option<String>) -> crate::error::Result<HeaderMap> {
18 let mut headers = HeaderMap::new();
19 let default_agent =
20 HeaderValue::from_static(concat!("gcloud-sdk-rs/", env!("CARGO_PKG_VERSION")));
21 headers.insert(USER_AGENT, default_agent.clone());
22 headers.insert(X_GOOG_API_CLIENT, default_agent);
23 if let Some(prefix) = cloud_resource_prefix {
24 headers.insert(
25 GOOGLE_CLOUD_RESOURCE_PREFIX,
26 HeaderValue::from_str(&prefix)?,
27 );
28 }
29 Ok(headers)
30}
31
32fn append_header_value(
35 headers: &HeaderMap,
36 name: &HeaderName,
37 extra: &str,
38) -> crate::error::Result<HeaderValue> {
39 let combined = match headers.get(name).and_then(|v| v.to_str().ok()) {
40 Some(current) => format!("{current} {extra}"),
41 None => extra.to_string(),
42 };
43 Ok(HeaderValue::from_str(&combined)?)
44}
45
46#[derive(Clone)]
47pub struct GoogleAuthMiddlewareService<T> {
48 inner: T,
49 token_generator: Arc<GoogleAuthTokenGenerator>,
50 headers: Arc<HeaderMap>,
52}
53
54impl<T> GoogleAuthMiddlewareService<T> {
55 pub fn new(
56 service: T,
57 token_generator: Arc<GoogleAuthTokenGenerator>,
58 cloud_resource_prefix: Option<String>,
59 ) -> crate::error::Result<GoogleAuthMiddlewareService<T>> {
60 Ok(GoogleAuthMiddlewareService {
61 inner: service,
62 token_generator,
63 headers: Arc::new(default_headers(cloud_resource_prefix)?),
64 })
65 }
66
67 pub fn set_user_agent(&mut self, user_agent: String) -> crate::error::Result<()> {
68 let value = HeaderValue::from_str(&user_agent)?;
69 Arc::make_mut(&mut self.headers).insert(USER_AGENT, value);
70 Ok(())
71 }
72
73 pub fn set_x_goog_api_client(&mut self, x_goog_api_client: String) -> crate::error::Result<()> {
74 let value = HeaderValue::from_str(&x_goog_api_client)?;
75 Arc::make_mut(&mut self.headers).insert(X_GOOG_API_CLIENT, value);
76 Ok(())
77 }
78
79 pub fn set_cloud_resource_prefix(
80 &mut self,
81 cloud_resource_prefix: String,
82 ) -> crate::error::Result<()> {
83 let value = HeaderValue::from_str(&cloud_resource_prefix)?;
84 Arc::make_mut(&mut self.headers).insert(GOOGLE_CLOUD_RESOURCE_PREFIX, value);
85 Ok(())
86 }
87
88 pub fn append_user_agent(&mut self, user_agent: String) -> crate::error::Result<()> {
89 let value = append_header_value(&self.headers, &USER_AGENT, &user_agent)?;
90 Arc::make_mut(&mut self.headers).insert(USER_AGENT, value);
91 Ok(())
92 }
93
94 pub fn append_x_goog_api_client(
95 &mut self,
96 x_goog_api_client: String,
97 ) -> crate::error::Result<()> {
98 let value = append_header_value(&self.headers, &X_GOOG_API_CLIENT, &x_goog_api_client)?;
99 Arc::make_mut(&mut self.headers).insert(X_GOOG_API_CLIENT, value);
100 Ok(())
101 }
102
103 pub fn set_additional_headers(&mut self, additional_headers: HeaderMap) {
104 Arc::make_mut(&mut self.headers).extend(additional_headers);
105 }
106}
107
108impl<T, RequestBody> Service<hyper::Request<RequestBody>> for GoogleAuthMiddlewareService<T>
109where
110 T: GrpcService<RequestBody> + Send + Clone + 'static,
111 T::Future: 'static + Send,
112 RequestBody: 'static + Send,
113 T::ResponseBody: 'static + Send,
114 T::Error: 'static + Send,
115{
116 type Response = hyper::Response<T::ResponseBody>;
117 type Error = Box<dyn std::error::Error + Send + Sync + 'static>;
118 type Future =
119 Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send + 'static>>;
120
121 fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
122 self.inner.poll_ready(cx).map_err(Into::into)
123 }
124
125 fn call(&mut self, mut req: hyper::Request<RequestBody>) -> Self::Future {
126 let generator = Arc::clone(&self.token_generator);
127 let headers = Arc::clone(&self.headers);
128
129 let clone = self.inner.clone();
132 let mut inner = std::mem::replace(&mut self.inner, clone);
133
134 Box::pin(async move {
135 let begin_time = Timestamp::now();
136 let authorization = generator.authorization_header().await.map_err(Box::new)?;
137 let token_generated_time = Timestamp::now();
138
139 let req_headers = req.headers_mut();
140 req_headers.insert(hyper::header::AUTHORIZATION, authorization);
141 for name in headers.keys() {
146 req_headers.remove(name);
147 }
148 req_headers.extend(
149 headers
150 .iter()
151 .map(|(name, value)| (name.clone(), value.clone())),
152 );
153
154 let req_uri = req.uri().clone();
155 inner
156 .call(req)
157 .map_ok(|x| {
158 let finished_time = Timestamp::now();
159 debug!(
160 %req_uri,
161 "OK: took {}ms (incl. token gen: {}ms)",
162 finished_time.duration_since(begin_time).as_millis(),
163 token_generated_time.duration_since(begin_time).as_millis()
164 );
165 x
166 })
167 .await
168 .map_err(|e| {
169 let finished_time = Timestamp::now();
170 error!(
171 %req_uri,
172 "Err: took {}ms (incl. token gen: {}ms)",
173 finished_time.duration_since(begin_time).as_millis(),
174 token_generated_time.duration_since(begin_time).as_millis()
175 );
176 e.into()
177 })
178 })
179 }
180}
181
182pub struct GoogleAuthMiddlewareLayer {
183 token_generator: Arc<GoogleAuthTokenGenerator>,
184 headers: Arc<HeaderMap>,
185}
186
187impl GoogleAuthMiddlewareLayer {
188 pub fn new(
189 token_generator: GoogleAuthTokenGenerator,
190 cloud_resource_prefix: Option<String>,
191 ) -> crate::error::Result<Self> {
192 Ok(GoogleAuthMiddlewareLayer {
193 token_generator: Arc::new(token_generator),
194 headers: Arc::new(default_headers(cloud_resource_prefix)?),
195 })
196 }
197
198 pub fn amend_user_agent(mut self, user_agent: String) -> crate::error::Result<Self> {
199 let value = append_header_value(&self.headers, &USER_AGENT, &user_agent)?;
200 Arc::make_mut(&mut self.headers).insert(USER_AGENT, value);
201 Ok(self)
202 }
203
204 pub fn amend_x_goog_api_client(
205 mut self,
206 x_goog_api_client: String,
207 ) -> crate::error::Result<Self> {
208 let value = append_header_value(&self.headers, &X_GOOG_API_CLIENT, &x_goog_api_client)?;
209 Arc::make_mut(&mut self.headers).insert(X_GOOG_API_CLIENT, value);
210 Ok(self)
211 }
212
213 pub fn set_additional_headers(&mut self, additional_headers: HeaderMap) {
214 Arc::make_mut(&mut self.headers).extend(additional_headers);
215 }
216}
217
218impl<S> Layer<S> for GoogleAuthMiddlewareLayer {
219 type Service = GoogleAuthMiddlewareService<S>;
220
221 fn layer(&self, service: S) -> GoogleAuthMiddlewareService<S> {
222 GoogleAuthMiddlewareService {
223 inner: service,
224 token_generator: Arc::clone(&self.token_generator),
225 headers: Arc::clone(&self.headers),
226 }
227 }
228}
229
230#[cfg(test)]
231mod tests {
232 use super::*;
233 use crate::token_source::{Source, Token, TokenSourceType};
234 use async_trait::async_trait;
235 use hyper::{Request, Response};
236 use jiff::{SignedDuration, Timestamp};
237 use secret_vault_value::SecretValue;
238 use std::convert::Infallible;
239
240 struct DummySource;
241
242 #[async_trait]
243 impl Source for DummySource {
244 async fn token(&self) -> crate::error::Result<Token> {
245 Ok(Token {
246 token_type: "Bearer".to_string(),
247 token: SecretValue::from("dummy-token"),
248 expiry: Timestamp::now() + SignedDuration::from_hours(1),
249 })
250 }
251 }
252
253 #[derive(Clone)]
254 struct DummyService {
255 tx: Arc<tokio::sync::mpsc::Sender<Request<String>>>,
256 }
257
258 impl Service<Request<String>> for DummyService {
259 type Response = Response<String>;
260 type Error = Infallible;
261 type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
262
263 fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
264 Poll::Ready(Ok(()))
265 }
266
267 fn call(&mut self, req: Request<String>) -> Self::Future {
268 let tx = self.tx.clone();
269 Box::pin(async move {
270 tx.send(req).await.unwrap();
271 Ok(Response::builder()
272 .status(200)
273 .body("".to_string())
274 .unwrap())
275 })
276 }
277 }
278
279 #[tokio::test]
280 async fn test_headers_presence() {
281 let token_generator = GoogleAuthTokenGenerator::new(
282 TokenSourceType::ExternalSource(Box::new(DummySource)),
283 vec![],
284 )
285 .await
286 .unwrap();
287
288 let (tx, mut rx) = tokio::sync::mpsc::channel(1);
289 let dummy_service = DummyService { tx: Arc::new(tx) };
290 let mut service =
291 GoogleAuthMiddlewareService::new(dummy_service, Arc::new(token_generator), None)
292 .unwrap();
293
294 let req = Request::builder()
295 .uri("http://example.com")
296 .body("".to_string())
297 .unwrap();
298
299 tower::Service::call(&mut service, req).await.unwrap();
300
301 let captured_req = rx.recv().await.unwrap();
302 let expected_default = format!("gcloud-sdk-rs/{}", env!("CARGO_PKG_VERSION"));
303 assert_eq!(
304 captured_req
305 .headers()
306 .get(hyper::header::USER_AGENT)
307 .unwrap(),
308 expected_default.as_str()
309 );
310 assert_eq!(
311 captured_req.headers().get("x-goog-api-client").unwrap(),
312 expected_default.as_str()
313 );
314 assert_eq!(
315 captured_req.headers().get("authorization").unwrap(),
316 "Bearer dummy-token"
317 );
318 }
319
320 #[tokio::test]
321 async fn authorization_header_is_marked_sensitive() {
322 let token_generator = GoogleAuthTokenGenerator::new(
323 TokenSourceType::ExternalSource(Box::new(DummySource)),
324 vec![],
325 )
326 .await
327 .unwrap();
328
329 let (tx, mut rx) = tokio::sync::mpsc::channel(1);
330 let dummy_service = DummyService { tx: Arc::new(tx) };
331 let mut service =
332 GoogleAuthMiddlewareService::new(dummy_service, Arc::new(token_generator), None)
333 .unwrap();
334
335 let req = Request::builder()
336 .uri("http://example.com")
337 .body("".to_string())
338 .unwrap();
339
340 tower::Service::call(&mut service, req).await.unwrap();
341
342 let captured_req = rx.recv().await.unwrap();
343 assert!(captured_req
344 .headers()
345 .get("authorization")
346 .unwrap()
347 .is_sensitive());
348 }
349
350 #[tokio::test]
351 async fn test_headers_amend() {
352 let token_generator = GoogleAuthTokenGenerator::new(
353 TokenSourceType::ExternalSource(Box::new(DummySource)),
354 vec![],
355 )
356 .await
357 .unwrap();
358
359 let (tx, mut rx) = tokio::sync::mpsc::channel(1);
360 let dummy_service = DummyService { tx: Arc::new(tx) };
361
362 let layer = GoogleAuthMiddlewareLayer::new(token_generator, None)
363 .unwrap()
364 .amend_user_agent("extra-ua".to_string())
365 .unwrap()
366 .amend_x_goog_api_client("extra-client".to_string())
367 .unwrap();
368
369 let mut service = layer.layer(dummy_service);
370
371 let req = Request::builder()
372 .uri("http://example.com")
373 .body("".to_string())
374 .unwrap();
375
376 tower::Service::call(&mut service, req).await.unwrap();
377
378 let captured_req = rx.recv().await.unwrap();
379 let expected_ua = format!("gcloud-sdk-rs/{} extra-ua", env!("CARGO_PKG_VERSION"));
380 let expected_client = format!("gcloud-sdk-rs/{} extra-client", env!("CARGO_PKG_VERSION"));
381
382 assert_eq!(
383 captured_req
384 .headers()
385 .get(hyper::header::USER_AGENT)
386 .unwrap(),
387 expected_ua.as_str()
388 );
389 assert_eq!(
390 captured_req.headers().get("x-goog-api-client").unwrap(),
391 expected_client.as_str()
392 );
393 }
394
395 #[tokio::test]
396 async fn invalid_user_agent_is_rejected_at_setter() {
397 let token_generator = GoogleAuthTokenGenerator::new(
398 TokenSourceType::ExternalSource(Box::new(DummySource)),
399 vec![],
400 )
401 .await
402 .unwrap();
403
404 let layer_result = GoogleAuthMiddlewareLayer::new(token_generator, None)
405 .unwrap()
406 .amend_user_agent("bad\nvalue".to_string());
407
408 match layer_result {
409 Err(e) => assert!(matches!(
410 e.into_kind(),
411 crate::error::ErrorKind::HeaderValue(_)
412 )),
413 Ok(_) => panic!("expected an invalid header value to be rejected"),
414 }
415 }
416
417 #[tokio::test]
418 async fn amended_clone_does_not_change_sibling() {
419 let token_generator = GoogleAuthTokenGenerator::new(
420 TokenSourceType::ExternalSource(Box::new(DummySource)),
421 vec![],
422 )
423 .await
424 .unwrap();
425
426 let (tx, mut rx) = tokio::sync::mpsc::channel(1);
427 let dummy_service = DummyService { tx: Arc::new(tx) };
428 let base_service =
429 GoogleAuthMiddlewareService::new(dummy_service, Arc::new(token_generator), None)
430 .unwrap();
431
432 let mut amended = base_service.clone();
433 amended.append_user_agent("extra".to_string()).unwrap();
434
435 let mut sibling = base_service.clone();
436
437 let req = Request::builder()
438 .uri("http://example.com")
439 .body("".to_string())
440 .unwrap();
441
442 tower::Service::call(&mut sibling, req).await.unwrap();
443
444 let captured_req = rx.recv().await.unwrap();
445 let expected_default = format!("gcloud-sdk-rs/{}", env!("CARGO_PKG_VERSION"));
446 assert_eq!(
447 captured_req
448 .headers()
449 .get(hyper::header::USER_AGENT)
450 .unwrap(),
451 expected_default.as_str()
452 );
453 }
454
455 #[tokio::test]
456 async fn test_additional_headers() {
457 let token_generator = GoogleAuthTokenGenerator::new(
458 TokenSourceType::ExternalSource(Box::new(DummySource)),
459 vec![],
460 )
461 .await
462 .unwrap();
463
464 let (tx, mut rx) = tokio::sync::mpsc::channel(1);
465 let dummy_service = DummyService { tx: Arc::new(tx) };
466 let mut service =
467 GoogleAuthMiddlewareService::new(dummy_service, Arc::new(token_generator), None)
468 .unwrap();
469 let mut test_headers = hyper::HeaderMap::new();
470 test_headers.insert("x-test-header", "test-value".parse().unwrap());
471 service.set_additional_headers(test_headers);
472
473 let req = Request::builder()
474 .uri("http://example.com")
475 .body("".to_string())
476 .unwrap();
477
478 tower::Service::call(&mut service, req).await.unwrap();
479
480 let captured_req = rx.recv().await.unwrap();
481 assert_eq!(
482 captured_req.headers().get("x-test-header").unwrap(),
483 "test-value"
484 );
485 }
486
487 #[tokio::test]
488 async fn additional_headers_keep_every_value_of_a_repeated_name() {
489 let token_generator = GoogleAuthTokenGenerator::new(
490 TokenSourceType::ExternalSource(Box::new(DummySource)),
491 vec![],
492 )
493 .await
494 .unwrap();
495
496 let (tx, mut rx) = tokio::sync::mpsc::channel(1);
497 let dummy_service = DummyService { tx: Arc::new(tx) };
498 let mut service =
499 GoogleAuthMiddlewareService::new(dummy_service, Arc::new(token_generator), None)
500 .unwrap();
501
502 let mut test_headers = hyper::HeaderMap::new();
503 test_headers.append("x-multi", "first".parse().unwrap());
504 test_headers.append("x-multi", "second".parse().unwrap());
505 service.set_additional_headers(test_headers);
506
507 let req = Request::builder()
508 .uri("http://example.com")
509 .body("".to_string())
510 .unwrap();
511
512 tower::Service::call(&mut service, req).await.unwrap();
513
514 let captured_req = rx.recv().await.unwrap();
515 let values: Vec<&str> = captured_req
516 .headers()
517 .get_all("x-multi")
518 .iter()
519 .map(|v| v.to_str().unwrap())
520 .collect();
521 assert_eq!(values, vec!["first", "second"]);
522 }
523}