1use std::net::{IpAddr, SocketAddr};
4use std::sync::Arc;
5use std::time::Duration;
6
7use arc_swap::ArcSwapOption;
8use async_trait::async_trait;
9use reqwest::header::HeaderValue;
10use tollgate_store::StoreError;
11use zeroize::Zeroizing;
12
13#[async_trait]
16pub trait BearerProvider: Send + Sync {
17 async fn token(&self) -> Result<BearerToken, StoreError>;
25}
26
27#[derive(Clone)]
29pub struct BearerToken(Zeroizing<String>);
30
31impl BearerToken {
32 pub fn new(token: impl Into<String>) -> Result<Self, StoreError> {
40 let token = Zeroizing::new(token.into());
41 if token.is_empty()
42 || token.len() > 16 * 1024 - 7
43 || !token.bytes().all(|b| b.is_ascii_graphic())
44 {
45 return Err(StoreError("invalid bearer token framing or length".into()));
46 }
47 Ok(Self(token))
48 }
49
50 pub(crate) fn header(&self) -> HeaderValue {
51 let framed = Zeroizing::new(format!("Bearer {}", self.0.as_str()));
53 let mut header =
54 HeaderValue::from_str(&framed).expect("validated bearer token is a header value");
55 header.set_sensitive(true);
56 header
57 }
58}
59
60impl std::fmt::Debug for BearerToken {
61 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
62 f.write_str("BearerToken([redacted])")
63 }
64}
65
66pub struct StaticBearer(ArcSwapOption<BearerToken>);
69
70impl StaticBearer {
71 pub fn new(token: BearerToken) -> Arc<Self> {
73 Arc::new(Self(ArcSwapOption::from(Some(Arc::new(token)))))
74 }
75 pub fn replace(&self, token: BearerToken) {
77 self.0.store(Some(Arc::new(token)));
78 }
79 pub fn revoke(&self) {
82 self.0.store(None);
83 }
84}
85
86#[async_trait]
87impl BearerProvider for StaticBearer {
88 async fn token(&self) -> Result<BearerToken, StoreError> {
89 self.0
90 .load_full()
91 .map(|token| (*token).clone())
92 .ok_or_else(|| StoreError("control-plane credential was revoked".into()))
93 }
94}
95
96pub struct GoogleIdentity {
100 client: reqwest::Client,
101 audience: String,
102 cached: tokio::sync::Mutex<Option<(tokio::time::Instant, BearerToken)>>,
103}
104
105impl GoogleIdentity {
106 pub fn new(audience: impl Into<String>) -> Result<Arc<Self>, StoreError> {
115 let audience = audience.into();
116 if audience.is_empty() {
117 return Err(StoreError(
118 "Google identity audience must not be empty".into(),
119 ));
120 }
121 let client = reqwest::Client::builder()
122 .no_proxy()
123 .redirect(reqwest::redirect::Policy::none())
124 .connect_timeout(Duration::from_secs(2))
125 .timeout(Duration::from_secs(5))
126 .build()
127 .map_err(|_| StoreError("cannot configure metadata client".into()))?;
128 Ok(Arc::new(Self {
129 client,
130 audience,
131 cached: tokio::sync::Mutex::new(None),
132 }))
133 }
134}
135
136#[async_trait]
137impl BearerProvider for GoogleIdentity {
138 async fn token(&self) -> Result<BearerToken, StoreError> {
139 self.token_from("http://metadata.google.internal/computeMetadata/v1/instance/service-accounts/default/identity", tokio::time::Instant::now).await
140 }
141}
142
143impl GoogleIdentity {
144 async fn token_from(
148 &self,
149 endpoint: &str,
150 clock: impl FnOnce() -> tokio::time::Instant,
151 ) -> Result<BearerToken, StoreError> {
152 let mut cached = self.cached.lock().await;
153 let now = clock();
154 if let Some((until, token)) = &*cached
155 && now < *until
156 {
157 return Ok(token.clone());
158 }
159 let mut response = self
160 .client
161 .get(endpoint)
162 .header("Metadata-Flavor", "Google")
163 .query(&[("audience", self.audience.as_str()), ("format", "full")])
164 .send()
165 .await
166 .map_err(|_| StoreError("Google metadata identity request failed".into()))?;
167 if !response.status().is_success()
168 || response
169 .headers()
170 .get("Metadata-Flavor")
171 .is_none_or(|v| v != "Google")
172 {
173 return Err(StoreError(
174 "Google metadata identity request was refused".into(),
175 ));
176 }
177 let mut bytes = Zeroizing::new(Vec::new());
178 while let Some(chunk) = response
179 .chunk()
180 .await
181 .map_err(|_| StoreError("Google metadata identity response interrupted".into()))?
182 {
183 if bytes.len() + chunk.len() > 16 * 1024 - 7 {
184 return Err(StoreError(
185 "Google metadata identity response too large".into(),
186 ));
187 }
188 bytes.extend_from_slice(&chunk);
189 }
190 let token = BearerToken::new(
191 std::str::from_utf8(&bytes)
192 .map_err(|_| StoreError("invalid metadata identity encoding".into()))?
193 .to_owned(),
194 )?;
195 *cached = Some((now + Duration::from_secs(60), token.clone()));
199 Ok(token)
200 }
201}
202
203pub struct HttpStoreConfig {
212 pub connect_timeout: Duration,
217 pub request_timeout: Duration,
225 pub root_ca_pem: Option<Vec<u8>>,
229 pub identity_pem: Option<Zeroizing<Vec<u8>>>,
233 pub bearer: Option<Arc<dyn BearerProvider>>,
236}
237
238impl Default for HttpStoreConfig {
239 fn default() -> Self {
240 Self {
241 connect_timeout: Duration::from_secs(2),
242 request_timeout: Duration::from_secs(10),
243 root_ca_pem: None,
244 identity_pem: None,
245 bearer: None,
246 }
247 }
248}
249
250pub(crate) fn client(
251 base: &str,
252 config: &HttpStoreConfig,
253) -> Result<(String, reqwest::Client), StoreError> {
254 let url =
255 reqwest::Url::parse(base).map_err(|_| StoreError("invalid control-plane URL".into()))?;
256 if !matches!(url.scheme(), "http" | "https")
257 || !url.username().is_empty()
258 || url.password().is_some()
259 || url.query().is_some()
260 || url.fragment().is_some()
261 {
262 return Err(StoreError(
263 "control-plane URL must be HTTP(S), with no userinfo, query or fragment".into(),
264 ));
265 }
266 let host = url
267 .host_str()
268 .ok_or_else(|| StoreError("control-plane URL needs a host".into()))?;
269 let localhost = host == "localhost";
270 let loopback = localhost
271 || host
272 .trim_matches(['[', ']'])
273 .parse::<IpAddr>()
274 .is_ok_and(|ip| match ip {
275 IpAddr::V4(ip) => ip.is_loopback(),
276 IpAddr::V6(ip) => {
277 ip.is_loopback() || ip.to_ipv4_mapped().is_some_and(|ip| ip.is_loopback())
278 }
279 });
280 if url.scheme() == "http"
281 && (!loopback || config.identity_pem.is_some() || config.root_ca_pem.is_some())
282 {
283 return Err(StoreError(
284 "plaintext is allowed only on loopback without TLS configuration".into(),
285 ));
286 }
287 for duration in [config.connect_timeout, config.request_timeout] {
288 if duration.is_zero() || std::time::Instant::now().checked_add(duration).is_none() {
289 return Err(StoreError(
290 "HTTP deadlines must be positive and representable".into(),
291 ));
292 }
293 }
294 let mut builder = reqwest::Client::builder()
295 .use_rustls_tls()
296 .no_proxy()
297 .redirect(reqwest::redirect::Policy::none())
298 .connect_timeout(config.connect_timeout)
299 .timeout(config.request_timeout);
300 if localhost {
303 builder = builder.resolve(
304 "localhost",
305 SocketAddr::from(([127, 0, 0, 1], url.port_or_known_default().unwrap_or(80))),
306 );
307 }
308 if let Some(pem) = &config.root_ca_pem {
309 let certificates = reqwest::Certificate::from_pem_bundle(pem)
310 .map_err(|_| StoreError("invalid control-plane CA PEM".into()))?;
311 if certificates.is_empty() {
312 return Err(StoreError("control-plane CA PEM is empty".into()));
313 }
314 builder = builder.tls_built_in_root_certs(false);
315 for certificate in certificates {
316 builder = builder.add_root_certificate(certificate);
317 }
318 }
319 if let Some(pem) = &config.identity_pem {
320 builder = builder.identity(
321 reqwest::Identity::from_pem(pem)
322 .map_err(|_| StoreError("invalid control-plane identity PEM".into()))?,
323 );
324 }
325 Ok((
326 url.as_str().trim_end_matches('/').to_owned(),
327 builder
328 .build()
329 .map_err(|_| StoreError("invalid HTTP/TLS client configuration".into()))?,
330 ))
331}
332
333#[cfg(test)]
334mod tests {
335 use super::*;
336 use tokio::io::{AsyncReadExt, AsyncWriteExt};
337
338 #[test]
339 fn client_validation_rejects_each_unsafe_url_component_independently() {
340 for url in [
341 "https://user@example.com",
342 "https://:password@example.com",
343 "https://example.com?query",
344 "https://example.com#fragment",
345 "http://example.com",
346 "http://[2001:db8::1]",
347 "http://[::ffff:192.0.2.1]",
348 ] {
349 assert!(client(url, &HttpStoreConfig::default()).is_err(), "{url}");
350 }
351 for url in [
352 "http://localhost",
353 "http://127.0.0.2",
354 "http://[::1]",
355 "http://[::ffff:127.0.0.1]",
356 ] {
357 assert!(client(url, &HttpStoreConfig::default()).is_ok(), "{url}");
358 for (root, identity) in [(true, false), (false, true), (true, true)] {
359 let config = HttpStoreConfig {
360 root_ca_pem: root.then(Vec::new),
361 identity_pem: identity.then(|| Zeroizing::new(Vec::new())),
362 ..Default::default()
363 };
364 let error = client(url, &config).err().unwrap();
365 assert_eq!(
366 error.0,
367 "plaintext is allowed only on loopback without TLS configuration"
368 );
369 }
370 }
371 }
372
373 struct Metadata {
374 endpoint: String,
375 response: Arc<std::sync::Mutex<Vec<u8>>>,
376 requests: tokio::sync::mpsc::UnboundedReceiver<String>,
377 task: tokio::task::JoinHandle<()>,
378 }
379 impl Drop for Metadata {
380 fn drop(&mut self) {
381 self.task.abort();
382 }
383 }
384 impl Metadata {
385 async fn start() -> Self {
386 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
387 let endpoint = format!(
388 "http://{}/computeMetadata/v1/instance/service-accounts/default/identity",
389 listener.local_addr().unwrap()
390 );
391 let response = Arc::new(std::sync::Mutex::new(Vec::new()));
392 let replies = response.clone();
393 let (requests, received) = tokio::sync::mpsc::unbounded_channel();
394 let task = tokio::spawn(async move {
395 loop {
396 let (mut stream, _) = listener.accept().await.unwrap();
397 let mut request = Vec::new();
398 while !request.ends_with(b"\r\n\r\n") {
399 let mut byte = [0];
400 if stream.read(&mut byte).await.unwrap() == 0 {
401 break;
402 }
403 request.extend_from_slice(&byte);
404 }
405 requests.send(String::from_utf8(request).unwrap()).unwrap();
406 let reply = replies.lock().unwrap().clone();
407 if stream.write_all(&reply).await.is_err() {
410 continue;
411 }
412 }
413 });
414 Self {
415 endpoint,
416 response,
417 requests: received,
418 task,
419 }
420 }
421 fn reply(&self, status: u16, flavor: Option<&str>, body: &[u8]) {
422 let mut reply = format!(
423 "HTTP/1.1 {status} Fixture\r\nConnection: close\r\nContent-Length: {}\r\n",
424 body.len()
425 );
426 if let Some(flavor) = flavor {
427 reply.push_str(&format!("Metadata-Flavor: {flavor}\r\n"));
428 }
429 reply.push_str("\r\n");
430 let mut bytes = reply.into_bytes();
431 bytes.extend_from_slice(body);
432 *self.response.lock().unwrap() = bytes;
433 }
434 }
435
436 #[tokio::test]
437 async fn google_metadata_cache_refreshes_at_its_exact_deadline_and_never_caches_failure() {
438 assert!(GoogleIdentity::new("").is_err());
439 let provider = GoogleIdentity::new("https://control.example.test").unwrap();
440 let mut server = Metadata::start().await;
441 server.reply(200, Some("Google"), b"fixture-token-one");
442 let now = tokio::time::Instant::now();
443 let first = provider.token_from(&server.endpoint, || now).await.unwrap();
444 assert_eq!(first.header(), "Bearer fixture-token-one");
445 let request = server.requests.recv().await.unwrap();
446 let path = request
447 .lines()
448 .next()
449 .unwrap()
450 .split_whitespace()
451 .nth(1)
452 .unwrap();
453 let url = reqwest::Url::parse(&format!("http://metadata.fixture{path}")).unwrap();
454 assert_eq!(
455 url.path(),
456 "/computeMetadata/v1/instance/service-accounts/default/identity"
457 );
458 assert_eq!(
459 url.query_pairs()
460 .collect::<std::collections::HashMap<_, _>>()
461 .get("audience")
462 .unwrap(),
463 "https://control.example.test"
464 );
465 assert_eq!(
466 url.query_pairs()
467 .collect::<std::collections::HashMap<_, _>>()
468 .get("format")
469 .unwrap(),
470 "full"
471 );
472 assert!(
473 request
474 .to_ascii_lowercase()
475 .contains("metadata-flavor: google\r\n")
476 );
477 server.reply(503, Some("Google"), b"unavailable");
478 let cached = provider
479 .token_from(&server.endpoint, || now + Duration::from_secs(59))
480 .await
481 .unwrap();
482 assert_eq!(cached.header(), first.header());
483 assert!(server.requests.try_recv().is_err());
484 assert!(
485 provider
486 .token_from(&server.endpoint, || now + Duration::from_secs(60))
487 .await
488 .is_err()
489 );
490 server.requests.recv().await.unwrap();
491 server.reply(200, Some("Google"), b"fixture-token-two");
492 let replacement = provider
493 .token_from(&server.endpoint, || now + Duration::from_secs(60))
494 .await
495 .unwrap();
496 assert_eq!(replacement.header(), "Bearer fixture-token-two");
497 server.requests.recv().await.unwrap();
498 assert!(server.requests.try_recv().is_err());
499 }
500
501 #[tokio::test]
502 async fn metadata_requires_success_google_provenance_and_a_complete_bounded_token() {
503 let server = Metadata::start().await;
504 for (status, flavor, body, accepted) in [
505 (200, Some("Google"), vec![b'a'; 16377], true),
506 (200, Some("Google"), vec![b'a'; 16378], false),
507 (200, Some("Google"), vec![b'a'; 32768], false),
508 (200, Some("Google"), vec![], false),
509 (200, Some("Google"), vec![0xff], false),
510 (200, Some("Google"), b"token\n".to_vec(), false),
511 (200, Some("Impostor"), b"token".to_vec(), false),
512 (200, None, b"token".to_vec(), false),
513 (503, Some("Google"), b"token".to_vec(), false),
514 (302, Some("Google"), b"token".to_vec(), false),
515 ] {
516 server.reply(status, flavor, &body);
517 let provider = GoogleIdentity::new("fixture-audience").unwrap();
518 let result = provider
519 .token_from(&server.endpoint, tokio::time::Instant::now)
520 .await;
521 assert_eq!(
522 result.is_ok(),
523 accepted,
524 "status={status} flavor={flavor:?} length={}",
525 body.len()
526 );
527 if body.len() > 16377 {
528 assert_eq!(
529 result.unwrap_err().0,
530 "Google metadata identity response too large"
531 );
532 }
533 }
534 *server.response.lock().unwrap() = b"HTTP/1.1 200 OK\r\nMetadata-Flavor: Google\r\nContent-Length: 100\r\nConnection: close\r\n\r\ntruncated".to_vec();
535 let provider = GoogleIdentity::new("fixture-audience").unwrap();
536 assert!(
537 provider
538 .token_from(&server.endpoint, tokio::time::Instant::now)
539 .await
540 .is_err()
541 );
542 }
543
544 #[test]
545 fn bearer_framing_is_bounded_sensitive_and_redacted() {
546 for invalid in ["", "one two", "one\ttwo", "one\ntwo", "\u{7f}"] {
547 assert!(BearerToken::new(invalid).is_err());
548 }
549 assert!(BearerToken::new("x".repeat(16378)).is_err());
550 let token = BearerToken::new("x".repeat(16377)).unwrap();
551 assert_eq!(token.header().as_bytes().len(), 16384);
552 assert!(token.header().is_sensitive());
553 assert_eq!(format!("{token:?}"), "BearerToken([redacted])");
554 }
555}