use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{Duration, Instant};
use serde::Deserialize;
use tokio::sync::Mutex;
use crate::credentials::Credentials;
use crate::error::ApiError;
use crate::secret::Secret;
pub(crate) const TOKEN_PATH: &str = "/api/v2/oauth/token";
const JWT_ASSERTION: &str = "urn:ietf:params:oauth:client-assertion-type:jwt-bearer";
pub(crate) const REFRESH_SKEW: Duration = Duration::from_secs(60);
const ASSUMED_LIFETIME: Duration = Duration::from_secs(300);
#[derive(Debug)]
pub(crate) struct Bearer {
pub value: Secret,
pub generation: Option<u64>,
}
#[derive(Debug)]
struct Minted {
value: Secret,
expires_at: Instant,
generation: u64,
}
#[derive(Debug, Deserialize)]
struct TokenResponse {
access_token: String,
#[serde(default)]
expires_in: Option<u64>,
}
#[derive(Debug)]
pub(crate) struct Tokens {
credentials: Credentials,
endpoint: String,
http: reqwest::Client,
minted: Mutex<Option<Minted>>,
generations: AtomicU64,
}
impl Tokens {
pub(crate) fn new(credentials: Credentials, base_url: &str, http: reqwest::Client) -> Self {
Self {
credentials,
endpoint: format!("{base_url}{TOKEN_PATH}"),
http,
minted: Mutex::new(None),
generations: AtomicU64::new(0),
}
}
pub(crate) const fn can_refresh(&self) -> bool {
!matches!(self.credentials, Credentials::ApiKey(_))
}
pub(crate) async fn bearer(&self) -> Result<Bearer, ApiError> {
let Credentials::ApiKey(key) = &self.credentials else {
return self.minted_bearer().await;
};
Ok(Bearer {
value: key.clone(),
generation: None,
})
}
async fn minted_bearer(&self) -> Result<Bearer, ApiError> {
let mut held = self.minted.lock().await;
if let Some(current) = held.as_ref()
&& Instant::now() + REFRESH_SKEW < current.expires_at
{
return Ok(Bearer {
value: current.value.clone(),
generation: Some(current.generation),
});
}
let fresh = self.exchange().await?;
let generation = self.generations.fetch_add(1, Ordering::Relaxed);
let bearer = Bearer {
value: fresh.0.clone(),
generation: Some(generation),
};
*held = Some(Minted {
value: fresh.0,
expires_at: Instant::now() + fresh.1,
generation,
});
Ok(bearer)
}
pub(crate) async fn evict(&self, generation: u64) {
let mut held = self.minted.lock().await;
if held.as_ref().is_some_and(|m| m.generation == generation) {
*held = None;
}
}
async fn exchange(&self) -> Result<(Secret, Duration), ApiError> {
let mut form: Vec<(&str, String)> = vec![("grant_type", "client_credentials".to_owned())];
match &self.credentials {
Credentials::ApiKey(_) => {
return Err(ApiError::Token(
"an API access token is used directly and is never exchanged".to_owned(),
));
}
Credentials::OauthClient {
client_id,
client_secret,
scopes,
} => {
form.push(("client_id", client_id.clone()));
form.push(("client_secret", client_secret.expose().to_owned()));
push_scopes(&mut form, scopes);
}
Credentials::Federated {
client_id,
jwt_file,
scopes,
} => {
let jwt =
std::fs::read_to_string(jwt_file).map_err(|source| ApiError::JwtFile {
path: jwt_file.clone(),
source,
})?;
if let Some(client_id) = client_id {
form.push(("client_id", client_id.clone()));
}
form.push(("client_assertion_type", JWT_ASSERTION.to_owned()));
form.push(("client_assertion", jwt.trim().to_owned()));
push_scopes(&mut form, scopes);
}
}
tracing::debug!(
endpoint = %self.endpoint,
kind = self.credentials.kind(),
"exchanging the control-plane credential for a token"
);
let request = format!("POST {TOKEN_PATH}");
let response = self
.http
.post(&self.endpoint)
.form(&form)
.send()
.await
.map_err(|source| ApiError::Transport {
request: request.clone(),
source,
})?;
let status = response.status();
let body = response.text().await.unwrap_or_default();
if !status.is_success() {
return Err(ApiError::Status {
request,
status: status.as_u16(),
message: crate::error::describe(status, &body),
retry_after: None,
});
}
let parsed: TokenResponse = serde_json::from_str(&body).map_err(|source| {
ApiError::Token(format!("the token endpoint answered with {source}"))
})?;
if parsed.access_token.trim().is_empty() {
return Err(ApiError::Token(
"the token endpoint answered with an empty access token".to_owned(),
));
}
let lifetime = parsed
.expires_in
.map_or(ASSUMED_LIFETIME, Duration::from_secs);
Ok((Secret::new(parsed.access_token), lifetime))
}
}
fn push_scopes(form: &mut Vec<(&'static str, String)>, scopes: &[String]) {
if !scopes.is_empty() {
form.push(("scope", scopes.join(" ")));
}
}