use base64::engine::general_purpose::URL_SAFE_NO_PAD;
use base64::Engine;
use sha2::{Digest, Sha256};
use super::oauth::{AuthError, Callback, Grant};
use super::urlencode::{encode_pairs, query_pairs};
use super::OAuthConfig;
use crate::protocol::WireRequest;
pub struct Pkce {
pub verifier: String,
pub challenge: String,
}
impl Pkce {
pub fn derive(verifier: impl Into<String>) -> Pkce {
let verifier = verifier.into();
let challenge = URL_SAFE_NO_PAD.encode(Sha256::digest(verifier.as_bytes()));
Pkce {
verifier,
challenge,
}
}
}
pub fn build_authorize_url(
cfg: &OAuthConfig,
pkce: &Pkce,
state: &str,
redirect_uri: &str,
) -> String {
let mut params = vec![
("response_type", "code"),
("client_id", cfg.client_id.as_str()),
("redirect_uri", redirect_uri),
("state", state),
("code_challenge", pkce.challenge.as_str()),
("code_challenge_method", "S256"),
];
if let Some(scope) = &cfg.scope {
params.push(("scope", scope.as_str()));
}
for (key, value) in &cfg.authorize_params {
params.push((key.as_str(), value.as_str()));
}
format!("{}?{}", cfg.authorize_url, encode_pairs(¶ms))
}
pub fn build_token_exchange_request(cfg: &OAuthConfig, grant: Grant) -> WireRequest {
let cid = cfg.client_id.as_str();
let pairs: Vec<(&str, &str)> = match &grant {
Grant::AuthCode {
code,
verifier,
redirect_uri,
} => vec![
("grant_type", "authorization_code"),
("code", code),
("redirect_uri", redirect_uri),
("code_verifier", verifier),
("client_id", cid),
],
Grant::Device { device_code } => vec![
("grant_type", "urn:ietf:params:oauth:grant-type:device_code"),
("device_code", device_code),
("client_id", cid),
],
Grant::Refresh { refresh_token } => vec![
("grant_type", "refresh_token"),
("refresh_token", refresh_token.expose()),
("client_id", cid),
],
};
form_post(&cfg.token_url, &pairs)
}
pub(crate) fn form_post(url: &str, pairs: &[(&str, &str)]) -> WireRequest {
let mut wire = WireRequest::new(url.to_owned(), encode_pairs(pairs).into_bytes());
wire.set_header("content-type", "application/x-www-form-urlencoded");
wire
}
pub fn parse_callback(query: &str, expected_state: &str) -> Result<Callback, AuthError> {
let (mut code, mut state, mut error) = (None, None, None);
for (key, value) in query_pairs(query) {
match key.as_str() {
"code" => code = Some(value),
"state" => state = Some(value),
"error" => error = Some(value),
_ => {}
}
}
if let Some(err) = error {
return Err(AuthError::Fatal(format!("authorization denied: {err}")));
}
let code = code.ok_or_else(|| AuthError::Fatal("callback missing `code`".to_owned()))?;
let state = state.ok_or_else(|| AuthError::Fatal("callback missing `state`".to_owned()))?;
if state != expected_state {
return Err(AuthError::Fatal(
"callback `state` mismatch (possible CSRF)".to_owned(),
));
}
Ok(Callback { code, state })
}
pub fn query_from_request_line(line: &str) -> Option<String> {
let target = line.split(' ').nth(1)?;
target.split_once('?').map(|(_, query)| query.to_owned())
}