use crate::{
errors::{Error, Result},
types::{AuthUrlResult, BuildAuthUrl, CallbackParams, EndSession, OidcProviderMetadata},
};
use rand::{distributions::Alphanumeric, Rng};
#[cfg(test)]
use std::collections::HashMap;
use url::Url;
pub fn build_auth_url(params: BuildAuthUrl) -> Result<AuthUrlResult> {
if params.issuer.is_empty() {
return Err(Error::InvalidParam("issuer cannot be empty"));
}
if params.client_id.is_empty() {
return Err(Error::InvalidParam("client_id cannot be empty"));
}
if params.redirect_uri.is_empty() {
return Err(Error::InvalidParam("redirect_uri cannot be empty"));
}
if params.code_challenge.is_empty() {
return Err(Error::InvalidParam("code_challenge cannot be empty"));
}
let auth_endpoint = params
.authorization_endpoint
.ok_or_else(|| Error::InvalidParam("authorization_endpoint is required. Please use OIDC discovery to obtain the correct endpoint"))?;
let mut url = Url::parse(&auth_endpoint)?;
let scope = if params.scope.is_empty() { "openid profile email" } else { ¶ms.scope };
let state = params.state.unwrap_or_else(generate_state);
let nonce = if scope.contains("openid") {
Some(params.nonce.unwrap_or_else(generate_nonce))
} else {
None
};
{
let mut query = url.query_pairs_mut();
query.append_pair("response_type", "code");
query.append_pair("client_id", ¶ms.client_id);
query.append_pair("redirect_uri", ¶ms.redirect_uri);
query.append_pair("scope", scope);
query.append_pair("state", &state);
query.append_pair("code_challenge", ¶ms.code_challenge);
query.append_pair("code_challenge_method", "S256");
if let Some(ref nonce) = nonce {
query.append_pair("nonce", nonce);
}
if let Some(prompt) = ¶ms.prompt {
query.append_pair("prompt", prompt);
}
if let Some(tenant) = ¶ms.tenant {
query.append_pair("tenant", tenant);
}
if let Some(extra) = ¶ms.extra_params {
for (key, value) in extra {
query.append_pair(key, value);
}
}
}
Ok(AuthUrlResult { url, state, nonce })
}
pub fn build_end_session_url(params: EndSession) -> Result<Url> {
if params.issuer.is_empty() {
return Err(Error::InvalidParam("issuer cannot be empty"));
}
if params.id_token_hint.is_empty() {
return Err(Error::InvalidParam("id_token_hint cannot be empty"));
}
let end_session_endpoint = if let Some(endpoint) = ¶ms.end_session_endpoint {
endpoint.clone()
} else {
if params.issuer.ends_with('/') {
format!("{}oidc/end_session", params.issuer)
} else {
format!("{}/oidc/end_session", params.issuer)
}
};
let mut url = Url::parse(&end_session_endpoint)?;
{
let mut query = url.query_pairs_mut();
query.append_pair("id_token_hint", ¶ms.id_token_hint);
if let Some(redirect_uri) = ¶ms.post_logout_redirect_uri {
query.append_pair("post_logout_redirect_uri", redirect_uri);
}
if let Some(state) = ¶ms.state {
query.append_pair("state", state);
}
}
Ok(url)
}
pub fn build_end_session_url_with_discovery(
mut params: EndSession,
metadata: &OidcProviderMetadata,
) -> Result<Url> {
if params.end_session_endpoint.is_none() {
params.end_session_endpoint = metadata.end_session_endpoint.clone();
}
build_end_session_url(params)
}
pub fn parse_callback_params(url: &str) -> CallbackParams {
let mut params =
CallbackParams { code: None, state: None, error: None, error_description: None };
if let Ok(parsed_url) = Url::parse(url) {
for (key, value) in parsed_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()),
"error_description" => params.error_description = Some(value.into_owned()),
_ => {} }
}
} else {
let query = if let Some(query_start) = url.find('?') {
&url[query_start + 1..]
} else if url.contains('=') {
url
} else {
""
};
if !query.is_empty() {
for pair in query.split('&') {
if let Some(eq_pos) = pair.find('=') {
let key = &pair[..eq_pos];
let value = &pair[eq_pos + 1..];
let decoded_value = urlencoding::decode(value).unwrap_or_else(|_| value.into());
match key {
"code" => params.code = Some(decoded_value.into_owned()),
"state" => params.state = Some(decoded_value.into_owned()),
"error" => params.error = Some(decoded_value.into_owned()),
"error_description" => {
params.error_description = Some(decoded_value.into_owned())
}
_ => {} }
}
}
}
}
params
}
#[allow(dead_code)]
pub fn build_auth_url_with_metadata(
metadata: &OidcProviderMetadata,
params: BuildAuthUrl,
) -> Result<AuthUrlResult> {
let mut url = Url::parse(&metadata.authorization_endpoint)?;
let state = params.state.unwrap_or_else(generate_state);
let nonce = if params.scope.contains("openid") {
Some(params.nonce.unwrap_or_else(generate_nonce))
} else {
None
};
{
let mut query = url.query_pairs_mut();
query.append_pair("response_type", "code");
query.append_pair("client_id", ¶ms.client_id);
query.append_pair("redirect_uri", ¶ms.redirect_uri);
query.append_pair("scope", ¶ms.scope);
query.append_pair("state", &state);
query.append_pair("code_challenge", ¶ms.code_challenge);
query.append_pair("code_challenge_method", "S256");
if let Some(ref nonce) = nonce {
query.append_pair("nonce", nonce);
}
if let Some(prompt) = ¶ms.prompt {
query.append_pair("prompt", prompt);
}
if let Some(tenant) = ¶ms.tenant {
query.append_pair("tenant", tenant);
}
if let Some(extra) = ¶ms.extra_params {
for (key, value) in extra {
query.append_pair(key, value);
}
}
}
Ok(AuthUrlResult { url, state, nonce })
}
fn generate_state() -> String {
rand::thread_rng().sample_iter(&Alphanumeric).take(32).map(char::from).collect()
}
fn generate_nonce() -> String {
rand::thread_rng().sample_iter(&Alphanumeric).take(32).map(char::from).collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_build_auth_url() {
let result = build_auth_url(BuildAuthUrl {
issuer: "https://auth.example.com".into(),
client_id: "test-client".into(),
redirect_uri: "https://app.example.com/callback".into(),
scope: "openid profile".into(),
code_challenge: "test_challenge".into(),
state: Some("test_state".into()),
nonce: Some("test_nonce".into()),
prompt: None,
extra_params: None,
tenant: None,
authorization_endpoint: Some("https://auth.example.com/oauth/authorize".into()),
})
.unwrap();
let url = result.url;
assert_eq!(result.state, "test_state");
assert_eq!(result.nonce, Some("test_nonce".to_string()));
let query: HashMap<_, _> = url.query_pairs().into_owned().collect();
assert_eq!(query.get("response_type"), Some(&"code".to_string()));
assert_eq!(query.get("client_id"), Some(&"test-client".to_string()));
assert_eq!(
query.get("redirect_uri"),
Some(&"https://app.example.com/callback".to_string())
);
assert_eq!(query.get("scope"), Some(&"openid profile".to_string()));
assert_eq!(query.get("state"), Some(&"test_state".to_string()));
assert_eq!(query.get("nonce"), Some(&"test_nonce".to_string()));
assert_eq!(query.get("code_challenge"), Some(&"test_challenge".to_string()));
assert_eq!(query.get("code_challenge_method"), Some(&"S256".to_string()));
}
#[test]
fn test_build_auth_url_auto_state_nonce() {
let result = build_auth_url(BuildAuthUrl {
issuer: "https://auth.example.com".into(),
client_id: "test-client".into(),
redirect_uri: "https://app.example.com/callback".into(),
scope: "openid profile".into(),
code_challenge: "test_challenge".into(),
state: None,
nonce: None,
prompt: None,
extra_params: None,
tenant: None,
authorization_endpoint: Some("https://auth.example.com/oauth/authorize".into()),
})
.unwrap();
let url = result.url;
assert_eq!(result.state.len(), 32);
assert_eq!(result.nonce.as_ref().unwrap().len(), 32);
let query: HashMap<_, _> = url.query_pairs().into_owned().collect();
assert!(query.contains_key("state"));
assert!(query.contains_key("nonce"));
assert_eq!(query.get("state").unwrap().len(), 32);
assert_eq!(query.get("nonce").unwrap().len(), 32);
}
#[test]
fn test_build_auth_url_missing_authorization_endpoint() {
let result = build_auth_url(BuildAuthUrl {
issuer: "https://auth.example.com".into(),
client_id: "test-client".into(),
redirect_uri: "https://app.example.com/callback".into(),
scope: "openid profile".into(),
code_challenge: "test_challenge".into(),
state: Some("test_state".into()),
nonce: Some("test_nonce".into()),
prompt: None,
extra_params: None,
tenant: None,
authorization_endpoint: None,
});
assert!(result.is_err());
match result {
Err(Error::InvalidParam(msg)) => {
assert!(msg.contains("authorization_endpoint is required"));
}
_ => panic!("Expected InvalidParam error"),
}
}
#[test]
fn test_parse_callback_params() {
let params =
parse_callback_params("https://app.example.com/callback?code=abc123&state=xyz456");
assert_eq!(params.code, Some("abc123".to_string()));
assert_eq!(params.state, Some("xyz456".to_string()));
assert_eq!(params.error, None);
assert_eq!(params.error_description, None);
}
#[test]
fn test_parse_callback_params_error() {
let params = parse_callback_params(
"https://app.example.com/callback?error=access_denied&error_description=User%20denied%20access"
);
assert_eq!(params.code, None);
assert_eq!(params.state, None);
assert_eq!(params.error, Some("access_denied".to_string()));
assert_eq!(params.error_description, Some("User denied access".to_string()));
}
#[test]
fn test_parse_callback_params_relative_url() {
let params = parse_callback_params("/callback?code=test&state=test");
assert_eq!(params.code, Some("test".to_string()));
assert_eq!(params.state, Some("test".to_string()));
}
#[test]
fn test_build_end_session_url() {
let url = build_end_session_url(EndSession {
issuer: "https://auth.example.com".into(),
id_token_hint: "test_token".into(),
post_logout_redirect_uri: Some("https://app.example.com".into()),
state: Some("logout_state".into()),
end_session_endpoint: None,
})
.unwrap();
let query: HashMap<_, _> = url.query_pairs().into_owned().collect();
assert_eq!(query.get("id_token_hint"), Some(&"test_token".to_string()));
assert_eq!(
query.get("post_logout_redirect_uri"),
Some(&"https://app.example.com".to_string())
);
assert_eq!(query.get("state"), Some(&"logout_state".to_string()));
}
}