use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::Mutex;
use super::jwt::{JwtConfig, Signer};
use super::Error;
const DEFAULT_TOKEN_URL: &str = "https://api.box.com/oauth2/token";
const AUTHORIZE_URL: &str = "https://account.box.com/api/oauth2/authorize";
const REFRESH_MARGIN: Duration = Duration::from_secs(60);
pub struct Auth {
source: Source,
}
enum Source {
Developer(String),
Form(FormSource),
OAuth(OAuthSource),
Jwt(Box<JwtSource>),
}
impl Auth {
pub fn developer_token(token: impl Into<String>) -> Auth {
Auth {
source: Source::Developer(token.into()),
}
}
pub fn client_credentials(config: CcgConfig) -> Auth {
let (subject_type, subject_id) = match &config.user_id {
Some(user) => ("user", user.clone()),
None => ("enterprise", config.enterprise_id.clone()),
};
let form = vec![
("grant_type".to_string(), "client_credentials".to_string()),
("client_id".to_string(), config.client_id),
("client_secret".to_string(), config.client_secret),
("box_subject_type".to_string(), subject_type.to_string()),
("box_subject_id".to_string(), subject_id),
];
Auth {
source: Source::Form(FormSource {
http: auth_http_client(),
token_url: config.token_url.unwrap_or_else(default_token_url),
form,
cached: Mutex::new(Cached::empty()),
}),
}
}
pub fn jwt(config: JwtConfig) -> Result<Auth, Error> {
let signer = Signer::new(&config)?;
Ok(Auth {
source: Source::Jwt(Box::new(JwtSource {
http: auth_http_client(),
token_url: config.token_url.unwrap_or_else(default_token_url),
client_id: config.client_id,
client_secret: config.client_secret,
signer,
cached: Mutex::new(Cached::empty()),
})),
})
}
pub fn oauth(config: OAuthConfig, refresh_token: impl Into<String>) -> Auth {
Self::oauth_source(config, refresh_token.into(), None)
}
pub fn oauth_with_store(
config: OAuthConfig,
refresh_token: impl Into<String>,
store: Arc<dyn RefreshTokenStore>,
) -> Auth {
Self::oauth_source(config, refresh_token.into(), Some(store))
}
fn oauth_source(
config: OAuthConfig,
refresh_token: String,
store: Option<Arc<dyn RefreshTokenStore>>,
) -> Auth {
Auth {
source: Source::OAuth(OAuthSource {
http: auth_http_client(),
token_url: config.token_url.clone().unwrap_or_else(default_token_url),
client_id: config.client_id,
client_secret: config.client_secret,
store,
state: Mutex::new(OAuthState {
token: String::new(),
expiry: Instant::now(),
refresh_token,
refresh_token_persisted: true,
}),
}),
}
}
pub(crate) async fn access_token(&self) -> Result<String, Error> {
match &self.source {
Source::Developer(token) => Ok(token.clone()),
Source::Form(source) => source.access_token().await,
Source::OAuth(source) => source.access_token().await,
Source::Jwt(source) => source.access_token().await,
}
}
pub(crate) async fn force_refresh(&self, stale: &str) -> Result<String, Error> {
match &self.source {
Source::Developer(token) => Ok(token.clone()),
Source::Form(source) => source.force_refresh(stale).await,
Source::OAuth(source) => source.force_refresh(stale).await,
Source::Jwt(source) => source.force_refresh(stale).await,
}
}
}
#[async_trait::async_trait]
pub trait RefreshTokenStore: Send + Sync {
async fn save(&self, refresh_token: &str) -> Result<(), Error>;
}
fn fresh_token(token: &str, expiry: Instant) -> Option<String> {
if !token.is_empty() && expiry.saturating_duration_since(Instant::now()) > REFRESH_MARGIN {
Some(token.to_string())
} else {
None
}
}
struct Cached {
token: String,
expiry: Instant,
}
impl Cached {
fn empty() -> Cached {
Cached {
token: String::new(),
expiry: Instant::now(),
}
}
fn fresh(&self) -> Option<String> {
fresh_token(&self.token, self.expiry)
}
fn store(&mut self, token: String, ttl: Duration) {
self.expiry = Instant::now() + ttl;
self.token = token;
}
}
struct FormSource {
http: reqwest::Client,
token_url: String,
form: Vec<(String, String)>,
cached: Mutex<Cached>,
}
impl FormSource {
async fn access_token(&self) -> Result<String, Error> {
let mut cached = self.cached.lock().await;
if let Some(token) = cached.fresh() {
return Ok(token);
}
self.refresh_locked(&mut cached).await
}
async fn force_refresh(&self, stale: &str) -> Result<String, Error> {
let mut cached = self.cached.lock().await;
if let Some(token) = cached.fresh() {
if token != stale {
return Ok(token);
}
}
self.refresh_locked(&mut cached).await
}
async fn refresh_locked(&self, cached: &mut Cached) -> Result<String, Error> {
let response = post_token_form(&self.http, &self.token_url, &self.form).await?;
cached.store(response.access_token.clone(), response.ttl());
Ok(response.access_token)
}
}
struct JwtSource {
http: reqwest::Client,
token_url: String,
client_id: String,
client_secret: String,
signer: Signer,
cached: Mutex<Cached>,
}
impl JwtSource {
async fn access_token(&self) -> Result<String, Error> {
let mut cached = self.cached.lock().await;
if let Some(token) = cached.fresh() {
return Ok(token);
}
self.refresh_locked(&mut cached).await
}
async fn force_refresh(&self, stale: &str) -> Result<String, Error> {
let mut cached = self.cached.lock().await;
if let Some(token) = cached.fresh() {
if token != stale {
return Ok(token);
}
}
self.refresh_locked(&mut cached).await
}
async fn refresh_locked(&self, cached: &mut Cached) -> Result<String, Error> {
let assertion = self.signer.assertion(&self.token_url)?;
let form = vec![
(
"grant_type".to_string(),
"urn:ietf:params:oauth:grant-type:jwt-bearer".to_string(),
),
("assertion".to_string(), assertion),
("client_id".to_string(), self.client_id.clone()),
("client_secret".to_string(), self.client_secret.clone()),
];
let response = post_token_form(&self.http, &self.token_url, &form).await?;
cached.store(response.access_token.clone(), response.ttl());
Ok(response.access_token)
}
}
struct OAuthSource {
http: reqwest::Client,
token_url: String,
client_id: String,
client_secret: String,
store: Option<Arc<dyn RefreshTokenStore>>,
state: Mutex<OAuthState>,
}
struct OAuthState {
token: String,
expiry: Instant,
refresh_token: String,
refresh_token_persisted: bool,
}
impl OAuthSource {
async fn access_token(&self) -> Result<String, Error> {
let mut state = self.state.lock().await;
self.retry_persist(&mut state).await;
if let Some(token) = fresh_token(&state.token, state.expiry) {
return Ok(token);
}
self.refresh_locked(&mut state).await
}
async fn force_refresh(&self, stale: &str) -> Result<String, Error> {
let mut state = self.state.lock().await;
self.retry_persist(&mut state).await;
if let Some(token) = fresh_token(&state.token, state.expiry) {
if token != stale {
return Ok(token);
}
}
self.refresh_locked(&mut state).await
}
async fn retry_persist(&self, state: &mut OAuthState) {
if state.refresh_token_persisted {
return;
}
match &self.store {
Some(store) => {
if store.save(&state.refresh_token).await.is_ok() {
state.refresh_token_persisted = true;
}
}
None => state.refresh_token_persisted = true,
}
}
async fn refresh_locked(&self, state: &mut OAuthState) -> Result<String, Error> {
let form = vec![
("grant_type".to_string(), "refresh_token".to_string()),
("refresh_token".to_string(), state.refresh_token.clone()),
("client_id".to_string(), self.client_id.clone()),
("client_secret".to_string(), self.client_secret.clone()),
];
let response = post_token_form(&self.http, &self.token_url, &form).await?;
state.token = response.access_token.clone();
state.expiry = Instant::now() + response.ttl();
if let Some(refresh) = &response.refresh_token {
state.refresh_token = refresh.clone();
state.refresh_token_persisted = self.store.is_none();
if let Some(store) = &self.store {
store.save(refresh).await?;
state.refresh_token_persisted = true;
}
}
Ok(response.access_token)
}
}
pub struct CcgConfig {
pub client_id: String,
pub client_secret: String,
pub enterprise_id: String,
pub user_id: Option<String>,
pub token_url: Option<String>,
}
#[derive(Clone)]
pub struct OAuthConfig {
pub client_id: String,
pub client_secret: String,
pub token_url: Option<String>,
}
impl OAuthConfig {
pub fn authorize_url(&self, redirect_uri: &str, state: &str) -> String {
let query = form_urlencode(&[
("response_type", "code"),
("client_id", &self.client_id),
("redirect_uri", redirect_uri),
("state", state),
]);
format!("{AUTHORIZE_URL}?{query}")
}
pub async fn exchange_code(&self, code: &str, redirect_uri: &str) -> Result<Auth, Error> {
let http = auth_http_client();
let token_url = self.token_url.clone().unwrap_or_else(default_token_url);
let form = vec![
("grant_type".to_string(), "authorization_code".to_string()),
("code".to_string(), code.to_string()),
("client_id".to_string(), self.client_id.clone()),
("client_secret".to_string(), self.client_secret.clone()),
("redirect_uri".to_string(), redirect_uri.to_string()),
];
let response = post_token_form(&http, &token_url, &form).await?;
let refresh_token = response.refresh_token.clone().ok_or_else(|| {
Error::new("gantryruntime: authorization-code exchange returned no refresh_token")
})?;
let ttl = response.ttl();
let source = OAuthSource {
http,
token_url,
client_id: self.client_id.clone(),
client_secret: self.client_secret.clone(),
store: None,
state: Mutex::new(OAuthState {
token: response.access_token,
expiry: Instant::now() + ttl,
refresh_token,
refresh_token_persisted: true,
}),
};
Ok(Auth {
source: Source::OAuth(source),
})
}
}
struct TokenResponse {
access_token: String,
refresh_token: Option<String>,
expires_in: u64,
}
impl TokenResponse {
fn ttl(&self) -> Duration {
Duration::from_secs(self.expires_in)
}
}
async fn post_token_form(
http: &reqwest::Client,
token_url: &str,
form: &[(String, String)],
) -> Result<TokenResponse, Error> {
let pairs: Vec<(&str, &str)> = form.iter().map(|(k, v)| (k.as_str(), v.as_str())).collect();
let body = form_urlencode(&pairs);
let response = http
.post(token_url)
.header("Content-Type", "application/x-www-form-urlencoded")
.header("Accept", "application/json")
.body(body)
.send()
.await?;
let status = response.status();
let bytes = response.bytes().await?;
if !status.is_success() {
let detail = String::from_utf8_lossy(&bytes);
return Err(Error::new(format!(
"gantryruntime: token endpoint returned {}: {}",
status.as_u16(),
detail.trim()
)));
}
let json: serde_json::Value = serde_json::from_slice(&bytes)?;
let access_token = json
.get("access_token")
.and_then(|v| v.as_str())
.map(|s| s.to_string())
.filter(|s| !s.is_empty())
.ok_or_else(|| Error::new("gantryruntime: token endpoint returned no access_token"))?;
Ok(TokenResponse {
access_token,
refresh_token: json
.get("refresh_token")
.and_then(|v| v.as_str())
.filter(|s| !s.is_empty())
.map(|s| s.to_string()),
expires_in: json.get("expires_in").and_then(|v| v.as_u64()).unwrap_or(0),
})
}
fn auth_http_client() -> reqwest::Client {
reqwest::Client::builder()
.timeout(Duration::from_secs(30))
.build()
.unwrap_or_default()
}
fn default_token_url() -> String {
DEFAULT_TOKEN_URL.to_string()
}
fn form_urlencode(pairs: &[(&str, &str)]) -> String {
let mut out = String::new();
for (name, value) in pairs {
if !out.is_empty() {
out.push('&');
}
percent_encode_into(&mut out, name);
out.push('=');
percent_encode_into(&mut out, value);
}
out
}
fn percent_encode_into(out: &mut String, value: &str) {
for byte in value.bytes() {
match byte {
b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'_' | b'.' | b'~' => {
out.push(byte as char)
}
b' ' => out.push('+'),
_ => out.push_str(&format!("%{byte:02X}")),
}
}
}