use std::time::{Duration, SystemTime, UNIX_EPOCH};
use base64::Engine;
use base64::engine::general_purpose::{STANDARD as BASE64_STANDARD, URL_SAFE_NO_PAD};
use sha2::{Digest, Sha256};
use url::Url;
use super::Auth;
use super::callback;
use super::pending;
use crate::error::{Error, Result};
use crate::store::OAuth2Token;
use tokio_util::sync::CancellationToken;
#[must_use]
pub fn get_oauth2_scopes() -> Vec<&'static str> {
vec![
"tweet.read",
"users.read",
"bookmark.read",
"follows.read",
"list.read",
"block.read",
"mute.read",
"like.read",
"users.email",
"dm.read",
"broadcast.read",
"tweet.write",
"tweet.moderate.write",
"follows.write",
"bookmark.write",
"block.write",
"mute.write",
"like.write",
"list.write",
"media.write",
"dm.write",
"broadcast.write",
"offline.access",
"space.read",
]
}
#[must_use]
pub fn generate_code_verifier_and_challenge() -> (String, String) {
let b: [u8; 32] = rand::random();
let verifier = URL_SAFE_NO_PAD.encode(b);
let mut hasher = Sha256::new();
hasher.update(verifier.as_bytes());
let challenge = URL_SAFE_NO_PAD.encode(hasher.finalize());
(verifier, challenge)
}
pub(crate) fn build_auth_url(auth: &Auth, state: &str, challenge: &str) -> Result<String> {
let scopes = get_oauth2_scopes().join(" ");
let mut auth_url =
Url::parse(auth.auth_url()).map_err(|e| Error::auth_with_cause("InvalidURL", &e))?;
auth_url
.query_pairs_mut()
.append_pair("response_type", "code")
.append_pair("client_id", auth.client_id())
.append_pair("redirect_uri", auth.redirect_uri())
.append_pair("scope", &scopes)
.append_pair("state", state)
.append_pair("code_challenge", challenge)
.append_pair("code_challenge_method", "S256");
Ok(auth_url.to_string())
}
pub(crate) async fn exchange_code_for_token(
auth: &mut Auth,
http: &reqwest::Client,
code: &str,
verifier: &str,
username: &str,
) -> Result<String> {
let token_resp = http
.post(auth.token_url())
.timeout(Duration::from_secs(auth.http_timeout_secs()))
.form(&[
("grant_type", "authorization_code"),
("code", code),
("redirect_uri", auth.redirect_uri()),
("client_id", auth.client_id()),
("code_verifier", verifier),
])
.basic_auth(auth.client_id(), Some(auth.client_secret()))
.send()
.await
.map_err(|e| Error::auth_with_cause("TokenExchangeError", &e))?;
let status = token_resp.status();
let token_data: serde_json::Value = token_resp
.json()
.await
.map_err(|e| Error::auth_with_cause("TokenExchangeError", &e))?;
if !status.is_success() {
let api_error = token_data["error"].as_str().unwrap_or("unknown");
let api_desc = token_data["error_description"].as_str().unwrap_or("");
return Err(Error::auth(format!(
"TokenExchangeError: HTTP {status} — {api_error}: {api_desc}"
)));
}
let access_token = token_data["access_token"]
.as_str()
.ok_or_else(|| Error::auth("TokenExchangeError: no access_token in response"))?
.to_string();
let refresh_token = token_data["refresh_token"]
.as_str()
.unwrap_or("")
.to_string();
let expires_in = token_data["expires_in"].as_u64().unwrap_or(7200);
let expiration_time = epoch_secs(SystemTime::now()) + expires_in;
let app_name = auth.app_name().to_string();
if username.is_empty() {
match auth.fetch_username(http, &access_token).await {
Ok(discovered) => {
auth.token_store.save_oauth2_token_for_app(
&app_name,
&discovered,
&access_token,
&refresh_token,
expiration_time,
)?;
}
Err(_) => {
tracing::warn!(
target: "xdk::auth",
"token exchange succeeded but /2/users/me lookup failed; token stored under unnamed slot"
);
auth.token_store.save_oauth2_token_unnamed_for_app(
&app_name,
&access_token,
&refresh_token,
expiration_time,
)?;
}
}
} else {
auth.token_store.save_oauth2_token_for_app(
&app_name,
username,
&access_token,
&refresh_token,
expiration_time,
)?;
}
let _ = auth
.token_store
.promote_to_default_if_first_credentialed(&app_name)?;
Ok(access_token)
}
pub async fn run_oauth2_flow<F>(
auth: &mut Auth,
http: &reqwest::Client,
username: &str,
cancel: CancellationToken,
browser_opener: F,
) -> Result<String>
where
F: Fn(&str) -> std::io::Result<()> + Send + Sync + 'static,
{
let state_bytes: [u8; 32] = rand::random();
let state = BASE64_STANDARD.encode(state_bytes);
let (verifier, challenge) = generate_code_verifier_and_challenge();
let auth_url_str = build_auth_url(auth, &state, &challenge)?;
let redirect_parsed =
Url::parse(auth.redirect_uri()).map_err(|e| Error::auth_with_cause("InvalidURL", &e))?;
let opener_failed = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false));
let opener_failed_for_closure = std::sync::Arc::clone(&opener_failed);
let cancel_for_closure = cancel.clone();
let on_bound = move || {
if browser_opener(&auth_url_str).is_err() {
opener_failed_for_closure.store(true, std::sync::atomic::Ordering::SeqCst);
cancel_for_closure.cancel();
}
};
let code_result =
callback::wait_for_callback_with(&redirect_parsed, &state, cancel, on_bound).await;
if opener_failed.load(std::sync::atomic::Ordering::SeqCst) {
return Err(Error::auth(
"browser-open failed; re-run with --no-browser to paste the URL manually",
));
}
let code = code_result?;
exchange_code_for_token(auth, http, &code, &verifier, username).await
}
pub fn run_remote_step1(auth: &Auth, pending_path: &std::path::Path) -> Result<String> {
if pending_path.exists() {
tracing::warn!(target: "xdk::auth", "overwriting previous pending auth flow");
}
let state_bytes: [u8; 32] = rand::random();
let state = BASE64_STANDARD.encode(state_bytes);
let (verifier, challenge) = generate_code_verifier_and_challenge();
let auth_url_str = build_auth_url(auth, &state, &challenge)?;
let now = std::time::SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let pending_state = pending::PendingOAuth2State {
code_verifier: verifier,
state,
client_id: auth.client_id().to_string(),
app_name: auth.app_name().to_string(),
created_at: now,
};
pending::save(&pending_state, pending_path)?;
Ok(auth_url_str)
}
pub async fn run_remote_step2(
auth: &mut Auth,
http: &reqwest::Client,
redirect_url: &str,
username: &str,
pending_path: &std::path::Path,
) -> Result<String> {
let pending_state = pending::load(pending_path)?;
if pending_state.client_id != auth.client_id() {
return Err(Error::auth(format!(
"AppMismatch: pending state was created for app {:?} (client_id: {}), \
but current context uses client_id: {}. Re-run step 1 with the correct --app",
pending_state.app_name,
pending_state.client_id,
auth.client_id()
)));
}
let parsed = Url::parse(redirect_url).map_err(|e| {
Error::auth_with_cause("InvalidRedirectURL: failed to parse redirect URL", &e)
})?;
let params: std::collections::HashMap<String, String> = parsed
.query_pairs()
.map(|(k, v)| (k.to_string(), v.to_string()))
.collect();
let state = params
.get("state")
.ok_or_else(|| Error::auth("MissingState: no 'state' parameter found in redirect URL"))?;
if *state != pending_state.state {
return Err(Error::auth(
"StateMismatch: the state parameter in the redirect URL does not match \
the pending auth flow. This may indicate a CSRF attack or that step 1 \
was re-run. Please start over with step 1",
));
}
let code = params.get("code").ok_or_else(|| {
Error::auth(
"MissingCode: no 'code' parameter found in redirect URL. \
Make sure you copied the full URL from your browser's address bar",
)
})?;
let access_token =
exchange_code_for_token(auth, http, code, &pending_state.code_verifier, username).await?;
pending::delete(pending_path)?;
Ok(access_token)
}
pub(crate) fn stored_oauth2_token(auth: &Auth, username: &str) -> Option<OAuth2Token> {
let app_name = auth.app_name().to_string();
let token = if username.is_empty() {
auth.token_store
.get_first_oauth2_token_for_app(&app_name)
.or_else(|| auth.token_store.get_oauth2_token_unnamed_for_app(&app_name))
} else {
auth.token_store
.get_oauth2_token_for_app(&app_name, username)
};
token.and_then(|t| t.oauth2.clone())
}
pub(crate) fn is_expired(token: &OAuth2Token) -> bool {
epoch_secs(SystemTime::now()) >= token.expiration_time
}
pub(crate) fn epoch_secs(at: SystemTime) -> u64 {
at.duration_since(UNIX_EPOCH).unwrap_or_default().as_secs()
}
pub(crate) struct RefreshedToken {
pub(crate) access_token: String,
pub(crate) refresh_token: Option<String>,
pub(crate) expires_at: SystemTime,
}
impl RefreshedToken {
pub(crate) fn expiration_time(&self) -> u64 {
epoch_secs(self.expires_at)
}
}
pub(crate) async fn refresh_grant(
http: &reqwest::Client,
token_url: &str,
timeout: Duration,
client_id: &str,
client_secret: &str,
refresh_token: &str,
) -> Result<RefreshedToken> {
let token_resp = http
.post(token_url)
.timeout(timeout)
.form(&[
("grant_type", "refresh_token"),
("refresh_token", refresh_token),
("client_id", client_id),
])
.basic_auth(client_id, Some(client_secret))
.send()
.await
.map_err(|e| Error::auth_with_cause("RefreshTokenError", &e))?;
let token_data: serde_json::Value = token_resp
.json()
.await
.map_err(|e| Error::auth_with_cause("RefreshTokenError", &e))?;
let access_token = token_data["access_token"]
.as_str()
.ok_or_else(|| Error::auth("RefreshTokenError: no access_token in response"))?
.to_string();
let refresh_token = token_data["refresh_token"]
.as_str()
.filter(|token| !token.is_empty())
.map(str::to_string);
let expires_in = token_data["expires_in"].as_u64().unwrap_or(7200);
Ok(RefreshedToken {
access_token,
refresh_token,
expires_at: SystemTime::now() + Duration::from_secs(expires_in),
})
}
pub async fn refresh_oauth2_token(
auth: &mut Auth,
http: &reqwest::Client,
username: &str,
) -> Result<String> {
let oauth2 = stored_oauth2_token(auth, username)
.ok_or_else(|| Error::auth(crate::error::NO_OAUTH2_TOKEN))?;
if !is_expired(&oauth2) {
return Ok(oauth2.access_token.clone());
}
let refreshed = refresh_grant(
http,
auth.token_url(),
Duration::from_secs(auth.http_timeout_secs()),
auth.client_id(),
auth.client_secret(),
&oauth2.refresh_token,
)
.await?;
let expiration_time = refreshed.expiration_time();
let new_access_token = refreshed.access_token;
let new_refresh_token = refreshed
.refresh_token
.unwrap_or_else(|| oauth2.refresh_token.clone());
let app_name = auth.app_name().to_string();
if username.is_empty() {
match auth.fetch_username(http, &new_access_token).await {
Ok(discovered) => {
auth.token_store.save_oauth2_token_for_app(
&app_name,
&discovered,
&new_access_token,
&new_refresh_token,
expiration_time,
)?;
}
Err(_) => {
tracing::warn!(
target: "xdk::auth",
"refresh succeeded but /2/users/me lookup failed; token stored under unnamed slot"
);
auth.token_store.save_oauth2_token_unnamed_for_app(
&app_name,
&new_access_token,
&new_refresh_token,
expiration_time,
)?;
}
}
} else {
auth.token_store.save_oauth2_token_for_app(
&app_name,
username,
&new_access_token,
&new_refresh_token,
expiration_time,
)?;
}
Ok(new_access_token)
}