use std::borrow::Cow;
use std::fmt;
use std::time::SystemTime;
use serde::{Deserialize, Serialize};
use crate::client::ClientId;
use crate::error::{ErrorCode, ErrorResponse};
use crate::scope::ScopeSet;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum ResponseType {
#[serde(rename = "code")]
Code,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum CodeChallengeMethod {
S256,
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
#[non_exhaustive]
pub struct AuthorizationRequest<'a> {
pub response_type: Option<Cow<'a, str>>,
pub client_id: Option<Cow<'a, str>>,
pub redirect_uri: Option<Cow<'a, str>>,
pub scope: Option<Cow<'a, str>>,
pub state: Option<Cow<'a, str>>,
pub code_challenge: Option<Cow<'a, str>>,
pub code_challenge_method: Option<Cow<'a, str>>,
pub resource: Vec<Cow<'a, str>>,
pub authorization_details: Option<Cow<'a, str>>,
#[cfg(feature = "consent")]
pub acr_values: Option<Cow<'a, str>>,
#[cfg(feature = "consent")]
pub max_age: Option<Cow<'a, str>>,
}
impl<'a> AuthorizationRequest<'a> {
pub fn from_pairs<I, K, V>(pairs: I) -> Self
where
I: IntoIterator<Item = (K, V)>,
K: AsRef<str>,
V: Into<Cow<'a, str>>,
{
let mut req = AuthorizationRequest::default();
for (k, v) in pairs {
if k.as_ref() == "resource" {
req.resource.push(v.into());
continue;
}
let slot = match k.as_ref() {
"response_type" => &mut req.response_type,
"client_id" => &mut req.client_id,
"redirect_uri" => &mut req.redirect_uri,
"scope" => &mut req.scope,
"state" => &mut req.state,
"code_challenge" => &mut req.code_challenge,
"code_challenge_method" => &mut req.code_challenge_method,
"authorization_details" => &mut req.authorization_details,
#[cfg(feature = "consent")]
"acr_values" => &mut req.acr_values,
#[cfg(feature = "consent")]
"max_age" => &mut req.max_age,
_ => continue,
};
if slot.is_none() {
*slot = Some(v.into());
}
}
req
}
}
pub(crate) fn is_valid_resource_indicator(value: &str) -> bool {
let bytes = value.as_bytes();
let colon = match value.find(':') {
Some(0) | None => return false,
Some(i) => i,
};
if !bytes[0].is_ascii_alphabetic() {
return false;
}
if !bytes[1..colon]
.iter()
.all(|b| b.is_ascii_alphanumeric() || matches!(b, b'+' | b'-' | b'.'))
{
return false;
}
bytes
.iter()
.all(|&b| b != b'#' && (0x21..=0x7e).contains(&b))
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ValidatedAuthorizationRequest {
pub client_id: ClientId,
pub redirect_uri: String,
pub redirect_uri_was_explicit: bool,
pub scope: ScopeSet,
pub state: Option<String>,
pub code_challenge: String,
pub code_challenge_method: CodeChallengeMethod,
pub issuer: String,
pub resource: Vec<String>,
#[cfg(feature = "rar")]
pub authorization_details: crate::rar::AuthorizationDetails,
#[cfg(feature = "consent")]
pub authentication_requirement: crate::consent::AuthenticationRequirement,
_sealed: Sealed,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct Sealed;
impl ValidatedAuthorizationRequest {
#[allow(clippy::too_many_arguments)]
pub(crate) fn new(
client_id: ClientId,
redirect_uri: String,
redirect_uri_was_explicit: bool,
scope: ScopeSet,
state: Option<String>,
code_challenge: String,
code_challenge_method: CodeChallengeMethod,
issuer: String,
resource: Vec<String>,
) -> Self {
ValidatedAuthorizationRequest {
client_id,
redirect_uri,
redirect_uri_was_explicit,
scope,
state,
code_challenge,
code_challenge_method,
issuer,
resource,
#[cfg(feature = "rar")]
authorization_details: crate::rar::AuthorizationDetails::none(),
#[cfg(feature = "consent")]
authentication_requirement: crate::consent::AuthenticationRequirement::none(),
_sealed: Sealed,
}
}
#[cfg(feature = "consent")]
pub(crate) fn set_authentication_requirement(
&mut self,
requirement: crate::consent::AuthenticationRequirement,
) {
self.authentication_requirement = requirement;
}
#[cfg(feature = "rar")]
pub(crate) fn set_authorization_details(&mut self, details: crate::rar::AuthorizationDetails) {
self.authorization_details = details;
}
pub fn denied(&self) -> AuthorizationErrorRedirect {
AuthorizationErrorRedirect {
redirect_uri: self.redirect_uri.clone(),
error: ErrorResponse::new(ErrorCode::AccessDenied),
state: self.state.clone(),
iss: self.issuer.clone(),
}
}
}
impl fmt::Debug for AuthorizationResponse {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("AuthorizationResponse")
.field("code", &"[redacted]")
.field("state", &self.state)
.field("iss", &self.iss)
.finish()
}
}
#[derive(Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct AuthorizationResponse {
pub code: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub state: Option<String>,
pub iss: String,
}
impl AuthorizationResponse {
pub fn location(&self, redirect_uri: &str) -> String {
let mut out = String::with_capacity(redirect_uri.len() + self.encoded_len());
out.push_str(redirect_uri);
let mut sep = query_separator(redirect_uri);
append_param(&mut out, &mut sep, "code", &self.code);
if let Some(state) = &self.state {
append_param(&mut out, &mut sep, "state", state);
}
append_param(&mut out, &mut sep, "iss", &self.iss);
out
}
fn encoded_len(&self) -> usize {
6 + self.code.len() * 3
+ self.state.as_ref().map_or(0, |s| 7 + s.len() * 3)
+ 5
+ self.iss.len() * 3
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct AuthorizationErrorRedirect {
pub redirect_uri: String,
pub error: ErrorResponse,
pub state: Option<String>,
pub iss: String,
}
impl AuthorizationErrorRedirect {
pub fn location(&self) -> String {
let mut out = String::with_capacity(self.redirect_uri.len() + 96 + self.iss.len() * 3);
out.push_str(&self.redirect_uri);
let mut sep = query_separator(&self.redirect_uri);
append_param(&mut out, &mut sep, "error", self.error.error.as_str());
if let Some(d) = &self.error.error_description {
append_param(&mut out, &mut sep, "error_description", d);
}
if let Some(u) = &self.error.error_uri {
append_param(&mut out, &mut sep, "error_uri", u);
}
if let Some(state) = &self.state {
append_param(&mut out, &mut sep, "state", state);
}
append_param(&mut out, &mut sep, "iss", &self.iss);
out
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum AuthorizationError {
Direct(ErrorResponse),
Redirect(AuthorizationErrorRedirect),
}
impl AuthorizationError {
pub fn http_status(&self) -> u16 {
match self {
AuthorizationError::Direct(e) if e.error == ErrorCode::ServerError => 500,
AuthorizationError::Direct(_) => 400,
AuthorizationError::Redirect(_) => 302,
}
}
}
impl std::fmt::Display for AuthorizationError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
AuthorizationError::Direct(e) => write!(f, "{e}"),
AuthorizationError::Redirect(r) => write!(f, "{}", r.error),
}
}
}
impl std::error::Error for AuthorizationError {}
#[derive(Clone, PartialEq, Eq, Serialize, Deserialize)]
pub enum AuthorizationCodeState {
Issued,
Consumed {
access_token: Option<String>,
refresh_token: Option<String>,
},
Replayed {
access_token: Option<String>,
refresh_token: Option<String>,
},
}
impl AuthorizationCodeState {
pub fn minted(&self) -> Option<(Option<&str>, Option<&str>)> {
match self {
AuthorizationCodeState::Issued => None,
AuthorizationCodeState::Consumed {
access_token,
refresh_token,
}
| AuthorizationCodeState::Replayed {
access_token,
refresh_token,
} => Some((access_token.as_deref(), refresh_token.as_deref())),
}
}
}
impl fmt::Debug for AuthorizationCodeState {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
AuthorizationCodeState::Issued => f.write_str("Issued"),
AuthorizationCodeState::Consumed {
access_token,
refresh_token,
}
| AuthorizationCodeState::Replayed {
access_token,
refresh_token,
} => f
.debug_struct(match self {
AuthorizationCodeState::Replayed { .. } => "Replayed",
_ => "Consumed",
})
.field("access_token", &access_token.as_ref().map(|_| "[redacted]"))
.field(
"refresh_token",
&refresh_token.as_ref().map(|_| "[redacted]"),
)
.finish(),
}
}
}
fn redirect_uri_was_explicit_default() -> bool {
true
}
fn grant_instant_default() -> SystemTime {
SystemTime::UNIX_EPOCH
}
#[derive(Clone, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct AuthorizationCodeRecord {
pub code: String,
pub client_id: ClientId,
pub redirect_uri: String,
#[serde(default = "redirect_uri_was_explicit_default")]
pub redirect_uri_was_explicit: bool,
pub scope: ScopeSet,
pub subject: String,
pub code_challenge: String,
pub code_challenge_method: CodeChallengeMethod,
pub resource: Vec<String>,
#[cfg(feature = "rar")]
#[serde(default)]
pub authorization_details: crate::rar::AuthorizationDetails,
#[serde(default = "grant_instant_default")]
pub issued_at: SystemTime,
pub expires_at: SystemTime,
pub state: AuthorizationCodeState,
#[cfg(feature = "consent")]
pub authentication: Option<Box<crate::consent::Authentication>>,
}
impl AuthorizationCodeRecord {
#[allow(clippy::too_many_arguments)]
pub fn new(
code: impl Into<String>,
client_id: ClientId,
redirect_uri: impl Into<String>,
scope: ScopeSet,
subject: impl Into<String>,
code_challenge: impl Into<String>,
expires_at: SystemTime,
) -> Self {
AuthorizationCodeRecord {
issued_at: SystemTime::UNIX_EPOCH,
code: code.into(),
client_id,
redirect_uri: redirect_uri.into(),
redirect_uri_was_explicit: true,
scope,
subject: subject.into(),
code_challenge: code_challenge.into(),
code_challenge_method: CodeChallengeMethod::S256,
resource: Vec::new(),
#[cfg(feature = "rar")]
authorization_details: crate::rar::AuthorizationDetails::none(),
expires_at,
state: AuthorizationCodeState::Issued,
#[cfg(feature = "consent")]
authentication: None,
}
}
}
impl fmt::Debug for AuthorizationCodeRecord {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let mut out = f.debug_struct("AuthorizationCodeRecord");
out.field("code", &"[redacted]")
.field("client_id", &self.client_id)
.field("redirect_uri", &self.redirect_uri)
.field("redirect_uri_was_explicit", &self.redirect_uri_was_explicit)
.field("scope", &self.scope)
.field("subject", &self.subject)
.field("code_challenge", &self.code_challenge)
.field("code_challenge_method", &self.code_challenge_method)
.field("resource", &self.resource);
#[cfg(feature = "rar")]
out.field("authorization_details", &self.authorization_details);
out.field("issued_at", &self.issued_at)
.field("expires_at", &self.expires_at)
.field("state", &self.state);
#[cfg(feature = "consent")]
out.field("authentication", &self.authentication);
out.finish()
}
}
pub(crate) fn query_separator(uri: &str) -> char {
if uri.contains('?') {
'&'
} else {
'?'
}
}
fn append_param(out: &mut String, sep: &mut char, name: &str, value: &str) {
out.push(*sep);
*sep = '&';
out.push_str(name);
out.push('=');
percent_encode_into(out, value);
}
fn percent_encode_into(out: &mut String, value: &str) {
const HEX: &[u8; 16] = b"0123456789ABCDEF";
for &b in value.as_bytes() {
if b.is_ascii_alphanumeric() || matches!(b, b'-' | b'.' | b'_' | b'~') {
out.push(b as char);
} else {
out.push('%');
out.push(HEX[(b >> 4) as usize] as char);
out.push(HEX[(b & 0x0f) as usize] as char);
}
}
}
#[cfg(test)]
#[path = "tests/authorization.rs"]
mod tests;