use std::time::Duration;
use base64::Engine;
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
use rand::distr::Alphanumeric;
use rand::{Rng, RngExt, rng};
use serde::Deserialize;
use sha2::{Digest, Sha256};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
use crate::auth::tokens::TokenSet;
use crate::error::ClientError;
pub struct AuthSuccess {
pub tokens: TokenSet,
pub id_token: Option<String>,
}
pub async fn run<F: FnOnce(&str)>(
http: &reqwest::Client,
authority_base: &str,
client_id: &str,
scope: &str,
on_open: F,
) -> Result<AuthSuccess, ClientError> {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.map_err(ClientError::Io)?;
let port = listener.local_addr().map_err(ClientError::Io)?.port();
let redirect_uri = format!("http://localhost:{port}");
let verifier = make_code_verifier();
let challenge = make_code_challenge(&verifier);
let state = make_random(32);
let authorize_url = build_authorize_url(
authority_base,
client_id,
&redirect_uri,
scope,
&challenge,
&state,
);
on_open(&authorize_url);
let CallbackParams {
code,
state: returned_state,
} = wait_for_callback(listener, Duration::from_secs(300), "Microsoft sign-in").await?;
if returned_state != state {
return Err(ClientError::Graph {
status: 400,
message: "OAuth state mismatch: possible CSRF or stale request".to_string(),
});
}
exchange_code_to_success(
http,
authority_base,
client_id,
&code,
&verifier,
&redirect_uri,
)
.await
}
pub(crate) async fn exchange_code_to_success(
http: &reqwest::Client,
authority_base: &str,
client_id: &str,
code: &str,
code_verifier: &str,
redirect_uri: &str,
) -> Result<AuthSuccess, ClientError> {
let tokens_response = exchange_code(
http,
authority_base,
client_id,
code,
code_verifier,
redirect_uri,
)
.await?;
Ok(AuthSuccess {
tokens: TokenSet {
access_token: tokens_response.access_token,
refresh_token: tokens_response.refresh_token.unwrap_or_default(),
expires_at: chrono::Utc::now()
+ chrono::Duration::seconds(tokens_response.expires_in.unwrap_or(3600)),
},
id_token: tokens_response.id_token,
})
}
pub(crate) fn build_authorize_url(
authority_base: &str,
client_id: &str,
redirect_uri: &str,
scope: &str,
challenge: &str,
state: &str,
) -> String {
build_authorize_url_with_hint(
authority_base,
client_id,
redirect_uri,
scope,
challenge,
state,
None,
)
}
pub(crate) fn build_authorize_url_with_hint(
authority_base: &str,
client_id: &str,
redirect_uri: &str,
scope: &str,
challenge: &str,
state: &str,
login_hint: Option<&str>,
) -> String {
let mut url = url::Url::parse(&format!("{authority_base}/oauth2/v2.0/authorize"))
.expect("authority_base is a valid URL");
{
let mut q = url.query_pairs_mut();
q.append_pair("client_id", client_id)
.append_pair("response_type", "code")
.append_pair("redirect_uri", redirect_uri)
.append_pair("response_mode", "query")
.append_pair("scope", scope)
.append_pair("state", state)
.append_pair("code_challenge", challenge)
.append_pair("code_challenge_method", "S256")
.append_pair("prompt", "select_account");
if let Some(hint) = login_hint.map(str::trim).filter(|h| !h.is_empty()) {
q.append_pair("login_hint", hint);
}
}
url.into()
}
pub(crate) struct CallbackParams {
pub(crate) code: String,
pub(crate) state: String,
}
fn timeout_label(timeout: Duration) -> String {
let secs = timeout.as_secs();
if secs >= 60 && secs.is_multiple_of(60) {
format!("{} min", secs / 60)
} else {
format!("{secs}s")
}
}
pub(crate) async fn wait_for_callback(
listener: TcpListener,
timeout: Duration,
label: &str,
) -> Result<CallbackParams, ClientError> {
let deadline = tokio::time::Instant::now() + timeout;
let timed_out = || ClientError::Graph {
status: 408,
message: format!(
"timed out waiting for browser sign-in ({})",
timeout_label(timeout)
),
};
let (mut stream, _) = tokio::time::timeout_at(deadline, listener.accept())
.await
.map_err(|_| timed_out())?
.map_err(ClientError::Io)?;
let mut buf = [0u8; 2048];
let n = tokio::time::timeout_at(deadline, stream.read(&mut buf))
.await
.map_err(|_| timed_out())?
.map_err(ClientError::Io)?;
let request = std::str::from_utf8(&buf[..n]).unwrap_or("");
let first_line = request.lines().next().unwrap_or("");
let path_and_query =
first_line
.split_whitespace()
.nth(1)
.ok_or_else(|| ClientError::Graph {
status: 400,
message: "malformed browser callback request".to_string(),
})?;
let query_start = path_and_query.find('?').unwrap_or(path_and_query.len());
let query = &path_and_query[query_start.saturating_add(1)..];
let pairs: Vec<(String, String)> = url::form_urlencoded::parse(query.as_bytes())
.map(|(k, v)| (k.into_owned(), v.into_owned()))
.collect();
let mut code: Option<String> = None;
let mut state: Option<String> = None;
let mut err: Option<String> = None;
let mut err_description: Option<String> = None;
for (k, v) in pairs {
match k.as_str() {
"code" => code = Some(v),
"state" => state = Some(v),
"error" => err = Some(v),
"error_description" => err_description = Some(v),
_ => {}
}
}
if let Some(e) = err {
let detail = err_description.unwrap_or_default();
write_html(&mut stream, &error_page_html(&e, &detail), "Sign-in failed")
.await
.ok();
return Err(ClientError::Graph {
status: 400,
message: format!("{label}: {e} ({detail})"),
});
}
let code = code.ok_or_else(|| ClientError::Graph {
status: 400,
message: "browser callback missing `code` parameter".to_string(),
})?;
let state = state.unwrap_or_default();
write_html(&mut stream, SUCCESS_HTML, "Signed in")
.await
.ok();
Ok(CallbackParams { code, state })
}
async fn write_html(
stream: &mut tokio::net::TcpStream,
body: &str,
title: &str,
) -> std::io::Result<()> {
let body_bytes = body.as_bytes();
let response = format!(
"HTTP/1.1 200 OK\r\n\
Content-Type: text/html; charset=utf-8\r\n\
Content-Length: {}\r\n\
Connection: close\r\n\
X-Title: {}\r\n\
\r\n",
body_bytes.len(),
title,
);
stream.write_all(response.as_bytes()).await?;
stream.write_all(body_bytes).await?;
stream.shutdown().await
}
const SUCCESS_HTML: &str = r#"<!doctype html><html><head><meta charset="utf-8"><title>Signed in</title>
<style>
body { font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", system-ui, sans-serif;
max-width: 480px; margin: 80px auto; text-align: center; color: #1d1d1f; }
.check { font-size: 48px; color: #34c759; }
h1 { font-size: 24px; margin: 16px 0 8px; }
p { color: #6e6e73; }
</style></head>
<body>
<div class="check">✓</div>
<h1>Signed in to pidge</h1>
<p>You can close this window and return to the terminal.</p>
</body></html>"#;
fn error_page_html(err: &str, description: &str) -> String {
let err = html_escape(err);
let description = html_escape(description);
format!(
r#"<!doctype html><html><head><meta charset="utf-8"><title>Sign-in failed</title>
<style>
body {{ font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", system-ui, sans-serif;
max-width: 540px; margin: 80px auto; color: #1d1d1f; }}
.x {{ font-size: 48px; color: #ff3b30; text-align: center; }}
h1 {{ font-size: 22px; margin: 16px 0 8px; text-align: center; }}
.detail {{ background: #f5f5f7; padding: 16px; border-radius: 8px; font-family: ui-monospace, monospace;
font-size: 13px; white-space: pre-wrap; word-break: break-word; }}
</style></head>
<body>
<div class="x">✕</div>
<h1>Sign-in failed</h1>
<p class="detail"><strong>{err}</strong>
{description}</p>
<p>You can close this window. Return to the terminal for next steps.</p>
</body></html>"#
)
}
fn html_escape(s: &str) -> String {
s.chars()
.fold(String::with_capacity(s.len()), |mut acc, c| {
match c {
'&' => acc.push_str("&"),
'<' => acc.push_str("<"),
'>' => acc.push_str(">"),
'"' => acc.push_str("""),
'\'' => acc.push_str("'"),
_ => acc.push(c),
}
acc
})
}
pub(crate) fn make_code_verifier() -> String {
let mut rng = rng();
(0..64).map(|_| rng.sample(Alphanumeric) as char).collect()
}
pub(crate) fn make_code_challenge(verifier: &str) -> String {
let mut hasher = Sha256::new();
hasher.update(verifier.as_bytes());
URL_SAFE_NO_PAD.encode(hasher.finalize())
}
pub(crate) fn make_random(byte_len: usize) -> String {
let mut buf = vec![0u8; byte_len];
rng().fill_bytes(&mut buf);
URL_SAFE_NO_PAD.encode(&buf)
}
#[derive(Debug, Deserialize)]
struct TokenResponse {
access_token: String,
refresh_token: Option<String>,
expires_in: Option<i64>,
id_token: Option<String>,
}
async fn exchange_code(
http: &reqwest::Client,
authority_base: &str,
client_id: &str,
code: &str,
code_verifier: &str,
redirect_uri: &str,
) -> Result<TokenResponse, ClientError> {
let url = format!("{authority_base}/oauth2/v2.0/token");
let params = [
("client_id", client_id),
("grant_type", "authorization_code"),
("code", code),
("code_verifier", code_verifier),
("redirect_uri", redirect_uri),
];
let resp = http.post(&url).form(¶ms).send().await?;
let status = resp.status();
if !status.is_success() {
let text = resp.text().await.unwrap_or_default();
return Err(ClientError::Graph {
status: status.as_u16(),
message: text,
});
}
Ok(resp.json().await?)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn timeout_label_uses_minutes_when_whole() {
assert_eq!(timeout_label(Duration::from_secs(300)), "5 min");
assert_eq!(timeout_label(Duration::from_secs(600)), "10 min");
assert_eq!(timeout_label(Duration::from_secs(45)), "45s");
}
#[test]
fn code_verifier_is_64_alphanumerics() {
let v = make_code_verifier();
assert_eq!(v.len(), 64);
assert!(v.chars().all(|c| c.is_ascii_alphanumeric()));
}
#[test]
fn challenge_matches_rfc_7636_example() {
let verifier = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk";
assert_eq!(
make_code_challenge(verifier),
"E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM"
);
}
#[test]
fn authorize_url_contains_required_params() {
let url = build_authorize_url(
"https://login.microsoftonline.com/common",
"client-id-here",
"http://localhost:47821",
"User.Read offline_access",
"challenge-here",
"state-here",
);
assert!(url.contains("client_id=client-id-here"));
assert!(url.contains("response_type=code"));
assert!(url.contains("code_challenge=challenge-here"));
assert!(url.contains("code_challenge_method=S256"));
assert!(url.contains("state=state-here"));
assert!(url.contains("redirect_uri=http%3A%2F%2Flocalhost%3A47821"));
assert!(url.contains("prompt=select_account"));
assert!(!url.contains("login_hint"));
}
#[test]
fn authorize_url_carries_the_login_hint_when_given() {
let url = build_authorize_url_with_hint(
"https://login.microsoftonline.com/common",
"client-id-here",
"http://localhost:47821",
"scope-here",
"challenge-here",
"state-here",
Some("jane.doe@example.com"),
);
assert!(url.contains("login_hint=jane.doe%40example.com"), "{url}");
assert!(url.contains("prompt=select_account"));
let blank = build_authorize_url_with_hint(
"https://login.microsoftonline.com/common",
"c",
"http://localhost:1",
"s",
"ch",
"st",
Some(" "),
);
assert!(!blank.contains("login_hint"));
}
#[test]
fn random_state_is_unique_per_call() {
let a = make_random(32);
let b = make_random(32);
assert_ne!(a, b);
assert_eq!(URL_SAFE_NO_PAD.decode(&a).unwrap().len(), 32);
}
#[test]
fn error_page_html_escapes_html_special_characters() {
let page = error_page_html("<script>alert(1)</script>", "quote\" apostrophe' amp&");
assert!(!page.contains("<script>"));
assert!(page.contains("<script>alert(1)</script>"));
assert!(page.contains("quote" apostrophe' amp&"));
}
}