use std::time::Duration;
use base64::Engine as _;
use reqwest::Method;
use serde::Deserialize;
use sha2::Digest as _;
use crate::client::Client;
use crate::error::Error;
use crate::http::{RequestSpec, Transport};
use crate::resources::auth::User;
const VERIFIER_CHARS: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-._~";
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct PkcePair {
pub verifier: String,
pub challenge: String,
}
pub fn generate_pkce_pair() -> PkcePair {
let verifier = random_string(64);
let digest = sha2::Sha256::digest(verifier.as_bytes());
let challenge = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(digest);
PkcePair {
verifier,
challenge,
}
}
pub fn generate_state() -> String {
random_string(32)
}
fn random_string(len: usize) -> String {
let mut rng = rand::rng();
(0..len)
.map(|_| {
let idx = rand::Rng::random_range(&mut rng, 0..VERIFIER_CHARS.len());
char::from(VERIFIER_CHARS.get(idx).copied().unwrap_or(b'a'))
})
.collect()
}
#[derive(Clone, Deserialize)]
#[non_exhaustive]
pub struct TokenPair {
pub access_token: String,
pub refresh_token: String,
}
impl TokenPair {
pub fn new(access_token: impl Into<String>, refresh_token: impl Into<String>) -> Self {
TokenPair {
access_token: access_token.into(),
refresh_token: refresh_token.into(),
}
}
}
impl std::fmt::Debug for TokenPair {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TokenPair")
.field("access_token", &"[redacted]")
.field("refresh_token", &"[redacted]")
.finish()
}
}
#[derive(Clone, Deserialize)]
#[non_exhaustive]
pub struct DeviceTokens {
pub access_token: String,
pub refresh_token: String,
pub user: User,
}
impl DeviceTokens {
pub fn tokens(&self) -> TokenPair {
TokenPair::new(self.access_token.clone(), self.refresh_token.clone())
}
}
impl std::fmt::Debug for DeviceTokens {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("DeviceTokens")
.field("access_token", &"[redacted]")
.field("refresh_token", &"[redacted]")
.field("user", &self.user)
.finish()
}
}
pub struct OAuth {
pub(crate) client: Client,
}
impl OAuth {
pub fn authorization_url(
&self,
app_id: impl Into<String>,
state: impl Into<String>,
code_challenge: impl Into<String>,
) -> AuthorizationUrlBuilder {
AuthorizationUrlBuilder {
base_url: self.client.transport.base_url.clone(),
app_id: app_id.into(),
state: state.into(),
code_challenge: code_challenge.into(),
redirect_uri: None,
}
}
pub async fn exchange_code(
&self,
code: impl Into<String>,
code_verifier: impl Into<String>,
) -> Result<DeviceTokens, Error> {
let spec =
RequestSpec::new(Method::POST, "/auth/device/token").json(&serde_json::json!({
"code": code.into(),
"code_verifier": code_verifier.into(),
}))?;
self.client.transport.execute(spec).await
}
pub async fn refresh_tokens(&self, refresh_token: &str) -> Result<TokenPair, Error> {
refresh_call(&self.client.transport, refresh_token).await
}
}
#[must_use = "the URL is only produced by .build()"]
pub struct AuthorizationUrlBuilder {
base_url: String,
app_id: String,
state: String,
code_challenge: String,
redirect_uri: Option<String>,
}
impl AuthorizationUrlBuilder {
pub fn redirect_uri(mut self, uri: impl Into<String>) -> Self {
self.redirect_uri = Some(uri.into());
self
}
pub fn build(self) -> Result<String, Error> {
let mut url = reqwest::Url::parse(&format!("{}/auth/device/login", self.base_url))
.map_err(|e| Error::Config(format!("invalid base URL: {e}")))?;
{
let mut query = url.query_pairs_mut();
query.append_pair("app_id", &self.app_id);
if let Some(redirect_uri) = &self.redirect_uri {
query.append_pair("redirect_uri", redirect_uri);
}
query.append_pair("state", &self.state);
query.append_pair("code_challenge", &self.code_challenge);
query.append_pair("code_challenge_method", "S256");
}
Ok(url.into())
}
}
async fn refresh_call(transport: &Transport, refresh_token: &str) -> Result<TokenPair, Error> {
let spec = RequestSpec::new(Method::POST, "/auth/device/refresh")
.json(&serde_json::json!({ "refresh_token": refresh_token }))?;
match transport.execute_unauthenticated::<TokenPair>(spec).await {
Ok(pair) => Ok(pair),
Err(Error::Api(e)) if e.status == 400 || e.status == 401 => Err(Error::SessionExpired),
Err(err) => Err(err),
}
}
type OnRefresh = Box<dyn Fn(&TokenPair) + Send + Sync>;
struct SessionState {
tokens: TokenPair,
generation: u64,
expires_at: Option<i64>,
}
pub struct Session {
state: tokio::sync::Mutex<SessionState>,
on_refresh: Option<OnRefresh>,
expiry_skew: Duration,
}
impl std::fmt::Debug for Session {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Session").finish_non_exhaustive()
}
}
impl Session {
pub fn new(tokens: TokenPair) -> Session {
let expires_at = decode_exp(&tokens.access_token);
Session {
state: tokio::sync::Mutex::new(SessionState {
tokens,
generation: 0,
expires_at,
}),
on_refresh: None,
expiry_skew: Duration::from_secs(30),
}
}
pub fn on_refresh(mut self, hook: impl Fn(&TokenPair) + Send + Sync + 'static) -> Session {
self.on_refresh = Some(Box::new(hook));
self
}
pub fn expiry_skew(mut self, skew: Duration) -> Session {
self.expiry_skew = skew;
self
}
pub async fn invalidate(&self) {
let mut state = self.state.lock().await;
state.expires_at = Some(0);
}
pub(crate) async fn fresh_token(&self, transport: &Transport) -> Result<(String, u64), Error> {
let mut state = self.state.lock().await;
let stale = state
.expires_at
.is_some_and(|exp| utc_now_secs() + self.expiry_skew.as_secs() as i64 >= exp);
let rotated = if stale {
Some(self.rotate_locked(&mut state, transport).await?)
} else {
None
};
let result = (state.tokens.access_token.clone(), state.generation);
drop(state);
if let Some(pair) = rotated {
self.notify(&pair);
}
Ok(result)
}
pub(crate) async fn refresh_stale(
&self,
transport: &Transport,
seen_generation: u64,
) -> Result<(), Error> {
let mut state = self.state.lock().await;
if state.generation != seen_generation {
return Ok(());
}
let pair = self.rotate_locked(&mut state, transport).await?;
drop(state);
self.notify(&pair);
Ok(())
}
async fn rotate_locked(
&self,
state: &mut SessionState,
transport: &Transport,
) -> Result<TokenPair, Error> {
let pair = refresh_call(transport, &state.tokens.refresh_token).await?;
state.expires_at = decode_exp(&pair.access_token);
state.tokens = pair.clone();
state.generation += 1;
Ok(pair)
}
fn notify(&self, pair: &TokenPair) {
if let Some(hook) = &self.on_refresh {
hook(pair);
}
}
}
fn utc_now_secs() -> i64 {
chrono::Utc::now().timestamp()
}
fn decode_exp(jwt: &str) -> Option<i64> {
let payload = jwt.split('.').nth(1)?;
let bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(payload)
.ok()?;
let value: serde_json::Value = serde_json::from_slice(&bytes).ok()?;
value.get("exp")?.as_i64()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn pkce_pair_matches_rfc_7636_s256() {
let pair = generate_pkce_pair();
assert_eq!(pair.verifier.len(), 64);
assert!(pair.verifier.bytes().all(|b| VERIFIER_CHARS.contains(&b)));
let digest = sha2::Sha256::digest(pair.verifier.as_bytes());
let expected = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(digest);
assert_eq!(pair.challenge, expected);
assert!(!pair.challenge.contains('='));
}
#[test]
fn exp_claim_decodes() {
let payload =
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(br#"{"exp":1234567890}"#);
let jwt = format!("h.{payload}.s");
assert_eq!(decode_exp(&jwt), Some(1234567890));
assert_eq!(decode_exp("not-a-jwt"), None);
}
}