Skip to main content

gcloud_sdk/
middleware.rs

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
32/// Appends `extra` to whatever `name` currently holds in `headers` (space separated),
33/// or sets it outright if absent.
34fn 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    /// Every header added to each request except `authorization`, already validated.
51    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        // tower's documented idiom for a `Clone` inner service: the instance we already
130        // polled ready goes into the future, and the service keeps a fresh clone.
131        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            // Each name in `headers` fully replaces whatever the request already carries
142            // under it (multi-valued or not); `iter()` yields one pair per value, so the
143            // names are cleared first and then appended rather than repeatedly `insert`ed,
144            // which would silently drop every value but the last for a repeated name.
145            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}