use std::collections::HashMap;
use std::fmt;
use base64::Engine;
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
use serde::Deserialize;
use serde::Serialize;
use sha2::Digest;
use sha2::Sha256;
use crate::Error;
use crate::Res;
use crate::error::AuthError;
use crate::io::remote::client::HttpClient;
use crate::io::storage::auth::OAuthClient;
use crate::io::storage::auth::Tokens;
use quilt_uri::Host;
pub struct OAuthParams {
pub code: String,
pub code_verifier: String,
pub redirect_uri: String,
pub client_id: String,
}
pub struct PkceChallenge {
pub code_verifier: String,
pub code_challenge: String,
}
#[must_use]
pub fn pkce_challenge() -> PkceChallenge {
let mut random_bytes = [0u8; 64];
getrandom::fill(&mut random_bytes).expect("failed to generate random bytes");
let code_verifier = URL_SAFE_NO_PAD.encode(random_bytes);
let code_challenge = URL_SAFE_NO_PAD.encode(Sha256::digest(code_verifier.as_bytes()));
PkceChallenge {
code_verifier,
code_challenge,
}
}
#[must_use]
pub fn random_state() -> String {
let mut bytes = [0u8; 16];
getrandom::fill(&mut bytes).expect("failed to generate random bytes");
URL_SAFE_NO_PAD.encode(bytes)
}
#[must_use]
pub fn catalog_authorize_url(host: &Host) -> String {
format!("https://{host}/connect/authorize")
}
#[must_use]
pub fn connect_host(host: &Host) -> String {
let s = host.to_string();
match s.split_once('.') {
Some((stack, domain)) => format!("{stack}-connect.{domain}"),
None => format!("{s}-connect"),
}
}
pub(super) fn connect_token_url(host: &Host) -> String {
format!("https://{}/auth/token", connect_host(host))
}
pub(super) fn connect_register_url(host: &Host) -> String {
format!("https://{}/auth/register", connect_host(host))
}
#[derive(Serialize)]
struct DcrRequest {
client_name: String,
redirect_uris: Vec<String>,
token_endpoint_auth_method: String,
}
#[derive(Deserialize)]
struct DcrResponse {
client_id: String,
}
pub(super) async fn register_client(
http_client: &impl HttpClient,
host: &Host,
redirect_uri: &str,
) -> Res<OAuthClient> {
let register_url = connect_register_url(host);
let request = DcrRequest {
client_name: "QuiltSync".to_string(),
redirect_uris: vec![redirect_uri.to_string()],
token_endpoint_auth_method: "none".to_string(),
};
let response: DcrResponse = http_client.post_json(®ister_url, &request).await?;
Ok(OAuthClient {
client_id: response.client_id,
redirect_uri: redirect_uri.to_string(),
})
}
pub(super) const DEFAULT_EXPIRES_IN: i64 = 3600;
fn default_expires_in() -> i64 {
DEFAULT_EXPIRES_IN
}
#[derive(Deserialize, Serialize)]
pub(super) struct OAuthTokenResponse {
pub(super) access_token: String,
#[serde(default)]
pub(super) refresh_token: Option<String>,
#[serde(default = "default_expires_in")]
pub(super) expires_in: i64,
}
impl fmt::Debug for OAuthTokenResponse {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("OAuthTokenResponse")
.field("expires_in", &self.expires_in)
.field("access_token", &"[REDACTED]")
.field(
"refresh_token",
&self.refresh_token.as_ref().map(|_| "[REDACTED]"),
)
.finish_non_exhaustive()
}
}
pub(super) async fn exchange_oauth_code(
http_client: &impl HttpClient,
host: &Host,
params: &OAuthParams,
) -> Res<Tokens> {
let token_url = connect_token_url(host);
let mut form_data: HashMap<String, String> = HashMap::new();
form_data.insert("grant_type".to_string(), "authorization_code".to_string());
form_data.insert("code".to_string(), params.code.clone());
form_data.insert("code_verifier".to_string(), params.code_verifier.clone());
form_data.insert("redirect_uri".to_string(), params.redirect_uri.clone());
form_data.insert("client_id".to_string(), params.client_id.clone());
let response: OAuthTokenResponse = http_client.post(&token_url, &form_data).await?;
let expires_at = chrono::Utc::now() + chrono::Duration::seconds(response.expires_in);
Ok(Tokens {
access_token: response.access_token,
refresh_token: response.refresh_token.ok_or_else(|| {
Error::Auth(
host.to_owned(),
AuthError::TokensExchange("server did not return a refresh token".to_string()),
)
})?,
expires_at,
})
}
pub(super) async fn refresh_oauth_tokens(
http_client: &impl HttpClient,
host: &Host,
refresh_token: &str,
client_id: &str,
) -> Res<Tokens> {
let token_url = connect_token_url(host);
let mut form_data: HashMap<String, String> = HashMap::new();
form_data.insert("grant_type".to_string(), "refresh_token".to_string());
form_data.insert("refresh_token".to_string(), refresh_token.to_string());
form_data.insert("client_id".to_string(), client_id.to_string());
let response: OAuthTokenResponse = http_client.post(&token_url, &form_data).await?;
let expires_at = chrono::Utc::now() + chrono::Duration::seconds(response.expires_in);
Ok(Tokens {
access_token: response.access_token,
refresh_token: response
.refresh_token
.unwrap_or_else(|| refresh_token.to_string()),
expires_at,
})
}
#[cfg(test)]
mod tests {
use super::*;
use async_trait::async_trait;
use test_log::test;
use crate::auth::test_utils::*;
#[test]
fn test_connect_host() {
let host: Host = "test.quilt.dev".parse().unwrap();
assert_eq!(connect_host(&host), "test-connect.quilt.dev");
}
#[test]
fn test_connect_token_url() {
let host: Host = "test.quilt.dev".parse().unwrap();
assert_eq!(
connect_token_url(&host),
"https://test-connect.quilt.dev/auth/token"
);
}
#[test(tokio::test)]
async fn test_exchange_oauth_code() {
let client = OAuthTestHttpClient::default();
let params = OAuthParams {
code: AUTH_CODE.to_string(),
code_verifier: CODE_VERIFIER.to_string(),
redirect_uri: REDIRECT_URI.to_string(),
client_id: CLIENT_ID.to_string(),
};
let tokens = exchange_oauth_code(&client, &get_host(), ¶ms)
.await
.unwrap();
assert_eq!(tokens.access_token, ACCESS_TOKEN);
assert_eq!(tokens.refresh_token, "oauth-refresh-token");
}
#[test]
fn test_pkce_challenge() {
let pkce = pkce_challenge();
assert_eq!(pkce.code_verifier.len(), 86);
assert_eq!(pkce.code_challenge.len(), 43);
let expected_challenge =
URL_SAFE_NO_PAD.encode(Sha256::digest(pkce.code_verifier.as_bytes()));
assert_eq!(pkce.code_challenge, expected_challenge);
let pkce2 = pkce_challenge();
assert_ne!(pkce.code_verifier, pkce2.code_verifier);
}
#[test]
fn test_pkce_verifier_charset_rfc7636() {
let pkce = pkce_challenge();
for ch in pkce.code_verifier.chars() {
assert!(
ch.is_ascii_alphanumeric() || matches!(ch, '-' | '.' | '_' | '~'),
"code_verifier contains char '{ch}' not allowed by RFC 7636 §4.1"
);
}
}
#[test(tokio::test)]
async fn test_refresh_oauth_tokens() -> Res {
let tokens = refresh_oauth_tokens(
&OAuthTestHttpClient::default(),
&get_host(),
REFRESH_TOKEN,
CLIENT_ID,
)
.await?;
assert_eq!(tokens.access_token, "refreshed-access-token");
assert_eq!(tokens.refresh_token, "new-refresh-token");
Ok(())
}
#[test(tokio::test)]
async fn test_refresh_oauth_tokens_retains_old_when_omitted() -> Res {
struct NoRefreshTokenClient;
#[async_trait]
impl HttpClient for NoRefreshTokenClient {
async fn get<T: serde::de::DeserializeOwned>(
&self,
_: &str,
_: Option<&str>,
) -> Res<T> {
unimplemented!()
}
async fn head(&self, _: &str) -> Res<reqwest::header::HeaderMap> {
unimplemented!()
}
async fn post<T: serde::de::DeserializeOwned>(
&self,
_: &str,
_: &HashMap<String, String>,
) -> Res<T> {
let resp = OAuthTokenResponse {
access_token: "new-access-token".to_string(),
refresh_token: None, expires_in: DEFAULT_EXPIRES_IN,
};
Ok(serde_json::from_value(serde_json::to_value(resp)?)?)
}
async fn post_json<
T: serde::de::DeserializeOwned,
B: serde::Serialize + Send + Sync,
>(
&self,
_: &str,
_: &B,
) -> Res<T> {
unimplemented!()
}
async fn post_json_auth<
T: serde::de::DeserializeOwned,
B: serde::Serialize + Send + Sync,
>(
&self,
_: &str,
_: &B,
_: &str,
) -> Res<T> {
unimplemented!()
}
}
let tokens =
refresh_oauth_tokens(&NoRefreshTokenClient, &get_host(), REFRESH_TOKEN, CLIENT_ID)
.await?;
assert_eq!(tokens.access_token, "new-access-token");
assert_eq!(tokens.refresh_token, REFRESH_TOKEN);
Ok(())
}
#[test(tokio::test)]
async fn test_exchange_oauth_code_errors_when_refresh_token_missing() {
struct NoRefreshTokenClient;
#[async_trait]
impl HttpClient for NoRefreshTokenClient {
async fn get<T: serde::de::DeserializeOwned>(
&self,
_: &str,
_: Option<&str>,
) -> Res<T> {
unimplemented!()
}
async fn head(&self, _: &str) -> Res<reqwest::header::HeaderMap> {
unimplemented!()
}
async fn post<T: serde::de::DeserializeOwned>(
&self,
_: &str,
_: &HashMap<String, String>,
) -> Res<T> {
let resp = OAuthTokenResponse {
access_token: ACCESS_TOKEN.to_string(),
refresh_token: None,
expires_in: DEFAULT_EXPIRES_IN,
};
Ok(serde_json::from_value(serde_json::to_value(resp)?)?)
}
async fn post_json<
T: serde::de::DeserializeOwned,
B: serde::Serialize + Send + Sync,
>(
&self,
_: &str,
_: &B,
) -> Res<T> {
unimplemented!()
}
async fn post_json_auth<
T: serde::de::DeserializeOwned,
B: serde::Serialize + Send + Sync,
>(
&self,
_: &str,
_: &B,
_: &str,
) -> Res<T> {
unimplemented!()
}
}
let params = OAuthParams {
code: AUTH_CODE.to_string(),
code_verifier: CODE_VERIFIER.to_string(),
redirect_uri: REDIRECT_URI.to_string(),
client_id: CLIENT_ID.to_string(),
};
let result = exchange_oauth_code(&NoRefreshTokenClient, &get_host(), ¶ms).await;
assert!(
matches!(result, Err(Error::Auth(_, AuthError::TokensExchange(_)))),
"expected TokensExchange error, got: {result:?}"
);
}
#[test]
fn test_oauth_token_response_missing_expires_in() {
let json = r#"{"access_token":"tok","refresh_token":"ref"}"#;
let resp: OAuthTokenResponse = serde_json::from_str(json).unwrap();
assert_eq!(resp.expires_in, DEFAULT_EXPIRES_IN);
}
#[test]
fn oauth_token_response_debug_redacts_secrets() {
let response = OAuthTokenResponse {
access_token: "secret-access".to_string(),
refresh_token: Some("secret-refresh".to_string()),
expires_in: 3600,
};
let output = format!("{response:?}");
assert!(output.contains("[REDACTED]"));
assert!(!output.contains("secret-access"));
assert!(!output.contains("secret-refresh"));
}
#[test]
fn oauth_token_response_debug_none_refresh_token() {
let response = OAuthTokenResponse {
access_token: "secret-access".to_string(),
refresh_token: None,
expires_in: 3600,
};
let output = format!("{response:?}");
assert!(output.contains("refresh_token: None"));
assert!(!output.contains("secret-access"));
}
}