use std::io::{BufRead, BufReader, Write};
use std::net::{TcpListener, TcpStream};
use std::process::{Command, Stdio};
use std::sync::{Mutex, PoisonError};
use std::thread;
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
use base64::Engine;
use secrecy::SecretString;
use serde::Deserialize;
use sha2::{Digest, Sha256};
use url::form_urlencoded;
use super::store::{self, StoredToken};
use super::{build_client, classify_status, validate, AuthError, Credential, Platform, Session};
fn authorize_endpoint(platform: &Platform) -> String {
format!("{}/auth/oauth/authorize", platform.base)
}
fn token_endpoint(platform: &Platform) -> String {
format!("{}/auth/oauth/token", platform.base)
}
fn logout_endpoint(platform: &Platform) -> String {
format!("{}/auth/oauth/logout", platform.base)
}
const LOOPBACK_HOST: &str = "127.0.0.1";
const CALLBACK_PATH: &str = "/callback";
const LOGIN_TIMEOUT: Duration = Duration::from_mins(5);
const POLL_INTERVAL: Duration = Duration::from_millis(200);
const READ_TIMEOUT: Duration = Duration::from_secs(5);
const REFRESH_SKEW_SECONDS: u64 = 60;
static REFRESH_GUARD: Mutex<()> = Mutex::new(());
pub(super) fn login(platform: &Platform) -> Result<Session, AuthError> {
let listener = bind_loopback()?;
let redirect_uri = redirect_uri(&listener)?;
let pkce = Pkce::new();
let state = new_opaque();
let url = build_authorize_url(platform, &redirect_uri, &pkce.challenge, &state)?;
announce_login(&url);
open_browser(&url);
let code = await_callback(platform, &listener, &state)?;
let tokens = exchange_code(platform, &code, &pkce.verifier, &redirect_uri)?;
persist_session(platform, tokens)
}
pub(super) fn token(platform: &Platform, stored: StoredToken) -> Result<Session, AuthError> {
if !stored.is_access_expired(now(), REFRESH_SKEW_SECONDS) {
return Ok(session_of(stored.access_token, stored.email));
}
refresh_guarded(platform, &Reuse::WhenFresh)
}
pub(super) fn refresh_stale(platform: &Platform, stale: &str) -> Result<Session, AuthError> {
refresh_guarded(platform, &Reuse::WhenReplaced(stale))
}
enum Reuse<'a> {
WhenReplaced(&'a str),
WhenFresh,
}
fn reusable(stored: &StoredToken, mode: &Reuse<'_>, now: u64) -> bool {
match mode {
Reuse::WhenReplaced(stale) => stored.access_token != *stale,
Reuse::WhenFresh => !stored.is_access_expired(now, REFRESH_SKEW_SECONDS),
}
}
fn refresh_guarded(platform: &Platform, mode: &Reuse<'_>) -> Result<Session, AuthError> {
let _guard = REFRESH_GUARD.lock().unwrap_or_else(PoisonError::into_inner);
let Some(stored) = load_stored(platform)? else {
return Err(AuthError::NotAuthenticated);
};
if reusable(&stored, mode, now()) {
return Ok(session_of(stored.access_token, stored.email));
}
let access_token = refresh_stored(platform, &stored)?;
Ok(session_of(access_token, stored.email))
}
pub(super) fn logout(platform: &Platform) -> Result<bool, AuthError> {
let revoked = match load_stored(platform)? {
Some(stored) => post_logout(platform, &stored.refresh_token).is_ok(),
None => true,
};
store::delete(platform.store_key.as_deref()).map_err(local)?;
Ok(revoked)
}
pub(super) fn stored_session(platform: &Platform) -> Result<Option<StoredToken>, AuthError> {
load_stored(platform)
}
pub(super) fn validated_token(
platform: &Platform,
stored: StoredToken,
) -> Result<Session, AuthError> {
let session = token(platform, stored)?;
let email = validate(platform, session.expose_token())?;
Ok(session_of(session.expose_token().to_owned(), email))
}
fn session_of(access_token: String, email: String) -> Session {
Session {
email,
source: Credential::Oauth,
token: SecretString::from(access_token),
}
}
struct Pkce {
verifier: String,
challenge: String,
}
impl Pkce {
fn new() -> Self {
let verifier = new_opaque();
let challenge = challenge_for(&verifier);
Self {
verifier,
challenge,
}
}
}
fn challenge_for(verifier: &str) -> String {
URL_SAFE_NO_PAD.encode(Sha256::digest(verifier.as_bytes()))
}
fn new_opaque() -> String {
let mut bytes = Vec::with_capacity(32);
bytes.extend_from_slice(uuid::Uuid::new_v4().as_bytes());
bytes.extend_from_slice(uuid::Uuid::new_v4().as_bytes());
URL_SAFE_NO_PAD.encode(&bytes)
}
fn bind_loopback() -> Result<TcpListener, AuthError> {
TcpListener::bind((LOOPBACK_HOST, 0)).map_err(local)
}
fn redirect_uri(listener: &TcpListener) -> Result<String, AuthError> {
let port = listener.local_addr().map_err(local)?.port();
Ok(format!("http://{LOOPBACK_HOST}:{port}{CALLBACK_PATH}"))
}
fn build_authorize_url(
platform: &Platform,
redirect_uri: &str,
challenge: &str,
state: &str,
) -> Result<String, AuthError> {
let mut url =
reqwest::Url::parse(&authorize_endpoint(platform)).map_err(|err| local(err.to_string()))?;
url.query_pairs_mut()
.append_pair("response_type", "code")
.append_pair("client_id", super::CLIENT_ID)
.append_pair("redirect_uri", redirect_uri)
.append_pair("code_challenge", challenge)
.append_pair("code_challenge_method", "S256")
.append_pair("state", state);
Ok(String::from(url))
}
#[allow(clippy::print_stderr)]
fn announce_login(url: &str) {
eprintln!("Opening your browser to log in. If it does not open, visit:\n{url}");
}
fn open_browser(url: &str) {
let opener = if cfg!(target_os = "macos") {
"open"
} else {
"xdg-open"
};
let spawned = Command::new(opener)
.arg(url)
.stdout(Stdio::null())
.stderr(Stdio::null())
.spawn();
if let Err(err) = spawned {
tracing::debug!(error = %err, "could not launch a browser; the login URL was printed");
}
}
fn await_callback(
platform: &Platform,
listener: &TcpListener,
state: &str,
) -> Result<String, AuthError> {
let deadline = Instant::now()
.checked_add(LOGIN_TIMEOUT)
.ok_or_else(|| local("clock overflow"))?;
listener.set_nonblocking(true).map_err(local)?;
while Instant::now() < deadline {
match listener.accept() {
Ok((stream, _)) => {
if let Some(code) = handle_stream(platform, stream, state)? {
return Ok(code);
}
}
Err(ref err) if err.kind() == std::io::ErrorKind::WouldBlock => {
thread::sleep(POLL_INTERVAL);
}
Err(err) => return Err(local(err.to_string())),
}
}
Err(local("timed out waiting for the browser redirect"))
}
fn handle_stream(
platform: &Platform,
mut stream: TcpStream,
expected_state: &str,
) -> Result<Option<String>, AuthError> {
let _ = stream.set_read_timeout(Some(READ_TIMEOUT));
let Some(request_line) = read_request_line(&stream) else {
return Ok(None);
};
let target = request_line.split_whitespace().nth(1).unwrap_or_default();
let params = parse_callback_query(target);
let outcome = classify_callback(¶ms, expected_state);
write_response(platform, &mut stream, matches!(outcome, Ok(Some(_))));
outcome
}
fn read_request_line(stream: &TcpStream) -> Option<String> {
let mut line = String::new();
match BufReader::new(stream).read_line(&mut line) {
Ok(0) | Err(_) => None,
Ok(_) => Some(line),
}
}
#[derive(Default)]
struct CallbackParams {
code: Option<String>,
state: Option<String>,
error: Option<String>,
}
fn parse_callback_query(target: &str) -> CallbackParams {
let mut params = CallbackParams::default();
let Ok(url) = reqwest::Url::parse("http://127.0.0.1/").and_then(|base| base.join(target))
else {
return params;
};
for (key, value) in url.query_pairs() {
match key.as_ref() {
"code" => params.code = Some(value.into_owned()),
"state" => params.state = Some(value.into_owned()),
"error" => params.error = Some(value.into_owned()),
_ => {}
}
}
params
}
fn classify_callback(
params: &CallbackParams,
expected_state: &str,
) -> Result<Option<String>, AuthError> {
if params.state.as_deref() != Some(expected_state) {
return Ok(None);
}
if params.error.is_some() {
return Err(AuthError::Invalid);
}
Ok(params.code.clone())
}
fn write_response(platform: &Platform, stream: &mut TcpStream, ok: bool) {
let outcome = if ok { "done" } else { "error" };
let location = format!("{}/auth/cli/{outcome}", platform.base);
let response = format!(
"HTTP/1.1 302 Found\r\nLocation: {location}\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
);
let _ = stream.write_all(response.as_bytes());
let _ = stream.flush();
}
struct Tokens {
access_token: String,
refresh_token: String,
expires_in: u64,
}
fn exchange_code(
platform: &Platform,
code: &str,
verifier: &str,
redirect_uri: &str,
) -> Result<Tokens, AuthError> {
post_token(
platform,
&[
("grant_type", "authorization_code"),
("code", code),
("code_verifier", verifier),
("client_id", super::CLIENT_ID),
("redirect_uri", redirect_uri),
],
)
}
fn post_refresh(platform: &Platform, refresh_token: &str) -> Result<Tokens, AuthError> {
post_token(
platform,
&[
("grant_type", "refresh_token"),
("refresh_token", refresh_token),
],
)
}
fn post_token(platform: &Platform, form: &[(&str, &str)]) -> Result<Tokens, AuthError> {
let response = post_form(platform, &token_endpoint(platform), form)?;
let status = response.status();
if !status.is_success() {
return Err(classify_status(status.as_u16()));
}
let body = response
.text()
.map_err(|err| AuthError::Transport(err.without_url().to_string()))?;
parse_token_response(&body)
}
fn post_form(
platform: &Platform,
endpoint: &str,
form: &[(&str, &str)],
) -> Result<reqwest::blocking::Response, AuthError> {
let body = form_urlencoded::Serializer::new(String::new())
.extend_pairs(form)
.finish();
build_client(platform)?
.post(endpoint)
.header("Content-Type", "application/x-www-form-urlencoded")
.header("User-Agent", super::user_agent())
.body(body)
.send()
.map_err(|err| AuthError::Transport(err.without_url().to_string()))
}
#[derive(Deserialize)]
struct TokenResponse {
access_token: Option<String>,
refresh_token: Option<String>,
expires_in: Option<u64>,
}
fn parse_token_response(body: &str) -> Result<Tokens, AuthError> {
let parsed: TokenResponse = serde_json::from_str(body)
.map_err(|_| AuthError::Transport("unexpected response from the platform".to_owned()))?;
match (parsed.access_token, parsed.refresh_token, parsed.expires_in) {
(Some(access_token), Some(refresh_token), Some(expires_in))
if !access_token.trim().is_empty() && !refresh_token.trim().is_empty() =>
{
Ok(Tokens {
access_token,
refresh_token,
expires_in,
})
}
_ => Err(AuthError::Invalid),
}
}
fn post_logout(platform: &Platform, refresh_token: &str) -> Result<(), AuthError> {
let _ = post_form(
platform,
&logout_endpoint(platform),
&[("refresh_token", refresh_token)],
)?;
Ok(())
}
fn persist_session(platform: &Platform, tokens: Tokens) -> Result<Session, AuthError> {
let email = validate(platform, &tokens.access_token)?;
save_tokens(platform, &tokens, &email)?;
Ok(Session {
email,
token: SecretString::from(tokens.access_token),
source: Credential::Oauth,
})
}
fn refresh_stored(platform: &Platform, stored: &StoredToken) -> Result<String, AuthError> {
let tokens = match post_refresh(platform, &stored.refresh_token) {
Ok(tokens) => tokens,
Err(AuthError::Invalid) => return reuse_or_forget(platform, stored),
Err(err) => return Err(err),
};
save_tokens(platform, &tokens, &stored.email)?;
Ok(tokens.access_token)
}
fn reuse_or_forget(platform: &Platform, rejected: &StoredToken) -> Result<String, AuthError> {
if let Some(access_token) = reusable_after_rotation(load_stored(platform)?.as_ref(), rejected) {
return Ok(access_token);
}
let _ = store::delete(platform.store_key.as_deref());
Err(AuthError::NotAuthenticated)
}
fn reusable_after_rotation(
current: Option<&StoredToken>,
rejected: &StoredToken,
) -> Option<String> {
current
.filter(|current| current.refresh_token != rejected.refresh_token)
.map(|current| current.access_token.clone())
}
fn save_tokens(platform: &Platform, tokens: &Tokens, email: &str) -> Result<(), AuthError> {
store::save(
platform.store_key.as_deref(),
&StoredToken {
access_token: tokens.access_token.clone(),
refresh_token: tokens.refresh_token.clone(),
expires_at: now().saturating_add(tokens.expires_in),
email: email.to_owned(),
},
)
.map_err(local)
}
fn load_stored(platform: &Platform) -> Result<Option<StoredToken>, AuthError> {
store::load(platform.store_key.as_deref()).map_err(local)
}
#[allow(clippy::needless_pass_by_value)]
fn local(detail: impl ToString) -> AuthError {
AuthError::Local(detail.to_string())
}
fn now() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map_or(0, |elapsed| elapsed.as_secs())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn challenge_matches_rfc7636_vector() {
let verifier = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk";
assert_eq!(
challenge_for(verifier),
"E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM"
);
}
#[test]
fn opaque_values_are_url_safe_and_long() {
let value = new_opaque();
assert!(value.len() >= 43);
assert!(value
.chars()
.all(|c| c.is_ascii_alphanumeric() || c == '-' || c == '_'));
}
#[test]
fn opaque_values_differ_between_calls() {
assert_ne!(new_opaque(), new_opaque());
}
#[test]
fn authorize_url_carries_pkce_and_state() {
let _guard = crate::auth::CONFIG_GUARD
.lock()
.unwrap_or_else(PoisonError::into_inner);
let url = build_authorize_url(
&Platform::for_tests(),
"http://127.0.0.1:1234/callback",
"chal",
"st",
)
.unwrap();
assert!(url.starts_with("https://app.fluidattacks.com/auth/oauth/authorize"));
assert!(url.contains("response_type=code"));
assert!(url.contains("code_challenge_method=S256"));
assert!(url.contains("code_challenge=chal"));
assert!(url.contains("state=st"));
assert!(url.contains("client_id=fluidattacks-cli"));
assert!(url.contains("redirect_uri=http%3A%2F%2F127.0.0.1%3A1234%2Fcallback"));
}
#[test]
fn parse_query_extracts_code_and_state() {
let params = parse_callback_query("/callback?code=abc&state=xyz");
assert_eq!(params.code.as_deref(), Some("abc"));
assert_eq!(params.state.as_deref(), Some("xyz"));
assert!(params.error.is_none());
}
#[test]
fn parse_query_without_query_is_empty() {
let params = parse_callback_query("/favicon.ico");
assert!(params.code.is_none() && params.state.is_none());
}
#[test]
fn classify_requires_matching_state() {
let params = CallbackParams {
code: Some("c".to_owned()),
state: Some("good".to_owned()),
error: None,
};
assert_eq!(
classify_callback(¶ms, "good").unwrap(),
Some("c".to_owned())
);
assert_eq!(classify_callback(¶ms, "bad").unwrap(), None);
}
#[test]
fn classify_ignores_stray_requests() {
let params = CallbackParams::default();
assert_eq!(classify_callback(¶ms, "s").unwrap(), None);
}
#[test]
fn classify_rejects_provider_error_with_matching_state() {
let params = CallbackParams {
code: None,
state: Some("s".to_owned()),
error: Some("access_denied".to_owned()),
};
assert!(matches!(
classify_callback(¶ms, "s"),
Err(AuthError::Invalid)
));
}
#[test]
fn parse_token_response_ok() {
let body =
r#"{"access_token":"a","refresh_token":"r","expires_in":3600,"token_type":"Bearer"}"#;
let tokens = parse_token_response(body).unwrap();
assert_eq!(tokens.access_token, "a");
assert_eq!(tokens.refresh_token, "r");
assert_eq!(tokens.expires_in, 3600);
}
#[test]
fn parse_token_response_missing_or_blank_is_invalid() {
assert!(matches!(
parse_token_response(r#"{"access_token":"a"}"#),
Err(AuthError::Invalid)
));
assert!(matches!(
parse_token_response(r#"{"access_token":"","refresh_token":"r","expires_in":1}"#),
Err(AuthError::Invalid)
));
}
#[test]
fn parse_token_response_non_json_is_transport() {
assert!(matches!(
parse_token_response("<html>502</html>"),
Err(AuthError::Transport(_))
));
}
#[test]
fn pkce_challenge_matches_its_verifier() {
let pkce = Pkce::new();
assert_eq!(pkce.challenge, challenge_for(&pkce.verifier));
}
#[test]
fn redirect_uri_is_loopback_callback() {
let listener = bind_loopback().unwrap();
let uri = redirect_uri(&listener).unwrap();
assert!(uri.starts_with("http://127.0.0.1:"));
assert!(uri.ends_with("/callback"));
}
#[test]
fn local_error_wraps_detail() {
assert!(matches!(local("boom"), AuthError::Local(detail) if detail == "boom"));
}
#[test]
fn now_is_after_epoch() {
assert!(now() > 1_600_000_000);
}
#[test]
fn handle_stream_returns_code_on_matching_state() {
use std::io::Read;
let listener = TcpListener::bind(("127.0.0.1", 0)).unwrap();
let addr = listener.local_addr().unwrap();
let client = thread::spawn(move || {
let mut stream = TcpStream::connect(addr).unwrap();
stream
.write_all(b"GET /callback?code=thecode&state=st HTTP/1.1\r\nHost: x\r\n\r\n")
.unwrap();
let mut buf = Vec::new();
let _ = stream.read_to_end(&mut buf);
buf
});
let (stream, _) = listener.accept().unwrap();
let code = handle_stream(&Platform::for_tests(), stream, "st").unwrap();
let response = client.join().unwrap();
assert_eq!(code.as_deref(), Some("thecode"));
let text = String::from_utf8_lossy(&response);
assert!(text.contains("302 Found"));
assert!(text.contains("/auth/cli/done"));
}
#[test]
fn handle_stream_ignores_state_mismatch() {
use std::io::Read;
let listener = TcpListener::bind(("127.0.0.1", 0)).unwrap();
let addr = listener.local_addr().unwrap();
let client = thread::spawn(move || {
let mut stream = TcpStream::connect(addr).unwrap();
stream
.write_all(b"GET /callback?code=c&state=wrong HTTP/1.1\r\n\r\n")
.unwrap();
let mut buf = Vec::new();
let _ = stream.read_to_end(&mut buf);
});
let (stream, _) = listener.accept().unwrap();
let outcome = handle_stream(&Platform::for_tests(), stream, "expected");
client.join().unwrap();
assert_eq!(outcome.unwrap(), None);
}
#[test]
fn handle_stream_ignores_empty_connection() {
let listener = TcpListener::bind(("127.0.0.1", 0)).unwrap();
let addr = listener.local_addr().unwrap();
let client = thread::spawn(move || {
let _stream = TcpStream::connect(addr).unwrap();
});
let (stream, _) = listener.accept().unwrap();
let outcome = handle_stream(&Platform::for_tests(), stream, "expected");
client.join().unwrap();
assert_eq!(outcome.unwrap(), None);
}
fn stored_with(access_token: &str, expires_at: u64) -> StoredToken {
StoredToken {
access_token: access_token.to_owned(),
refresh_token: "ref".to_owned(),
expires_at,
email: "u@fluidattacks.com".to_owned(),
}
}
#[test]
fn reuse_when_replaced_spots_another_callers_refresh() {
let stored = stored_with("fresh", 1_000);
assert!(reusable(&stored, &Reuse::WhenReplaced("rejected"), 0));
assert!(!reusable(&stored, &Reuse::WhenReplaced("fresh"), 0));
}
#[test]
fn reuse_when_fresh_follows_the_expiry_skew() {
let stored = stored_with("acc", 1_000);
assert!(reusable(&stored, &Reuse::WhenFresh, 900));
assert!(!reusable(
&stored,
&Reuse::WhenFresh,
1_000 - REFRESH_SKEW_SECONDS
));
assert!(!reusable(&stored, &Reuse::WhenFresh, 2_000));
}
#[test]
fn a_session_another_process_refreshed_is_kept() {
let rejected = stored_with("acc", 1_000);
let mut winner = stored_with("fresh", 5_000);
winner.refresh_token = "rotated".to_owned();
assert_eq!(
reusable_after_rotation(Some(&winner), &rejected),
Some("fresh".to_owned())
);
}
#[test]
fn a_session_nobody_refreshed_is_forgotten() {
let rejected = stored_with("acc", 1_000);
assert_eq!(reusable_after_rotation(Some(&rejected), &rejected), None);
assert_eq!(reusable_after_rotation(None, &rejected), None);
}
}