Skip to main content

gcloud_sdk/
middleware.rs

1use crate::token_source::auth_token_generator::GoogleAuthTokenGenerator;
2use futures::{Future, TryFutureExt};
3use jiff::Timestamp;
4use std::pin::Pin;
5use std::sync::Arc;
6use std::task::{Context, Poll};
7use tonic::client::GrpcService;
8use tower::Service;
9use tower_layer::Layer;
10use tracing::*;
11
12#[derive(Clone)]
13pub struct GoogleAuthMiddlewareService<T>
14where
15    T: Clone,
16{
17    google_service: Option<T>,
18    token_generator: Arc<GoogleAuthTokenGenerator>,
19    cloud_resource_prefix: Option<String>,
20    user_agent: String,
21    x_goog_api_client: String,
22    additional_headers: hyper::header::HeaderMap,
23}
24
25impl<T> GoogleAuthMiddlewareService<T>
26where
27    T: Clone,
28{
29    pub fn new(
30        service: T,
31        token_generator: Arc<GoogleAuthTokenGenerator>,
32        cloud_resource_prefix: Option<String>,
33    ) -> GoogleAuthMiddlewareService<T> {
34        GoogleAuthMiddlewareService {
35            google_service: Some(service),
36            token_generator,
37            cloud_resource_prefix,
38            user_agent: format!("gcloud-sdk-rs/{}", env!("CARGO_PKG_VERSION")),
39            x_goog_api_client: format!("gcloud-sdk-rs/{}", env!("CARGO_PKG_VERSION")),
40            additional_headers: hyper::header::HeaderMap::new(),
41        }
42    }
43
44    pub fn set_user_agent(&mut self, user_agent: String) {
45        self.user_agent = user_agent;
46    }
47
48    pub fn set_x_goog_api_client(&mut self, x_goog_api_client: String) {
49        self.x_goog_api_client = x_goog_api_client;
50    }
51
52    pub fn append_user_agent(&mut self, user_agent: String) {
53        self.user_agent = format!("{} {}", self.user_agent, user_agent);
54    }
55
56    pub fn append_x_goog_api_client(&mut self, x_goog_api_client: String) {
57        self.x_goog_api_client = format!("{} {}", self.x_goog_api_client, x_goog_api_client);
58    }
59
60    pub fn set_additional_headers(&mut self, additional_headers: hyper::HeaderMap) {
61        self.additional_headers = additional_headers;
62    }
63}
64
65impl<T, RequestBody> Service<hyper::Request<RequestBody>> for GoogleAuthMiddlewareService<T>
66where
67    T: GrpcService<RequestBody> + Send + Clone + 'static,
68    T::Future: 'static + Send,
69    RequestBody: 'static + Send,
70    T::ResponseBody: 'static + Send,
71    T::Error: 'static + Send,
72{
73    type Response = hyper::Response<T::ResponseBody>;
74    type Error = Box<dyn std::error::Error + Send + Sync + 'static>;
75    type Future =
76        Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send + 'static>>;
77
78    fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
79        if let Some(ref mut google_service) = self.google_service.as_mut() {
80            google_service.poll_ready(cx).map_err(|e| e.into())
81        } else {
82            Poll::Pending
83        }
84    }
85
86    fn call(&mut self, mut req: hyper::Request<RequestBody>) -> Self::Future {
87        let generator = self.token_generator.clone();
88        let cloud_resource_prefix = self.cloud_resource_prefix.clone();
89        let user_agent = self.user_agent.clone();
90        let x_goog_api_client = self.x_goog_api_client.clone();
91        let additional_headers = self.additional_headers.clone();
92
93        if let Some(mut google_service) = self.google_service.take() {
94            self.google_service = Some(google_service.clone());
95            Box::pin(async move {
96                let begin_time = Timestamp::now();
97                let token = generator.create_token().await.map_err(Box::new)?;
98                let token_generated_time = Timestamp::now();
99                let headers = req.headers_mut();
100                headers.insert("authorization", token.header_value().parse()?);
101                if let Some(cloud_resource_prefix_value) = cloud_resource_prefix {
102                    headers.insert(
103                        "google-cloud-resource-prefix",
104                        cloud_resource_prefix_value.parse()?,
105                    );
106                }
107                headers.insert(hyper::header::USER_AGENT, user_agent.parse()?);
108                headers.insert("x-goog-api-client", x_goog_api_client.parse()?);
109
110                for (maybe_k, v) in additional_headers.into_iter() {
111                    if let Some(k) = maybe_k {
112                        headers.insert(k, v);
113                    }
114                }
115
116                let req_uri_str = req.uri().to_string();
117                google_service
118                    .call(req)
119                    .map_ok(|x| {
120                        let finished_time = Timestamp::now();
121                        debug!(
122                            "OK: {} took {}ms (incl. token gen: {}ms)",
123                            req_uri_str,
124                            finished_time.duration_since(begin_time).as_millis(),
125                            token_generated_time.duration_since(begin_time).as_millis()
126                        );
127                        x
128                    })
129                    .await
130                    .map_err(|e| {
131                        let finished_time = Timestamp::now();
132                        error!(
133                            "Err: {} took {}ms (incl. token gen: {}ms)",
134                            req_uri_str,
135                            finished_time.duration_since(begin_time).as_millis(),
136                            token_generated_time.duration_since(begin_time).as_millis()
137                        );
138                        e.into()
139                    })
140            })
141        } else {
142            panic!("Should never happen, system error");
143        }
144    }
145}
146
147pub struct GoogleAuthMiddlewareLayer {
148    pub token_generator: Arc<GoogleAuthTokenGenerator>,
149    pub cloud_resource_prefix: Option<String>,
150    pub user_agent: String,
151    pub x_goog_api_client: String,
152    pub additional_headers: hyper::header::HeaderMap,
153}
154
155impl GoogleAuthMiddlewareLayer {
156    pub fn new(
157        token_generator: GoogleAuthTokenGenerator,
158        cloud_resource_prefix: Option<String>,
159    ) -> Self {
160        GoogleAuthMiddlewareLayer {
161            token_generator: Arc::new(token_generator),
162            cloud_resource_prefix,
163            user_agent: format!("gcloud-sdk-rs/{}", env!("CARGO_PKG_VERSION")),
164            x_goog_api_client: format!("gcloud-sdk-rs/{}", env!("CARGO_PKG_VERSION")),
165            additional_headers: hyper::header::HeaderMap::new(),
166        }
167    }
168
169    pub fn amend_user_agent(mut self, user_agent: String) -> Self {
170        self.user_agent = format!("{} {}", self.user_agent, user_agent);
171        self
172    }
173
174    pub fn amend_x_goog_api_client(mut self, x_goog_api_client: String) -> Self {
175        self.x_goog_api_client = format!("{} {}", self.x_goog_api_client, x_goog_api_client);
176        self
177    }
178
179    pub fn set_additional_headers(&mut self, additional_headers: hyper::HeaderMap) {
180        self.additional_headers = additional_headers;
181    }
182}
183
184impl<S> Layer<S> for GoogleAuthMiddlewareLayer
185where
186    S: Clone,
187{
188    type Service = GoogleAuthMiddlewareService<S>;
189
190    fn layer(&self, service: S) -> GoogleAuthMiddlewareService<S> {
191        let mut middleware_service = GoogleAuthMiddlewareService::new(
192            service,
193            self.token_generator.clone(),
194            self.cloud_resource_prefix.clone(),
195        );
196        middleware_service.set_user_agent(self.user_agent.clone());
197        middleware_service.set_x_goog_api_client(self.x_goog_api_client.clone());
198        middleware_service.set_additional_headers(self.additional_headers.clone());
199        middleware_service
200    }
201}
202
203#[cfg(test)]
204mod tests {
205    use super::*;
206    use crate::token_source::{Source, Token, TokenSourceType};
207    use async_trait::async_trait;
208    use hyper::{Request, Response};
209    use jiff::{SignedDuration, Timestamp};
210    use secret_vault_value::SecretValue;
211    use std::convert::Infallible;
212
213    struct DummySource;
214
215    #[async_trait]
216    impl Source for DummySource {
217        async fn token(&self) -> crate::error::Result<Token> {
218            Ok(Token {
219                token_type: "Bearer".to_string(),
220                token: SecretValue::from("dummy-token"),
221                expiry: Timestamp::now() + SignedDuration::from_hours(1),
222            })
223        }
224    }
225
226    #[derive(Clone)]
227    struct DummyService {
228        tx: Arc<tokio::sync::mpsc::Sender<Request<String>>>,
229    }
230
231    impl Service<Request<String>> for DummyService {
232        type Response = Response<String>;
233        type Error = Infallible;
234        type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
235
236        fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
237            Poll::Ready(Ok(()))
238        }
239
240        fn call(&mut self, req: Request<String>) -> Self::Future {
241            let tx = self.tx.clone();
242            Box::pin(async move {
243                tx.send(req).await.unwrap();
244                Ok(Response::builder()
245                    .status(200)
246                    .body("".to_string())
247                    .unwrap())
248            })
249        }
250    }
251
252    #[tokio::test]
253    async fn test_headers_presence() {
254        let token_generator = GoogleAuthTokenGenerator::new(
255            TokenSourceType::ExternalSource(Box::new(DummySource)),
256            vec![],
257        )
258        .await
259        .unwrap();
260
261        let (tx, mut rx) = tokio::sync::mpsc::channel(1);
262        let dummy_service = DummyService { tx: Arc::new(tx) };
263        let mut service =
264            GoogleAuthMiddlewareService::new(dummy_service, Arc::new(token_generator), None);
265
266        let req = Request::builder()
267            .uri("http://example.com")
268            .body("".to_string())
269            .unwrap();
270
271        tower::Service::call(&mut service, req).await.unwrap();
272
273        let captured_req = rx.recv().await.unwrap();
274        let expected_default = format!("gcloud-sdk-rs/{}", env!("CARGO_PKG_VERSION"));
275        assert_eq!(
276            captured_req
277                .headers()
278                .get(hyper::header::USER_AGENT)
279                .unwrap(),
280            expected_default.as_str()
281        );
282        assert_eq!(
283            captured_req.headers().get("x-goog-api-client").unwrap(),
284            expected_default.as_str()
285        );
286        assert_eq!(
287            captured_req.headers().get("authorization").unwrap(),
288            "Bearer dummy-token"
289        );
290    }
291
292    #[tokio::test]
293    async fn test_headers_amend() {
294        let token_generator = GoogleAuthTokenGenerator::new(
295            TokenSourceType::ExternalSource(Box::new(DummySource)),
296            vec![],
297        )
298        .await
299        .unwrap();
300
301        let (tx, mut rx) = tokio::sync::mpsc::channel(1);
302        let dummy_service = DummyService { tx: Arc::new(tx) };
303
304        let layer = GoogleAuthMiddlewareLayer::new(token_generator, None)
305            .amend_user_agent("extra-ua".to_string())
306            .amend_x_goog_api_client("extra-client".to_string());
307
308        let mut service = layer.layer(dummy_service);
309
310        let req = Request::builder()
311            .uri("http://example.com")
312            .body("".to_string())
313            .unwrap();
314
315        tower::Service::call(&mut service, req).await.unwrap();
316
317        let captured_req = rx.recv().await.unwrap();
318        let expected_ua = format!("gcloud-sdk-rs/{} extra-ua", env!("CARGO_PKG_VERSION"));
319        let expected_client = format!("gcloud-sdk-rs/{} extra-client", env!("CARGO_PKG_VERSION"));
320
321        assert_eq!(
322            captured_req
323                .headers()
324                .get(hyper::header::USER_AGENT)
325                .unwrap(),
326            expected_ua.as_str()
327        );
328        assert_eq!(
329            captured_req.headers().get("x-goog-api-client").unwrap(),
330            expected_client.as_str()
331        );
332    }
333
334    #[tokio::test]
335    async fn test_additional_headers() {
336        let token_generator = GoogleAuthTokenGenerator::new(
337            TokenSourceType::ExternalSource(Box::new(DummySource)),
338            vec![],
339        )
340        .await
341        .unwrap();
342
343        let (tx, mut rx) = tokio::sync::mpsc::channel(1);
344        let dummy_service = DummyService { tx: Arc::new(tx) };
345        let mut service =
346            GoogleAuthMiddlewareService::new(dummy_service, Arc::new(token_generator), None);
347        let mut test_headers = hyper::HeaderMap::new();
348        test_headers.insert("x-test-header", "test-value".parse().unwrap());
349        service.set_additional_headers(test_headers);
350
351        let req = Request::builder()
352            .uri("http://example.com")
353            .body("".to_string())
354            .unwrap();
355
356        tower::Service::call(&mut service, req).await.unwrap();
357
358        let captured_req = rx.recv().await.unwrap();
359        assert_eq!(
360            captured_req.headers().get("x-test-header").unwrap(),
361            "test-value"
362        );
363    }
364}