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}