1use std::sync::Arc;
13use std::time::{Duration, Instant};
14
15use tokio::sync::Mutex;
16
17use super::jwt::{JwtConfig, Signer};
18use super::Error;
19
20const DEFAULT_TOKEN_URL: &str = "https://api.box.com/oauth2/token";
22
23const AUTHORIZE_URL: &str = "https://account.box.com/api/oauth2/authorize";
25
26const REFRESH_MARGIN: Duration = Duration::from_secs(60);
29
30pub struct Auth {
34 source: Source,
35}
36
37enum Source {
39 Developer(String),
41 Form(FormSource),
43 OAuth(OAuthSource),
45 Jwt(Box<JwtSource>),
48 Custom(Arc<dyn TokenSource>),
50}
51
52#[async_trait::async_trait]
57pub trait TokenSource: Send + Sync {
58 async fn access_token(&self) -> Result<String, Error>;
61
62 async fn force_refresh(&self, stale: &str) -> Result<String, Error>;
67}
68
69impl Auth {
70 pub fn developer_token(token: impl Into<String>) -> Auth {
72 Auth {
73 source: Source::Developer(token.into()),
74 }
75 }
76
77 pub fn client_credentials(config: CcgConfig) -> Auth {
81 let (subject_type, subject_id) = match &config.user_id {
82 Some(user) => ("user", user.clone()),
83 None => ("enterprise", config.enterprise_id.clone()),
84 };
85 let form = vec![
86 ("grant_type".to_string(), "client_credentials".to_string()),
87 ("client_id".to_string(), config.client_id),
88 ("client_secret".to_string(), config.client_secret),
89 ("box_subject_type".to_string(), subject_type.to_string()),
90 ("box_subject_id".to_string(), subject_id),
91 ];
92 Auth {
93 source: Source::Form(FormSource {
94 http: auth_http_client(),
95 token_url: config.token_url.unwrap_or_else(default_token_url),
96 form,
97 cached: Mutex::new(Cached::empty()),
98 }),
99 }
100 }
101
102 pub fn jwt(config: JwtConfig) -> Result<Auth, Error> {
109 let signer = Signer::new(&config)?;
110 Ok(Auth {
111 source: Source::Jwt(Box::new(JwtSource {
112 http: auth_http_client(),
113 token_url: config.token_url.unwrap_or_else(default_token_url),
114 client_id: config.client_id,
115 client_secret: config.client_secret,
116 signer,
117 cached: Mutex::new(Cached::empty()),
118 })),
119 })
120 }
121
122 pub fn oauth(config: OAuthConfig, refresh_token: impl Into<String>) -> Auth {
129 Self::oauth_source(config, refresh_token.into(), None)
130 }
131
132 pub fn oauth_with_store(
136 config: OAuthConfig,
137 refresh_token: impl Into<String>,
138 store: Arc<dyn RefreshTokenStore>,
139 ) -> Auth {
140 Self::oauth_source(config, refresh_token.into(), Some(store))
141 }
142
143 pub fn custom(source: Arc<dyn TokenSource>) -> Auth {
148 Auth {
149 source: Source::Custom(source),
150 }
151 }
152
153 fn oauth_source(
154 config: OAuthConfig,
155 refresh_token: String,
156 store: Option<Arc<dyn RefreshTokenStore>>,
157 ) -> Auth {
158 Auth {
159 source: Source::OAuth(OAuthSource {
160 http: auth_http_client(),
161 token_url: config.token_url.clone().unwrap_or_else(default_token_url),
162 client_id: config.client_id,
163 client_secret: config.client_secret,
164 store,
165 state: Mutex::new(OAuthState {
166 token: String::new(),
167 expiry: Instant::now(),
168 refresh_token,
169 refresh_token_persisted: true,
171 }),
172 }),
173 }
174 }
175
176 pub(crate) async fn access_token(&self) -> Result<String, Error> {
178 match &self.source {
179 Source::Developer(token) => Ok(token.clone()),
180 Source::Form(source) => source.access_token().await,
181 Source::OAuth(source) => source.access_token().await,
182 Source::Jwt(source) => source.access_token().await,
183 Source::Custom(source) => source.access_token().await,
184 }
185 }
186
187 pub(crate) async fn force_refresh(&self, stale: &str) -> Result<String, Error> {
193 match &self.source {
194 Source::Developer(token) => Ok(token.clone()),
195 Source::Form(source) => source.force_refresh(stale).await,
196 Source::OAuth(source) => source.force_refresh(stale).await,
197 Source::Jwt(source) => source.force_refresh(stale).await,
198 Source::Custom(source) => source.force_refresh(stale).await,
199 }
200 }
201}
202
203#[async_trait::async_trait]
215pub trait RefreshTokenStore: Send + Sync {
216 async fn save(&self, refresh_token: &str) -> Result<(), Error>;
218}
219
220fn fresh_token(token: &str, expiry: Instant) -> Option<String> {
223 if !token.is_empty() && expiry.saturating_duration_since(Instant::now()) > REFRESH_MARGIN {
224 Some(token.to_string())
225 } else {
226 None
227 }
228}
229
230struct Cached {
232 token: String,
233 expiry: Instant,
234}
235
236impl Cached {
237 fn empty() -> Cached {
238 Cached {
239 token: String::new(),
240 expiry: Instant::now(),
241 }
242 }
243
244 fn fresh(&self) -> Option<String> {
245 fresh_token(&self.token, self.expiry)
246 }
247
248 fn store(&mut self, token: String, ttl: Duration) {
249 self.expiry = Instant::now() + ttl;
250 self.token = token;
251 }
252}
253
254struct FormSource {
256 http: reqwest::Client,
257 token_url: String,
258 form: Vec<(String, String)>,
259 cached: Mutex<Cached>,
260}
261
262impl FormSource {
263 async fn access_token(&self) -> Result<String, Error> {
264 let mut cached = self.cached.lock().await;
265 if let Some(token) = cached.fresh() {
266 return Ok(token);
267 }
268 self.refresh_locked(&mut cached).await
269 }
270
271 async fn force_refresh(&self, stale: &str) -> Result<String, Error> {
272 let mut cached = self.cached.lock().await;
273 if let Some(token) = cached.fresh() {
275 if token != stale {
276 return Ok(token);
277 }
278 }
279 self.refresh_locked(&mut cached).await
280 }
281
282 async fn refresh_locked(&self, cached: &mut Cached) -> Result<String, Error> {
283 let response = post_token_form(&self.http, &self.token_url, &self.form).await?;
284 cached.store(response.access_token.clone(), response.ttl());
285 Ok(response.access_token)
286 }
287}
288
289struct JwtSource {
293 http: reqwest::Client,
294 token_url: String,
295 client_id: String,
296 client_secret: String,
297 signer: Signer,
298 cached: Mutex<Cached>,
299}
300
301impl JwtSource {
302 async fn access_token(&self) -> Result<String, Error> {
303 let mut cached = self.cached.lock().await;
304 if let Some(token) = cached.fresh() {
305 return Ok(token);
306 }
307 self.refresh_locked(&mut cached).await
308 }
309
310 async fn force_refresh(&self, stale: &str) -> Result<String, Error> {
311 let mut cached = self.cached.lock().await;
312 if let Some(token) = cached.fresh() {
314 if token != stale {
315 return Ok(token);
316 }
317 }
318 self.refresh_locked(&mut cached).await
319 }
320
321 async fn refresh_locked(&self, cached: &mut Cached) -> Result<String, Error> {
322 let assertion = self.signer.assertion(&self.token_url)?;
323 let form = vec![
324 (
325 "grant_type".to_string(),
326 "urn:ietf:params:oauth:grant-type:jwt-bearer".to_string(),
327 ),
328 ("assertion".to_string(), assertion),
329 ("client_id".to_string(), self.client_id.clone()),
330 ("client_secret".to_string(), self.client_secret.clone()),
331 ];
332 let response = post_token_form(&self.http, &self.token_url, &form).await?;
333 cached.store(response.access_token.clone(), response.ttl());
334 Ok(response.access_token)
335 }
336}
337
338struct OAuthSource {
342 http: reqwest::Client,
343 token_url: String,
344 client_id: String,
345 client_secret: String,
346 store: Option<Arc<dyn RefreshTokenStore>>,
347 state: Mutex<OAuthState>,
348}
349
350struct OAuthState {
351 token: String,
352 expiry: Instant,
353 refresh_token: String,
354 refresh_token_persisted: bool,
358}
359
360impl OAuthSource {
361 async fn access_token(&self) -> Result<String, Error> {
362 let mut state = self.state.lock().await;
363 self.retry_persist(&mut state).await;
364 if let Some(token) = fresh_token(&state.token, state.expiry) {
365 return Ok(token);
366 }
367 self.refresh_locked(&mut state).await
368 }
369
370 async fn force_refresh(&self, stale: &str) -> Result<String, Error> {
371 let mut state = self.state.lock().await;
372 self.retry_persist(&mut state).await;
373 if let Some(token) = fresh_token(&state.token, state.expiry) {
375 if token != stale {
376 return Ok(token);
377 }
378 }
379 self.refresh_locked(&mut state).await
380 }
381
382 async fn retry_persist(&self, state: &mut OAuthState) {
387 if state.refresh_token_persisted {
388 return;
389 }
390 match &self.store {
391 Some(store) => {
392 if store.save(&state.refresh_token).await.is_ok() {
393 state.refresh_token_persisted = true;
394 }
395 }
396 None => state.refresh_token_persisted = true,
397 }
398 }
399
400 async fn refresh_locked(&self, state: &mut OAuthState) -> Result<String, Error> {
401 let form = vec![
402 ("grant_type".to_string(), "refresh_token".to_string()),
403 ("refresh_token".to_string(), state.refresh_token.clone()),
404 ("client_id".to_string(), self.client_id.clone()),
405 ("client_secret".to_string(), self.client_secret.clone()),
406 ];
407 let response = post_token_form(&self.http, &self.token_url, &form).await?;
408 state.token = response.access_token.clone();
409 state.expiry = Instant::now() + response.ttl();
410 if let Some(refresh) = &response.refresh_token {
411 state.refresh_token = refresh.clone();
412 state.refresh_token_persisted = self.store.is_none();
417 if let Some(store) = &self.store {
418 store.save(refresh).await?;
419 state.refresh_token_persisted = true;
420 }
421 }
422 Ok(response.access_token)
423 }
424}
425
426#[derive(Clone, Default)]
431pub struct CcgConfig {
432 pub client_id: String,
433 pub client_secret: String,
434 pub enterprise_id: String,
435 pub user_id: Option<String>,
437 pub token_url: Option<String>,
439}
440
441#[derive(Clone)]
446pub struct OAuthConfig {
447 pub client_id: String,
448 pub client_secret: String,
449 pub token_url: Option<String>,
451}
452
453impl OAuthConfig {
454 pub fn authorize_url(&self, redirect_uri: &str, state: &str) -> String {
457 let query = form_urlencode(&[
458 ("response_type", "code"),
459 ("client_id", &self.client_id),
460 ("redirect_uri", redirect_uri),
461 ("state", state),
462 ]);
463 format!("{AUTHORIZE_URL}?{query}")
464 }
465
466 pub async fn exchange_code(&self, code: &str, redirect_uri: &str) -> Result<Auth, Error> {
469 let http = auth_http_client();
470 let token_url = self.token_url.clone().unwrap_or_else(default_token_url);
471 let form = vec![
472 ("grant_type".to_string(), "authorization_code".to_string()),
473 ("code".to_string(), code.to_string()),
474 ("client_id".to_string(), self.client_id.clone()),
475 ("client_secret".to_string(), self.client_secret.clone()),
476 ("redirect_uri".to_string(), redirect_uri.to_string()),
477 ];
478 let response = post_token_form(&http, &token_url, &form).await?;
479 let refresh_token = response.refresh_token.clone().ok_or_else(|| {
480 Error::new("gantryruntime: authorization-code exchange returned no refresh_token")
481 })?;
482 let ttl = response.ttl();
483 let source = OAuthSource {
484 http,
485 token_url,
486 client_id: self.client_id.clone(),
487 client_secret: self.client_secret.clone(),
488 store: None,
489 state: Mutex::new(OAuthState {
490 token: response.access_token,
491 expiry: Instant::now() + ttl,
492 refresh_token,
493 refresh_token_persisted: true,
494 }),
495 };
496 Ok(Auth {
497 source: Source::OAuth(source),
498 })
499 }
500}
501
502struct TokenResponse {
504 access_token: String,
505 refresh_token: Option<String>,
506 expires_in: u64,
507}
508
509impl TokenResponse {
510 fn ttl(&self) -> Duration {
511 Duration::from_secs(self.expires_in)
512 }
513}
514
515async fn post_token_form(
518 http: &reqwest::Client,
519 token_url: &str,
520 form: &[(String, String)],
521) -> Result<TokenResponse, Error> {
522 let pairs: Vec<(&str, &str)> = form.iter().map(|(k, v)| (k.as_str(), v.as_str())).collect();
523 let body = form_urlencode(&pairs);
524 let response = http
525 .post(token_url)
526 .header("Content-Type", "application/x-www-form-urlencoded")
527 .header("Accept", "application/json")
528 .body(body)
529 .send()
530 .await?;
531 let status = response.status();
532 let bytes = response.bytes().await?;
533 if !status.is_success() {
534 let detail = String::from_utf8_lossy(&bytes);
535 return Err(Error::new(format!(
536 "gantryruntime: token endpoint returned {}: {}",
537 status.as_u16(),
538 detail.trim()
539 )));
540 }
541 let json: serde_json::Value = serde_json::from_slice(&bytes)?;
542 let access_token = json
543 .get("access_token")
544 .and_then(|v| v.as_str())
545 .map(|s| s.to_string())
546 .filter(|s| !s.is_empty())
547 .ok_or_else(|| Error::new("gantryruntime: token endpoint returned no access_token"))?;
548 Ok(TokenResponse {
549 access_token,
550 refresh_token: json
551 .get("refresh_token")
552 .and_then(|v| v.as_str())
553 .filter(|s| !s.is_empty())
554 .map(|s| s.to_string()),
555 expires_in: json.get("expires_in").and_then(|v| v.as_u64()).unwrap_or(0),
556 })
557}
558
559fn auth_http_client() -> reqwest::Client {
561 reqwest::Client::builder()
562 .timeout(Duration::from_secs(30))
563 .build()
564 .unwrap_or_default()
565}
566
567fn default_token_url() -> String {
568 DEFAULT_TOKEN_URL.to_string()
569}
570
571fn form_urlencode(pairs: &[(&str, &str)]) -> String {
573 let mut out = String::new();
574 for (name, value) in pairs {
575 if !out.is_empty() {
576 out.push('&');
577 }
578 percent_encode_into(&mut out, name);
579 out.push('=');
580 percent_encode_into(&mut out, value);
581 }
582 out
583}
584
585fn percent_encode_into(out: &mut String, value: &str) {
588 for byte in value.bytes() {
589 match byte {
590 b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'_' | b'.' | b'~' => {
591 out.push(byte as char)
592 }
593 b' ' => out.push('+'),
594 _ => out.push_str(&format!("%{byte:02X}")),
595 }
596 }
597}