use core::fmt;
use crate::crypto::constant_time::constant_time_eq;
use crate::util::validation::{is_valid_client_id, is_valid_redirect_uri, is_valid_scope};
#[doc(alias = "client_type")]
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[non_exhaustive]
pub enum ClientType {
Confidential,
Public,
}
impl ClientType {
#[must_use]
#[inline]
pub fn is_confidential(self) -> bool {
matches!(self, Self::Confidential)
}
#[must_use]
#[inline]
pub fn is_public(self) -> bool {
matches!(self, Self::Public)
}
}
#[doc(alias = "oauth_client")]
#[derive(Debug, Clone)]
pub struct RegisteredClient {
client_id: String,
client_type: ClientType,
redirect_uris: Vec<String>,
allowed_scopes: Vec<String>,
require_pkce: bool,
active: bool,
}
impl RegisteredClient {
pub fn builder(
client_id: impl Into<String>,
client_type: ClientType,
) -> RegisteredClientBuilder {
RegisteredClientBuilder {
client_id: client_id.into(),
client_type,
redirect_uris: Vec::new(),
allowed_scopes: Vec::new(),
require_pkce: true,
active: true,
}
}
#[must_use]
#[inline]
pub fn client_id(&self) -> &str {
&self.client_id
}
#[must_use]
#[inline]
pub fn client_type(&self) -> ClientType {
self.client_type
}
#[must_use]
#[inline]
pub fn redirect_uris(&self) -> &[String] {
&self.redirect_uris
}
#[must_use]
#[inline]
pub fn allowed_scopes(&self) -> &[String] {
&self.allowed_scopes
}
#[must_use]
#[inline]
pub fn require_pkce(&self) -> bool {
self.require_pkce
}
#[must_use]
#[inline]
pub fn active(&self) -> bool {
self.active
}
#[must_use]
pub fn is_registered_redirect_uri(&self, uri: &str) -> bool {
self.redirect_uris
.iter()
.any(|registered| constant_time_eq(registered.as_bytes(), uri.as_bytes()))
}
#[must_use]
pub fn allows_scopes(&self, requested: &str) -> bool {
requested
.split(' ')
.filter(|s| !s.is_empty())
.all(|s| self.allowed_scopes.iter().any(|allowed| allowed == s))
}
}
#[doc(alias = "client_builder")]
#[derive(Debug, Clone)]
#[must_use = "a builder does nothing until `build` is called"]
pub struct RegisteredClientBuilder {
client_id: String,
client_type: ClientType,
redirect_uris: Vec<String>,
allowed_scopes: Vec<String>,
require_pkce: bool,
active: bool,
}
impl RegisteredClientBuilder {
#[inline]
pub fn redirect_uri(mut self, uri: impl Into<String>) -> Self {
self.redirect_uris.push(uri.into());
self
}
#[inline]
pub fn allowed_scope(mut self, scope: impl Into<String>) -> Self {
self.allowed_scopes.push(scope.into());
self
}
#[inline]
pub fn require_pkce(mut self, required: bool) -> Self {
self.require_pkce = required;
self
}
#[inline]
pub fn active(mut self, active: bool) -> Self {
self.active = active;
self
}
pub fn build(self) -> Result<RegisteredClient, RegisteredClientError> {
if !is_valid_client_id(&self.client_id) {
return Err(RegisteredClientError {
kind: RegisteredClientErrorKind::InvalidClientId,
});
}
if !self.redirect_uris.iter().all(|u| is_valid_redirect_uri(u)) {
return Err(RegisteredClientError {
kind: RegisteredClientErrorKind::InvalidRedirectUri,
});
}
if !self.allowed_scopes.iter().all(|s| is_valid_scope(s)) {
return Err(RegisteredClientError {
kind: RegisteredClientErrorKind::InvalidScope,
});
}
Ok(RegisteredClient {
client_id: self.client_id,
client_type: self.client_type,
redirect_uris: self.redirect_uris,
allowed_scopes: self.allowed_scopes,
require_pkce: self.require_pkce,
active: self.active,
})
}
}
#[allow(clippy::enum_variant_names)]
#[derive(Debug, Clone, PartialEq, Eq)]
enum RegisteredClientErrorKind {
InvalidClientId,
InvalidRedirectUri,
InvalidScope,
}
#[doc(alias = "registered_client_error")]
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RegisteredClientError {
kind: RegisteredClientErrorKind,
}
impl RegisteredClientError {
#[must_use]
#[inline]
pub fn is_invalid_client_id(&self) -> bool {
self.kind == RegisteredClientErrorKind::InvalidClientId
}
#[must_use]
#[inline]
pub fn is_invalid_redirect_uri(&self) -> bool {
self.kind == RegisteredClientErrorKind::InvalidRedirectUri
}
#[must_use]
#[inline]
pub fn is_invalid_scope(&self) -> bool {
self.kind == RegisteredClientErrorKind::InvalidScope
}
}
impl fmt::Display for RegisteredClientError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let msg = match self.kind {
RegisteredClientErrorKind::InvalidClientId => "invalid client_id",
RegisteredClientErrorKind::InvalidRedirectUri => {
"invalid redirect URI (must be HTTPS or loopback http, no fragment)"
}
RegisteredClientErrorKind::InvalidScope => "invalid scope",
};
write!(f, "registered client: {msg}")
}
}
impl std::error::Error for RegisteredClientError {}
#[cfg(test)]
mod tests {
use super::*;
fn client() -> RegisteredClient {
RegisteredClient::builder("c1", ClientType::Confidential)
.redirect_uri("https://app.example.com/callback")
.redirect_uri("https://app.example.com/callback2")
.allowed_scope("openid")
.allowed_scope("profile")
.allowed_scope("email")
.build()
.expect("valid test client")
}
#[test]
fn accessors() {
let c = client();
assert_eq!(c.client_id(), "c1");
assert_eq!(c.client_type(), ClientType::Confidential);
assert!(c.client_type().is_confidential());
assert!(!c.client_type().is_public());
assert_eq!(c.redirect_uris().len(), 2);
assert_eq!(c.allowed_scopes(), &["openid", "profile", "email"]);
assert!(c.require_pkce());
assert!(c.active());
}
#[test]
fn redirect_uri_exact_match() {
let c = client();
assert!(c.is_registered_redirect_uri("https://app.example.com/callback"));
assert!(c.is_registered_redirect_uri("https://app.example.com/callback2"));
}
#[test]
fn redirect_uri_rejects_non_exact() {
let c = client();
assert!(!c.is_registered_redirect_uri("https://app.example.com/callback/"));
assert!(!c.is_registered_redirect_uri("https://app.example.com/callbac"));
assert!(!c.is_registered_redirect_uri("https://app.example.com/callback?x=1"));
assert!(!c.is_registered_redirect_uri("https://evil.example.com/callback"));
assert!(!c.is_registered_redirect_uri(""));
}
#[test]
fn scope_subset() {
let c = client();
assert!(c.allows_scopes("openid"));
assert!(c.allows_scopes("openid profile"));
assert!(c.allows_scopes("openid profile email"));
assert!(c.allows_scopes("")); }
#[test]
fn scope_superset_rejected() {
let c = client();
assert!(!c.allows_scopes("openid admin"));
assert!(!c.allows_scopes("offline_access"));
}
#[test]
fn scope_handles_extra_whitespace() {
let c = client();
assert!(c.allows_scopes("openid profile"));
assert!(c.allows_scopes(" openid "));
}
#[test]
fn builder_flags() {
let public = RegisteredClient::builder("spa", ClientType::Public)
.redirect_uri("https://spa.example.com/cb")
.require_pkce(false)
.active(false)
.build()
.expect("valid test client");
assert!(public.client_type().is_public());
assert!(!public.require_pkce());
assert!(!public.active());
}
#[test]
fn build_allows_no_redirect_uri_for_m2m_client() {
let client = RegisteredClient::builder("svc_c1", ClientType::Confidential)
.allowed_scope("read")
.require_pkce(false)
.build()
.expect("m2m client builds without redirect URIs");
assert!(client.redirect_uris().is_empty());
assert!(!client.is_registered_redirect_uri("https://app.example.com/cb"));
}
#[test]
fn build_rejects_plaintext_http_redirect() {
let err = RegisteredClient::builder("c1", ClientType::Confidential)
.redirect_uri("http://evil.example.com/cb")
.build()
.unwrap_err();
assert!(err.is_invalid_redirect_uri());
}
#[test]
fn build_rejects_redirect_with_fragment() {
let err = RegisteredClient::builder("c1", ClientType::Confidential)
.redirect_uri("https://app.example.com/cb#frag")
.build()
.unwrap_err();
assert!(err.is_invalid_redirect_uri());
}
#[test]
fn build_rejects_invalid_client_id() {
let err = RegisteredClient::builder("", ClientType::Confidential)
.redirect_uri("https://app.example.com/cb")
.build()
.unwrap_err();
assert!(err.is_invalid_client_id());
}
#[test]
fn build_allows_loopback_http_redirect() {
RegisteredClient::builder("native", ClientType::Public)
.redirect_uri("http://127.0.0.1:8080/cb")
.require_pkce(true)
.build()
.expect("loopback http is valid");
}
}