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}
49
50impl Auth {
51 pub fn developer_token(token: impl Into<String>) -> Auth {
53 Auth {
54 source: Source::Developer(token.into()),
55 }
56 }
57
58 pub fn client_credentials(config: CcgConfig) -> Auth {
62 let (subject_type, subject_id) = match &config.user_id {
63 Some(user) => ("user", user.clone()),
64 None => ("enterprise", config.enterprise_id.clone()),
65 };
66 let form = vec![
67 ("grant_type".to_string(), "client_credentials".to_string()),
68 ("client_id".to_string(), config.client_id),
69 ("client_secret".to_string(), config.client_secret),
70 ("box_subject_type".to_string(), subject_type.to_string()),
71 ("box_subject_id".to_string(), subject_id),
72 ];
73 Auth {
74 source: Source::Form(FormSource {
75 http: auth_http_client(),
76 token_url: config.token_url.unwrap_or_else(default_token_url),
77 form,
78 cached: Mutex::new(Cached::empty()),
79 }),
80 }
81 }
82
83 pub fn jwt(config: JwtConfig) -> Result<Auth, Error> {
90 let signer = Signer::new(&config)?;
91 Ok(Auth {
92 source: Source::Jwt(Box::new(JwtSource {
93 http: auth_http_client(),
94 token_url: config.token_url.unwrap_or_else(default_token_url),
95 client_id: config.client_id,
96 client_secret: config.client_secret,
97 signer,
98 cached: Mutex::new(Cached::empty()),
99 })),
100 })
101 }
102
103 pub fn oauth(config: OAuthConfig, refresh_token: impl Into<String>) -> Auth {
110 Self::oauth_source(config, refresh_token.into(), None)
111 }
112
113 pub fn oauth_with_store(
117 config: OAuthConfig,
118 refresh_token: impl Into<String>,
119 store: Arc<dyn RefreshTokenStore>,
120 ) -> Auth {
121 Self::oauth_source(config, refresh_token.into(), Some(store))
122 }
123
124 fn oauth_source(
125 config: OAuthConfig,
126 refresh_token: String,
127 store: Option<Arc<dyn RefreshTokenStore>>,
128 ) -> Auth {
129 Auth {
130 source: Source::OAuth(OAuthSource {
131 http: auth_http_client(),
132 token_url: config.token_url.clone().unwrap_or_else(default_token_url),
133 client_id: config.client_id,
134 client_secret: config.client_secret,
135 store,
136 state: Mutex::new(OAuthState {
137 token: String::new(),
138 expiry: Instant::now(),
139 refresh_token,
140 refresh_token_persisted: true,
142 }),
143 }),
144 }
145 }
146
147 pub(crate) async fn access_token(&self) -> Result<String, Error> {
149 match &self.source {
150 Source::Developer(token) => Ok(token.clone()),
151 Source::Form(source) => source.access_token().await,
152 Source::OAuth(source) => source.access_token().await,
153 Source::Jwt(source) => source.access_token().await,
154 }
155 }
156
157 pub(crate) async fn force_refresh(&self, stale: &str) -> Result<String, Error> {
163 match &self.source {
164 Source::Developer(token) => Ok(token.clone()),
165 Source::Form(source) => source.force_refresh(stale).await,
166 Source::OAuth(source) => source.force_refresh(stale).await,
167 Source::Jwt(source) => source.force_refresh(stale).await,
168 }
169 }
170}
171
172#[async_trait::async_trait]
184pub trait RefreshTokenStore: Send + Sync {
185 async fn save(&self, refresh_token: &str) -> Result<(), Error>;
187}
188
189fn fresh_token(token: &str, expiry: Instant) -> Option<String> {
192 if !token.is_empty() && expiry.saturating_duration_since(Instant::now()) > REFRESH_MARGIN {
193 Some(token.to_string())
194 } else {
195 None
196 }
197}
198
199struct Cached {
201 token: String,
202 expiry: Instant,
203}
204
205impl Cached {
206 fn empty() -> Cached {
207 Cached {
208 token: String::new(),
209 expiry: Instant::now(),
210 }
211 }
212
213 fn fresh(&self) -> Option<String> {
214 fresh_token(&self.token, self.expiry)
215 }
216
217 fn store(&mut self, token: String, ttl: Duration) {
218 self.expiry = Instant::now() + ttl;
219 self.token = token;
220 }
221}
222
223struct FormSource {
225 http: reqwest::Client,
226 token_url: String,
227 form: Vec<(String, String)>,
228 cached: Mutex<Cached>,
229}
230
231impl FormSource {
232 async fn access_token(&self) -> Result<String, Error> {
233 let mut cached = self.cached.lock().await;
234 if let Some(token) = cached.fresh() {
235 return Ok(token);
236 }
237 self.refresh_locked(&mut cached).await
238 }
239
240 async fn force_refresh(&self, stale: &str) -> Result<String, Error> {
241 let mut cached = self.cached.lock().await;
242 if let Some(token) = cached.fresh() {
244 if token != stale {
245 return Ok(token);
246 }
247 }
248 self.refresh_locked(&mut cached).await
249 }
250
251 async fn refresh_locked(&self, cached: &mut Cached) -> Result<String, Error> {
252 let response = post_token_form(&self.http, &self.token_url, &self.form).await?;
253 cached.store(response.access_token.clone(), response.ttl());
254 Ok(response.access_token)
255 }
256}
257
258struct JwtSource {
262 http: reqwest::Client,
263 token_url: String,
264 client_id: String,
265 client_secret: String,
266 signer: Signer,
267 cached: Mutex<Cached>,
268}
269
270impl JwtSource {
271 async fn access_token(&self) -> Result<String, Error> {
272 let mut cached = self.cached.lock().await;
273 if let Some(token) = cached.fresh() {
274 return Ok(token);
275 }
276 self.refresh_locked(&mut cached).await
277 }
278
279 async fn force_refresh(&self, stale: &str) -> Result<String, Error> {
280 let mut cached = self.cached.lock().await;
281 if let Some(token) = cached.fresh() {
283 if token != stale {
284 return Ok(token);
285 }
286 }
287 self.refresh_locked(&mut cached).await
288 }
289
290 async fn refresh_locked(&self, cached: &mut Cached) -> Result<String, Error> {
291 let assertion = self.signer.assertion(&self.token_url)?;
292 let form = vec![
293 (
294 "grant_type".to_string(),
295 "urn:ietf:params:oauth:grant-type:jwt-bearer".to_string(),
296 ),
297 ("assertion".to_string(), assertion),
298 ("client_id".to_string(), self.client_id.clone()),
299 ("client_secret".to_string(), self.client_secret.clone()),
300 ];
301 let response = post_token_form(&self.http, &self.token_url, &form).await?;
302 cached.store(response.access_token.clone(), response.ttl());
303 Ok(response.access_token)
304 }
305}
306
307struct OAuthSource {
311 http: reqwest::Client,
312 token_url: String,
313 client_id: String,
314 client_secret: String,
315 store: Option<Arc<dyn RefreshTokenStore>>,
316 state: Mutex<OAuthState>,
317}
318
319struct OAuthState {
320 token: String,
321 expiry: Instant,
322 refresh_token: String,
323 refresh_token_persisted: bool,
327}
328
329impl OAuthSource {
330 async fn access_token(&self) -> Result<String, Error> {
331 let mut state = self.state.lock().await;
332 self.retry_persist(&mut state).await;
333 if let Some(token) = fresh_token(&state.token, state.expiry) {
334 return Ok(token);
335 }
336 self.refresh_locked(&mut state).await
337 }
338
339 async fn force_refresh(&self, stale: &str) -> Result<String, Error> {
340 let mut state = self.state.lock().await;
341 self.retry_persist(&mut state).await;
342 if let Some(token) = fresh_token(&state.token, state.expiry) {
344 if token != stale {
345 return Ok(token);
346 }
347 }
348 self.refresh_locked(&mut state).await
349 }
350
351 async fn retry_persist(&self, state: &mut OAuthState) {
356 if state.refresh_token_persisted {
357 return;
358 }
359 match &self.store {
360 Some(store) => {
361 if store.save(&state.refresh_token).await.is_ok() {
362 state.refresh_token_persisted = true;
363 }
364 }
365 None => state.refresh_token_persisted = true,
366 }
367 }
368
369 async fn refresh_locked(&self, state: &mut OAuthState) -> Result<String, Error> {
370 let form = vec![
371 ("grant_type".to_string(), "refresh_token".to_string()),
372 ("refresh_token".to_string(), state.refresh_token.clone()),
373 ("client_id".to_string(), self.client_id.clone()),
374 ("client_secret".to_string(), self.client_secret.clone()),
375 ];
376 let response = post_token_form(&self.http, &self.token_url, &form).await?;
377 state.token = response.access_token.clone();
378 state.expiry = Instant::now() + response.ttl();
379 if let Some(refresh) = &response.refresh_token {
380 state.refresh_token = refresh.clone();
381 state.refresh_token_persisted = self.store.is_none();
386 if let Some(store) = &self.store {
387 store.save(refresh).await?;
388 state.refresh_token_persisted = true;
389 }
390 }
391 Ok(response.access_token)
392 }
393}
394
395#[derive(Clone, Default)]
400pub struct CcgConfig {
401 pub client_id: String,
402 pub client_secret: String,
403 pub enterprise_id: String,
404 pub user_id: Option<String>,
406 pub token_url: Option<String>,
408}
409
410#[derive(Clone)]
415pub struct OAuthConfig {
416 pub client_id: String,
417 pub client_secret: String,
418 pub token_url: Option<String>,
420}
421
422impl OAuthConfig {
423 pub fn authorize_url(&self, redirect_uri: &str, state: &str) -> String {
426 let query = form_urlencode(&[
427 ("response_type", "code"),
428 ("client_id", &self.client_id),
429 ("redirect_uri", redirect_uri),
430 ("state", state),
431 ]);
432 format!("{AUTHORIZE_URL}?{query}")
433 }
434
435 pub async fn exchange_code(&self, code: &str, redirect_uri: &str) -> Result<Auth, Error> {
438 let http = auth_http_client();
439 let token_url = self.token_url.clone().unwrap_or_else(default_token_url);
440 let form = vec![
441 ("grant_type".to_string(), "authorization_code".to_string()),
442 ("code".to_string(), code.to_string()),
443 ("client_id".to_string(), self.client_id.clone()),
444 ("client_secret".to_string(), self.client_secret.clone()),
445 ("redirect_uri".to_string(), redirect_uri.to_string()),
446 ];
447 let response = post_token_form(&http, &token_url, &form).await?;
448 let refresh_token = response.refresh_token.clone().ok_or_else(|| {
449 Error::new("gantryruntime: authorization-code exchange returned no refresh_token")
450 })?;
451 let ttl = response.ttl();
452 let source = OAuthSource {
453 http,
454 token_url,
455 client_id: self.client_id.clone(),
456 client_secret: self.client_secret.clone(),
457 store: None,
458 state: Mutex::new(OAuthState {
459 token: response.access_token,
460 expiry: Instant::now() + ttl,
461 refresh_token,
462 refresh_token_persisted: true,
463 }),
464 };
465 Ok(Auth {
466 source: Source::OAuth(source),
467 })
468 }
469}
470
471struct TokenResponse {
473 access_token: String,
474 refresh_token: Option<String>,
475 expires_in: u64,
476}
477
478impl TokenResponse {
479 fn ttl(&self) -> Duration {
480 Duration::from_secs(self.expires_in)
481 }
482}
483
484async fn post_token_form(
487 http: &reqwest::Client,
488 token_url: &str,
489 form: &[(String, String)],
490) -> Result<TokenResponse, Error> {
491 let pairs: Vec<(&str, &str)> = form.iter().map(|(k, v)| (k.as_str(), v.as_str())).collect();
492 let body = form_urlencode(&pairs);
493 let response = http
494 .post(token_url)
495 .header("Content-Type", "application/x-www-form-urlencoded")
496 .header("Accept", "application/json")
497 .body(body)
498 .send()
499 .await?;
500 let status = response.status();
501 let bytes = response.bytes().await?;
502 if !status.is_success() {
503 let detail = String::from_utf8_lossy(&bytes);
504 return Err(Error::new(format!(
505 "gantryruntime: token endpoint returned {}: {}",
506 status.as_u16(),
507 detail.trim()
508 )));
509 }
510 let json: serde_json::Value = serde_json::from_slice(&bytes)?;
511 let access_token = json
512 .get("access_token")
513 .and_then(|v| v.as_str())
514 .map(|s| s.to_string())
515 .filter(|s| !s.is_empty())
516 .ok_or_else(|| Error::new("gantryruntime: token endpoint returned no access_token"))?;
517 Ok(TokenResponse {
518 access_token,
519 refresh_token: json
520 .get("refresh_token")
521 .and_then(|v| v.as_str())
522 .filter(|s| !s.is_empty())
523 .map(|s| s.to_string()),
524 expires_in: json.get("expires_in").and_then(|v| v.as_u64()).unwrap_or(0),
525 })
526}
527
528fn auth_http_client() -> reqwest::Client {
530 reqwest::Client::builder()
531 .timeout(Duration::from_secs(30))
532 .build()
533 .unwrap_or_default()
534}
535
536fn default_token_url() -> String {
537 DEFAULT_TOKEN_URL.to_string()
538}
539
540fn form_urlencode(pairs: &[(&str, &str)]) -> String {
542 let mut out = String::new();
543 for (name, value) in pairs {
544 if !out.is_empty() {
545 out.push('&');
546 }
547 percent_encode_into(&mut out, name);
548 out.push('=');
549 percent_encode_into(&mut out, value);
550 }
551 out
552}
553
554fn percent_encode_into(out: &mut String, value: &str) {
557 for byte in value.bytes() {
558 match byte {
559 b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'_' | b'.' | b'~' => {
560 out.push(byte as char)
561 }
562 b' ' => out.push('+'),
563 _ => out.push_str(&format!("%{byte:02X}")),
564 }
565 }
566}