use std::sync::{Arc, RwLock};
use std::time::Duration;
use serde::Deserialize;
use crate::error::{Error, Result};
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Challenge {
pub x_redirect_server: String,
pub x_token_server: String,
}
pub fn parse_www_authenticate(header: &str) -> Option<Challenge> {
let trimmed = header.trim();
match trimmed.get(..6) {
Some(prefix) if prefix.eq_ignore_ascii_case("bearer") => {}
_ => return None,
}
let mut x_redirect_server = None;
let mut x_token_server = None;
for part in trimmed.split(',') {
let Some((key, value)) = part.split_once('=') else {
continue;
};
let key = key
.trim()
.rsplit(char::is_whitespace)
.next()
.unwrap_or("")
.trim();
let value = value.trim().trim_matches('"');
match key.to_ascii_lowercase().as_str() {
"x_redirect_server" => x_redirect_server = Some(value.to_string()),
"x_token_server" => x_token_server = Some(value.to_string()),
_ => {}
}
}
Some(Challenge {
x_redirect_server: x_redirect_server.unwrap_or_default(),
x_token_server: x_token_server?,
})
}
pub trait RedirectHandler: Send + Sync {
fn redirect(&self, url: &str) -> Result<()>;
}
pub struct BrowserRedirectHandler;
impl RedirectHandler for BrowserRedirectHandler {
fn redirect(&self, url: &str) -> Result<()> {
eprintln!("Open the following URL in a browser to authenticate:\n{url}");
let _ = open::that(url);
Ok(())
}
}
pub struct OAuth2State {
pub(crate) token: RwLock<Option<String>>,
pub(crate) acquire: tokio::sync::Mutex<()>,
pub(crate) handler: Arc<dyn RedirectHandler>,
pub(crate) max_poll_attempts: usize,
pub(crate) poll_timeout: Duration,
}
impl OAuth2State {
pub fn new(
handler: Arc<dyn RedirectHandler>,
max_poll_attempts: usize,
poll_timeout: Duration,
) -> Self {
Self {
token: RwLock::new(None),
acquire: tokio::sync::Mutex::new(()),
handler,
max_poll_attempts,
poll_timeout,
}
}
pub fn cached_token(&self) -> Option<String> {
self.token.read().unwrap().clone()
}
}
#[derive(Deserialize)]
struct TokenResponse {
token: Option<String>,
#[serde(rename = "nextUri")]
next_uri: Option<String>,
error: Option<String>,
}
pub(crate) async fn run_flow(
client: &reqwest::Client,
state: &OAuth2State,
challenge: &Challenge,
) -> Result<String> {
if !challenge.x_redirect_server.is_empty() {
state.handler.redirect(&challenge.x_redirect_server)?;
}
let deadline = tokio::time::Instant::now() + state.poll_timeout;
let mut url = challenge.x_token_server.clone();
for _ in 0..state.max_poll_attempts {
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
if remaining.is_zero() {
break;
}
let body = match tokio::time::timeout(remaining, async {
let resp = client.get(&url).send().await?;
resp.json::<TokenResponse>().await
})
.await
{
Ok(result) => result?, Err(_elapsed) => break, };
if let Some(err) = body.error {
return Err(Error::OAuth2(format!(
"token endpoint returned error: {err}"
)));
}
if let Some(token) = body.token {
return Ok(token);
}
match body.next_uri {
Some(next) => url = next,
None => {
return Err(Error::OAuth2(
"token endpoint response had neither token nor nextUri".to_string(),
))
}
}
}
Err(Error::OAuth2(format!(
"authentication did not complete within {} attempts / {:?}",
state.max_poll_attempts, state.poll_timeout
)))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parses_standard_challenge() {
let h = r#"Bearer x_redirect_server="https://c/oauth2/token/initiate/abc", x_token_server="https://c/oauth2/token/abc""#;
let c = parse_www_authenticate(h).expect("should parse");
assert_eq!(c.x_redirect_server, "https://c/oauth2/token/initiate/abc");
assert_eq!(c.x_token_server, "https://c/oauth2/token/abc");
}
#[test]
fn tolerates_bearer_prefixed_key_quirk() {
let h = r#"Bearer x_redirect_server="https://c/i", x_token_server="https://c/t""#;
let c = parse_www_authenticate(h).expect("should parse");
assert_eq!(c.x_redirect_server, "https://c/i");
assert_eq!(c.x_token_server, "https://c/t");
}
#[test]
fn ignores_param_order_and_extra_params() {
let h = r#"Bearer realm="trino", x_token_server="https://c/t", x_redirect_server="https://c/i""#;
let c = parse_www_authenticate(h).expect("should parse");
assert_eq!(c.x_token_server, "https://c/t");
assert_eq!(c.x_redirect_server, "https://c/i");
}
#[test]
fn none_when_no_token_server() {
let h = r#"Bearer x_redirect_server="https://c/i""#;
assert!(parse_www_authenticate(h).is_none());
}
#[test]
fn none_when_not_bearer() {
assert!(parse_www_authenticate(r#"Basic realm="trino""#).is_none());
}
#[test]
fn none_on_non_ascii_header_without_panicking() {
assert!(parse_www_authenticate("aaaaaé x_token_server=\"https://c/t\"").is_none());
}
use std::sync::Arc;
use std::time::Duration;
#[test]
fn state_defaults_have_no_token() {
let state = OAuth2State::new(
Arc::new(BrowserRedirectHandler),
10,
Duration::from_secs(120),
);
assert!(state.cached_token().is_none());
}
#[test]
fn state_stores_and_reads_token() {
let state = OAuth2State::new(
Arc::new(BrowserRedirectHandler),
10,
Duration::from_secs(120),
);
*state.token.write().unwrap() = Some("tok".to_string());
assert_eq!(state.cached_token().as_deref(), Some("tok"));
}
struct RecordingHandler {
seen: std::sync::Mutex<Vec<String>>,
}
impl RedirectHandler for RecordingHandler {
fn redirect(&self, url: &str) -> Result<()> {
self.seen.lock().unwrap().push(url.to_string());
Ok(())
}
}
#[tokio::test]
async fn run_flow_follows_next_uri_then_returns_token() {
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/token/step1"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"nextUri": format!("{}/token/step2", server.uri())
})))
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/token/step2"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"token": "final-token"
})))
.mount(&server)
.await;
let handler = Arc::new(RecordingHandler {
seen: std::sync::Mutex::new(vec![]),
});
let state = OAuth2State::new(handler.clone(), 10, Duration::from_secs(30));
let challenge = Challenge {
x_redirect_server: "https://login.example/redirect".to_string(),
x_token_server: format!("{}/token/step1", server.uri()),
};
let token = run_flow(&reqwest::Client::new(), &state, &challenge)
.await
.expect("flow should succeed");
assert_eq!(token, "final-token");
assert_eq!(
handler.seen.lock().unwrap().as_slice(),
&["https://login.example/redirect".to_string()]
);
}
#[tokio::test]
async fn run_flow_surfaces_error_field() {
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/token/err"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"error": "access_denied"
})))
.mount(&server)
.await;
let state = OAuth2State::new(
Arc::new(BrowserRedirectHandler),
10,
Duration::from_secs(30),
);
let challenge = Challenge {
x_redirect_server: String::new(),
x_token_server: format!("{}/token/err", server.uri()),
};
let err = run_flow(&reqwest::Client::new(), &state, &challenge)
.await
.unwrap_err();
match err {
crate::error::Error::OAuth2(msg) => assert!(msg.contains("access_denied")),
other => panic!("expected OAuth2 error, got {other:?}"),
}
}
#[tokio::test]
async fn run_flow_bounded_by_poll_timeout_when_server_stalls() {
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/token/stall"))
.respond_with(
ResponseTemplate::new(200)
.set_delay(Duration::from_secs(30))
.set_body_json(serde_json::json!({ "token": "never-arrives-in-time" })),
)
.mount(&server)
.await;
let state = OAuth2State::new(
Arc::new(BrowserRedirectHandler),
10,
Duration::from_millis(150),
);
let challenge = Challenge {
x_redirect_server: String::new(),
x_token_server: format!("{}/token/stall", server.uri()),
};
let err = run_flow(&reqwest::Client::new(), &state, &challenge)
.await
.unwrap_err();
match err {
crate::error::Error::OAuth2(msg) => assert!(msg.contains("did not complete")),
other => panic!("expected OAuth2 timeout error, got {other:?}"),
}
}
}