Import upstream v0.16.22, stripped
Upstream commit: 474dd0229cb20cf513036619781ed97bd8073c3f Enterprise-only files removed or emptied: 63 Enterprise-only snippets removed: 117 in 50 files Dangling module declarations removed: 5 Cargo edits turning enterprise off: 14 Verification: clean Enterprise feature gates left for rebuilt features: 19 in 18 files Produced by tools/fork/strip.py. The full report is in docs/fork/strip-reports/ on main.
This commit is contained in:
@@ -0,0 +1,225 @@
|
||||
/*
|
||||
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
|
||||
*
|
||||
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
|
||||
*/
|
||||
|
||||
use super::{
|
||||
SCOPE_CALENDARS, SCOPE_CONTACTS, SCOPE_MAIL, SCOPE_OFFLINE_ACCESS, SCOPE_OPENID,
|
||||
crypto::SymmetricEncrypt,
|
||||
};
|
||||
use base64::{Engine, engine::general_purpose};
|
||||
use store::blake3;
|
||||
use utils::codec::leb128::{Leb128Iterator, Leb128Vec};
|
||||
|
||||
const CLIENT_ID_HEADER: &str = "swc1.";
|
||||
const CLIENT_ID_KEY_CONTEXT: &str = "stalwart-oauth-client-id-sw1";
|
||||
const CLIENT_ID_VERSION: u8 = 1;
|
||||
|
||||
const SCOPE_BITS: &[&str] = &[
|
||||
SCOPE_OPENID,
|
||||
SCOPE_OFFLINE_ACCESS,
|
||||
SCOPE_MAIL,
|
||||
SCOPE_CONTACTS,
|
||||
SCOPE_CALENDARS,
|
||||
];
|
||||
|
||||
pub fn scopes_to_mask(scope: &str) -> u64 {
|
||||
let mut mask = 0u64;
|
||||
for scope in scope.split_ascii_whitespace() {
|
||||
if let Some(bit) = SCOPE_BITS.iter().position(|known| *known == scope) {
|
||||
mask |= 1 << bit;
|
||||
}
|
||||
}
|
||||
mask
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Default)]
|
||||
pub struct ClientMeta {
|
||||
pub redirect_uris: Vec<String>,
|
||||
pub scope_mask: u64,
|
||||
pub client_name: Option<String>,
|
||||
}
|
||||
|
||||
pub fn encode_client_id(key: &[u8], meta: &ClientMeta) -> Result<String, String> {
|
||||
let client_name = meta.client_name.as_deref().unwrap_or_default();
|
||||
|
||||
let mut payload = Vec::with_capacity(
|
||||
24 + meta
|
||||
.redirect_uris
|
||||
.iter()
|
||||
.map(|u| u.len() + 2)
|
||||
.sum::<usize>()
|
||||
+ client_name.len(),
|
||||
);
|
||||
payload.push(CLIENT_ID_VERSION);
|
||||
payload.push_leb128(meta.redirect_uris.len());
|
||||
for uri in &meta.redirect_uris {
|
||||
payload.push_leb128(uri.len());
|
||||
payload.extend_from_slice(uri.as_bytes());
|
||||
}
|
||||
payload.push_leb128(meta.scope_mask);
|
||||
payload.push_leb128(client_name.len());
|
||||
payload.extend_from_slice(client_name.as_bytes());
|
||||
|
||||
let digest = blake3::hash(&payload);
|
||||
let nonce = &digest.as_bytes()[..SymmetricEncrypt::NONCE_LEN];
|
||||
let ciphertext =
|
||||
SymmetricEncrypt::new(key, CLIENT_ID_KEY_CONTEXT).encrypt_with_aad(&payload, nonce, &[])?;
|
||||
|
||||
let mut body = Vec::with_capacity(nonce.len() + ciphertext.len());
|
||||
body.extend_from_slice(nonce);
|
||||
body.extend_from_slice(&ciphertext);
|
||||
|
||||
let mut out = String::with_capacity(CLIENT_ID_HEADER.len() + body.len().div_ceil(3) * 4);
|
||||
out.push_str(CLIENT_ID_HEADER);
|
||||
general_purpose::URL_SAFE_NO_PAD.encode_string(&body, &mut out);
|
||||
|
||||
Ok(out)
|
||||
}
|
||||
|
||||
pub fn decode_client_id(key: &[u8], client_id: &str) -> Option<ClientMeta> {
|
||||
let body = general_purpose::URL_SAFE_NO_PAD
|
||||
.decode(client_id.strip_prefix(CLIENT_ID_HEADER)?.as_bytes())
|
||||
.ok()?;
|
||||
if body.len() < SymmetricEncrypt::NONCE_LEN + SymmetricEncrypt::ENCRYPT_TAG_LEN {
|
||||
return None;
|
||||
}
|
||||
let (nonce, ciphertext) = body.split_at(SymmetricEncrypt::NONCE_LEN);
|
||||
let payload = SymmetricEncrypt::new(key, CLIENT_ID_KEY_CONTEXT)
|
||||
.decrypt_with_aad(ciphertext, nonce, &[])
|
||||
.ok()?;
|
||||
|
||||
let mut bytes = payload.iter();
|
||||
if bytes.next().copied()? != CLIENT_ID_VERSION {
|
||||
return None;
|
||||
}
|
||||
|
||||
let uri_count: usize = bytes.next_leb128()?;
|
||||
if uri_count > u8::MAX as usize {
|
||||
return None;
|
||||
}
|
||||
let mut redirect_uris = Vec::with_capacity(uri_count);
|
||||
for _ in 0..uri_count {
|
||||
redirect_uris.push(take_string(&mut bytes)?);
|
||||
}
|
||||
let scope_mask: u64 = bytes.next_leb128()?;
|
||||
let client_name = take_string(&mut bytes)?;
|
||||
|
||||
Some(ClientMeta {
|
||||
redirect_uris,
|
||||
scope_mask,
|
||||
client_name: (!client_name.is_empty()).then_some(client_name),
|
||||
})
|
||||
}
|
||||
|
||||
fn take_string(bytes: &mut std::slice::Iter<'_, u8>) -> Option<String> {
|
||||
let len: usize = bytes.next_leb128()?;
|
||||
let slice = bytes.as_slice();
|
||||
if slice.len() < len {
|
||||
return None;
|
||||
}
|
||||
let value = String::from_utf8(slice[..len].to_vec()).ok()?;
|
||||
if len > 0 {
|
||||
bytes.nth(len - 1)?;
|
||||
}
|
||||
Some(value)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
const KEY: &[u8] = b"a-test-encryption-key-of-some-length";
|
||||
|
||||
fn sample() -> ClientMeta {
|
||||
ClientMeta {
|
||||
redirect_uris: vec![
|
||||
"http://127.0.0.1/cb".to_string(),
|
||||
"com.example.app:/oauth".to_string(),
|
||||
],
|
||||
scope_mask: scopes_to_mask(&format!("{SCOPE_OFFLINE_ACCESS} {SCOPE_MAIL}")),
|
||||
client_name: Some("Example Client".to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn round_trip_preserves_all_fields() {
|
||||
for meta in [
|
||||
sample(),
|
||||
ClientMeta {
|
||||
redirect_uris: vec!["http://[::1]/".to_string()],
|
||||
scope_mask: 0,
|
||||
client_name: None,
|
||||
},
|
||||
ClientMeta::default(),
|
||||
] {
|
||||
let client_id = encode_client_id(KEY, &meta).unwrap();
|
||||
assert!(client_id.starts_with(CLIENT_ID_HEADER));
|
||||
assert_eq!(decode_client_id(KEY, &client_id), Some(meta));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn scope_mask_is_order_independent_and_drops_unknown() {
|
||||
assert_eq!(
|
||||
scopes_to_mask(&format!("{SCOPE_MAIL} {SCOPE_OFFLINE_ACCESS}")),
|
||||
scopes_to_mask(&format!("{SCOPE_OFFLINE_ACCESS} {SCOPE_MAIL}"))
|
||||
);
|
||||
assert_eq!(
|
||||
scopes_to_mask(&format!("{SCOPE_MAIL} custom:unknown")),
|
||||
scopes_to_mask(SCOPE_MAIL)
|
||||
);
|
||||
assert_eq!(scopes_to_mask("totally unknown"), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn identical_input_is_deterministic() {
|
||||
let meta = sample();
|
||||
assert_eq!(
|
||||
encode_client_id(KEY, &meta).unwrap(),
|
||||
encode_client_id(KEY, &meta).unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn wrong_key_is_rejected() {
|
||||
let client_id = encode_client_id(KEY, &sample()).unwrap();
|
||||
assert_eq!(
|
||||
decode_client_id(b"a-completely-different-key-value!", &client_id),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tampering_is_rejected() {
|
||||
let client_id = encode_client_id(KEY, &sample()).unwrap();
|
||||
let (header, body_b64) = client_id.split_at(CLIENT_ID_HEADER.len());
|
||||
let mut body = general_purpose::URL_SAFE_NO_PAD.decode(body_b64).unwrap();
|
||||
for idx in 0..body.len() {
|
||||
let mut tampered = body.clone();
|
||||
tampered[idx] ^= 0x01;
|
||||
let forged = format!(
|
||||
"{header}{}",
|
||||
general_purpose::URL_SAFE_NO_PAD.encode(&tampered)
|
||||
);
|
||||
assert_eq!(decode_client_id(KEY, &forged), None, "byte {idx}");
|
||||
}
|
||||
body[0] ^= 0x00;
|
||||
assert!(decode_client_id(KEY, &client_id).is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn malformed_input_never_panics() {
|
||||
for case in [
|
||||
"",
|
||||
"swc1.",
|
||||
"swc1.!!!",
|
||||
"swc1.AAAA",
|
||||
"wrong.AAAA",
|
||||
"swc1.AAAAAAAAAAAAAAAAAAAAAAAAAAAA",
|
||||
] {
|
||||
assert_eq!(decode_client_id(KEY, case), None, "{case:?}");
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,203 @@
|
||||
/*
|
||||
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
|
||||
*
|
||||
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
|
||||
*/
|
||||
|
||||
use crate::{
|
||||
config::{EcKeyCurve, build_ecdsa_pem, build_rsa_keypair},
|
||||
manager::application::Resource,
|
||||
};
|
||||
use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD};
|
||||
use jsonwebtoken::{
|
||||
Algorithm, EncodingKey,
|
||||
jwk::{
|
||||
AlgorithmParameters, CommonParameters, EllipticCurve, EllipticCurveKeyParameters,
|
||||
EllipticCurveKeyType, Jwk, JwkSet, KeyAlgorithm, OctetKeyParameters, OctetKeyType,
|
||||
PublicKeyUse, RSAKeyParameters, RSAKeyType,
|
||||
},
|
||||
};
|
||||
use registry::schema::{enums::JwtSignatureAlgorithm, prelude::ObjectType, structs::OidcProvider};
|
||||
use std::borrow::Cow;
|
||||
use store::{
|
||||
rand::{RngExt, distr::Alphanumeric, rng},
|
||||
registry::bootstrap::Bootstrap,
|
||||
};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct OAuthConfig {
|
||||
pub oauth_key: String,
|
||||
pub oauth_expiry_user_code: u64,
|
||||
pub oauth_expiry_auth_code: u64,
|
||||
pub oauth_expiry_token: u64,
|
||||
pub oauth_expiry_refresh_token: u64,
|
||||
pub oauth_expiry_refresh_token_renew: u64,
|
||||
pub oauth_max_auth_attempts: u32,
|
||||
|
||||
pub allow_anonymous_client_registration: bool,
|
||||
pub require_client_authentication: bool,
|
||||
|
||||
pub oidc_expiry_id_token: u64,
|
||||
pub oidc_signing_secret: EncodingKey,
|
||||
pub oidc_signature_algorithm: Algorithm,
|
||||
pub oidc_jwks: Resource<Vec<u8>>,
|
||||
}
|
||||
|
||||
impl OAuthConfig {
|
||||
pub async fn parse(bp: &mut Bootstrap) -> Self {
|
||||
let auth = bp.setting_infallible::<OidcProvider>().await;
|
||||
|
||||
let oidc_signature_algorithm = match auth.signature_algorithm {
|
||||
JwtSignatureAlgorithm::Es256 => Algorithm::ES256,
|
||||
JwtSignatureAlgorithm::Es384 => Algorithm::ES384,
|
||||
JwtSignatureAlgorithm::Ps256 => Algorithm::PS256,
|
||||
JwtSignatureAlgorithm::Ps384 => Algorithm::PS384,
|
||||
JwtSignatureAlgorithm::Ps512 => Algorithm::PS512,
|
||||
JwtSignatureAlgorithm::Rs256 => Algorithm::RS256,
|
||||
JwtSignatureAlgorithm::Rs384 => Algorithm::RS384,
|
||||
JwtSignatureAlgorithm::Rs512 => Algorithm::RS512,
|
||||
JwtSignatureAlgorithm::Hs256 => Algorithm::HS256,
|
||||
JwtSignatureAlgorithm::Hs384 => Algorithm::HS384,
|
||||
JwtSignatureAlgorithm::Hs512 => Algorithm::HS512,
|
||||
};
|
||||
|
||||
let rand_key = rng()
|
||||
.sample_iter(Alphanumeric)
|
||||
.take(64)
|
||||
.map(char::from)
|
||||
.collect::<String>();
|
||||
|
||||
let signature_key = auth
|
||||
.signature_key
|
||||
.secret()
|
||||
.await
|
||||
.map_err(|err| {
|
||||
bp.build_error(ObjectType::OidcProvider.singleton(), err);
|
||||
})
|
||||
.unwrap_or(Cow::Borrowed(rand_key.as_str()));
|
||||
|
||||
let fallback_key = || {
|
||||
(
|
||||
EncodingKey::from_secret(rand_key.as_bytes()),
|
||||
AlgorithmParameters::OctetKey(OctetKeyParameters {
|
||||
key_type: OctetKeyType::Octet,
|
||||
value: URL_SAFE_NO_PAD.encode(&rand_key),
|
||||
})
|
||||
.into(),
|
||||
)
|
||||
};
|
||||
|
||||
let (oidc_signing_secret, algorithm) = match oidc_signature_algorithm {
|
||||
Algorithm::HS256 | Algorithm::HS384 | Algorithm::HS512 => {
|
||||
(EncodingKey::from_secret(signature_key.as_bytes()), None)
|
||||
}
|
||||
Algorithm::RS256
|
||||
| Algorithm::RS384
|
||||
| Algorithm::RS512
|
||||
| Algorithm::PS256
|
||||
| Algorithm::PS384
|
||||
| Algorithm::PS512 => parse_rsa_key(&auth)
|
||||
.await
|
||||
.map_err(|err| {
|
||||
bp.build_error(ObjectType::OidcProvider.singleton(), err);
|
||||
})
|
||||
.map(|(secret, alg)| (secret, Some(alg)))
|
||||
.unwrap_or_else(|_| fallback_key()),
|
||||
Algorithm::ES256 | Algorithm::ES384 => parse_ecdsa_key(&auth, oidc_signature_algorithm)
|
||||
.await
|
||||
.map_err(|err| {
|
||||
bp.build_error(ObjectType::OidcProvider.singleton(), err);
|
||||
})
|
||||
.map(|(secret, alg)| (secret, Some(alg)))
|
||||
.unwrap_or_else(|_| fallback_key()),
|
||||
_ => {
|
||||
bp.build_error(
|
||||
ObjectType::OidcProvider.singleton(),
|
||||
format!("Unsupported OIDC signature algorithm {oidc_signature_algorithm:?}"),
|
||||
);
|
||||
fallback_key()
|
||||
}
|
||||
};
|
||||
|
||||
let oidc_jwks = Resource {
|
||||
content_type: "application/json".into(),
|
||||
contents: serde_json::to_string(&JwkSet {
|
||||
keys: algorithm
|
||||
.into_iter()
|
||||
.map(|algorithm| Jwk {
|
||||
common: CommonParameters {
|
||||
public_key_use: PublicKeyUse::Signature.into(),
|
||||
key_algorithm: KeyAlgorithm::from(oidc_signature_algorithm).into(),
|
||||
key_id: "default".to_string().into(),
|
||||
..Default::default()
|
||||
},
|
||||
algorithm,
|
||||
})
|
||||
.collect(),
|
||||
})
|
||||
.unwrap_or_default()
|
||||
.into_bytes(),
|
||||
};
|
||||
|
||||
OAuthConfig {
|
||||
oauth_key: auth
|
||||
.encryption_key
|
||||
.secret()
|
||||
.await
|
||||
.map_err(|err| bp.build_error(ObjectType::OidcProvider.singleton(), err))
|
||||
.map_or_else(|_| rand_key.clone(), Cow::into_owned),
|
||||
oauth_expiry_user_code: auth.user_code_expiry.as_secs(),
|
||||
oauth_expiry_auth_code: auth.auth_code_expiry.as_secs(),
|
||||
oauth_expiry_token: auth.access_token_expiry.as_secs(),
|
||||
oauth_expiry_refresh_token: auth.refresh_token_expiry.as_secs(),
|
||||
oauth_expiry_refresh_token_renew: auth.refresh_token_renewal.as_secs(),
|
||||
oauth_max_auth_attempts: auth.auth_code_max_attempts as u32,
|
||||
oidc_expiry_id_token: auth.id_token_expiry.as_secs(),
|
||||
allow_anonymous_client_registration: auth.anonymous_client_registration,
|
||||
require_client_authentication: auth.require_client_registration,
|
||||
oidc_signing_secret,
|
||||
oidc_signature_algorithm,
|
||||
oidc_jwks,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn parse_rsa_key(auth: &OidcProvider) -> Result<(EncodingKey, AlgorithmParameters), String> {
|
||||
let rsa_key = build_rsa_keypair(auth.signature_key.secret().await?.as_ref())?;
|
||||
|
||||
let rsa_key_params = RSAKeyParameters {
|
||||
key_type: RSAKeyType::RSA,
|
||||
n: URL_SAFE_NO_PAD.encode(&rsa_key.modulus),
|
||||
e: URL_SAFE_NO_PAD.encode(&rsa_key.exponent),
|
||||
};
|
||||
|
||||
Ok((
|
||||
EncodingKey::from_rsa_der(&rsa_key.pkcs1_der),
|
||||
AlgorithmParameters::RSA(rsa_key_params),
|
||||
))
|
||||
}
|
||||
|
||||
async fn parse_ecdsa_key(
|
||||
auth: &OidcProvider,
|
||||
oidc_signature_algorithm: Algorithm,
|
||||
) -> Result<(EncodingKey, AlgorithmParameters), String> {
|
||||
let (curve, ec_curve) = match oidc_signature_algorithm {
|
||||
Algorithm::ES256 => (EllipticCurve::P256, EcKeyCurve::P256),
|
||||
Algorithm::ES384 => (EllipticCurve::P384, EcKeyCurve::P384),
|
||||
_ => unreachable!(),
|
||||
};
|
||||
|
||||
let ecdsa_key = build_ecdsa_pem(ec_curve, auth.signature_key.secret().await?.as_ref())?;
|
||||
|
||||
let ecdsa_key_params = EllipticCurveKeyParameters {
|
||||
key_type: EllipticCurveKeyType::EC,
|
||||
curve,
|
||||
x: URL_SAFE_NO_PAD.encode(&ecdsa_key.x),
|
||||
y: URL_SAFE_NO_PAD.encode(&ecdsa_key.y),
|
||||
};
|
||||
|
||||
Ok((
|
||||
EncodingKey::from_ec_der(&ecdsa_key.pkcs8_der),
|
||||
AlgorithmParameters::EllipticCurve(ecdsa_key_params),
|
||||
))
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
/*
|
||||
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
|
||||
*
|
||||
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
|
||||
*/
|
||||
|
||||
use aes_gcm_siv::{
|
||||
Aes256GcmSiv, Key, KeyInit, Nonce,
|
||||
aead::{Aead, Payload},
|
||||
};
|
||||
use store::blake3;
|
||||
|
||||
pub struct SymmetricEncrypt {
|
||||
aes: Aes256GcmSiv,
|
||||
}
|
||||
|
||||
impl SymmetricEncrypt {
|
||||
pub const ENCRYPT_TAG_LEN: usize = 16;
|
||||
pub const NONCE_LEN: usize = 12;
|
||||
|
||||
pub fn new(key: &[u8], context: &str) -> Self {
|
||||
SymmetricEncrypt {
|
||||
aes: Aes256GcmSiv::new(&Key::<Aes256GcmSiv>::from(blake3::derive_key(context, key))),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn encrypt_with_aad(
|
||||
&self,
|
||||
bytes: &[u8],
|
||||
nonce: &[u8],
|
||||
aad: &[u8],
|
||||
) -> Result<Vec<u8>, String> {
|
||||
self.aes
|
||||
.encrypt(
|
||||
<&Nonce>::try_from(nonce).map_err(|e| e.to_string())?,
|
||||
Payload { msg: bytes, aad },
|
||||
)
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
pub fn decrypt_with_aad(
|
||||
&self,
|
||||
bytes: &[u8],
|
||||
nonce: &[u8],
|
||||
aad: &[u8],
|
||||
) -> Result<Vec<u8>, String> {
|
||||
self.aes
|
||||
.decrypt(
|
||||
<&Nonce>::try_from(nonce).map_err(|e| e.to_string())?,
|
||||
Payload { msg: bytes, aad },
|
||||
)
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,81 @@
|
||||
/*
|
||||
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
|
||||
*
|
||||
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
|
||||
*/
|
||||
|
||||
use crate::{Server, auth::AccessToken};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use trc::{AddContext, AuthEvent, EventType};
|
||||
|
||||
#[derive(Debug, Default, Clone, Eq, PartialEq, Deserialize, Serialize)]
|
||||
pub struct OAuthIntrospect {
|
||||
#[serde(default)]
|
||||
pub active: bool,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub scope: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub client_id: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub username: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub token_type: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub exp: Option<i64>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub iat: Option<i64>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub nbf: Option<i64>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub sub: Option<String>,
|
||||
}
|
||||
|
||||
impl Server {
|
||||
pub async fn introspect_access_token(
|
||||
&self,
|
||||
token: &str,
|
||||
access_token: &AccessToken,
|
||||
) -> trc::Result<OAuthIntrospect> {
|
||||
match self.validate_access_token(None, token).await {
|
||||
Ok(token_info) => Ok(OAuthIntrospect {
|
||||
active: true,
|
||||
username: self
|
||||
.account(access_token.account_id())
|
||||
.await
|
||||
.caused_by(trc::location!())?
|
||||
.name()
|
||||
.to_string()
|
||||
.into(),
|
||||
token_type: Some("bearer".into()),
|
||||
exp: Some(token_info.expiry as i64),
|
||||
iat: Some(token_info.issued_at as i64),
|
||||
..Default::default()
|
||||
}),
|
||||
Err(err)
|
||||
if matches!(
|
||||
err.event_type(),
|
||||
EventType::Auth(AuthEvent::Error) | EventType::Auth(AuthEvent::TokenExpired)
|
||||
) =>
|
||||
{
|
||||
Ok(OAuthIntrospect::default())
|
||||
}
|
||||
Err(err) => Err(err),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
/*
|
||||
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
|
||||
*
|
||||
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
|
||||
*/
|
||||
|
||||
pub mod client_id;
|
||||
pub mod config;
|
||||
pub mod crypto;
|
||||
pub mod introspect;
|
||||
pub mod oidc;
|
||||
pub mod registration;
|
||||
pub mod token;
|
||||
|
||||
pub const DEVICE_CODE_LEN: usize = 40;
|
||||
pub const USER_CODE_LEN: usize = 8;
|
||||
pub const RANDOM_CODE_LEN: usize = 32;
|
||||
pub const CLIENT_ID_MAX_LEN: usize = 2048;
|
||||
|
||||
pub const USER_CODE_ALPHABET: &[u8] = b"ABCDEFGHJKLMNPQRSTUVWXYZ23456789"; // No 0, O, I, 1
|
||||
|
||||
pub const SCOPE_OPENID: &str = "openid";
|
||||
pub const SCOPE_OFFLINE_ACCESS: &str = "offline_access";
|
||||
pub const SCOPE_MAIL: &str = "urn:ietf:params:oauth:scope:mail";
|
||||
pub const SCOPE_CONTACTS: &str = "urn:ietf:params:oauth:scope:contacts";
|
||||
pub const SCOPE_CALENDARS: &str = "urn:ietf:params:oauth:scope:calendars";
|
||||
|
||||
pub const SUPPORTED_SCOPES: &[&str] = &[
|
||||
SCOPE_OPENID,
|
||||
SCOPE_OFFLINE_ACCESS,
|
||||
SCOPE_MAIL,
|
||||
SCOPE_CONTACTS,
|
||||
SCOPE_CALENDARS,
|
||||
];
|
||||
|
||||
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
|
||||
pub enum GrantType {
|
||||
AccessToken,
|
||||
RefreshToken,
|
||||
LiveTracing,
|
||||
LiveMetrics,
|
||||
LiveDelivery,
|
||||
Rsvp,
|
||||
}
|
||||
|
||||
impl GrantType {
|
||||
pub fn as_str(&self) -> &'static str {
|
||||
match self {
|
||||
GrantType::AccessToken => "access_token",
|
||||
GrantType::RefreshToken => "refresh_token",
|
||||
GrantType::LiveTracing => "live_tracing",
|
||||
GrantType::LiveMetrics => "live_metrics",
|
||||
GrantType::LiveDelivery => "live_delivery",
|
||||
GrantType::Rsvp => "rsvp",
|
||||
}
|
||||
}
|
||||
|
||||
pub fn id(&self) -> u8 {
|
||||
match self {
|
||||
GrantType::AccessToken => 0,
|
||||
GrantType::RefreshToken => 1,
|
||||
GrantType::LiveTracing => 2,
|
||||
GrantType::LiveMetrics => 3,
|
||||
GrantType::LiveDelivery => 4,
|
||||
GrantType::Rsvp => 5,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn from_id(id: u8) -> Option<Self> {
|
||||
match id {
|
||||
0 => Some(GrantType::AccessToken),
|
||||
1 => Some(GrantType::RefreshToken),
|
||||
2 => Some(GrantType::LiveTracing),
|
||||
3 => Some(GrantType::LiveMetrics),
|
||||
4 => Some(GrantType::LiveDelivery),
|
||||
5 => Some(GrantType::Rsvp),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,182 @@
|
||||
/*
|
||||
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
|
||||
*
|
||||
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
|
||||
*/
|
||||
|
||||
use std::fmt;
|
||||
|
||||
use jsonwebtoken::Header;
|
||||
|
||||
use serde::{
|
||||
Deserialize, Deserializer, Serialize,
|
||||
de::{self, Visitor},
|
||||
};
|
||||
use store::write::now;
|
||||
|
||||
use crate::Server;
|
||||
|
||||
#[derive(Debug, Default, Clone, Eq, PartialEq, Deserialize, Serialize)]
|
||||
pub struct Userinfo {
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub sub: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub name: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub given_name: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub family_name: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub middle_name: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub nickname: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub preferred_username: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub profile: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub picture: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub website: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub email: Option<String>,
|
||||
|
||||
#[serde(default, deserialize_with = "any_bool")]
|
||||
#[serde(skip_serializing_if = "std::ops::Not::not")]
|
||||
pub email_verified: bool,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub zoneinfo: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub locale: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub updated_at: Option<i64>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Default, Eq, PartialEq, Deserialize, Serialize)]
|
||||
pub struct StandardClaims {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
#[serde(default)]
|
||||
pub nonce: Option<String>,
|
||||
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
#[serde(default)]
|
||||
pub preferred_username: Option<String>,
|
||||
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
#[serde(default)]
|
||||
pub email: Option<String>,
|
||||
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
#[serde(default)]
|
||||
pub description: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
struct IdTokenClaims {
|
||||
iss: String,
|
||||
sub: String,
|
||||
aud: String,
|
||||
nbf: i64,
|
||||
iat: i64,
|
||||
exp: i64,
|
||||
|
||||
#[serde(flatten)]
|
||||
private: StandardClaims,
|
||||
}
|
||||
|
||||
impl Server {
|
||||
pub fn issue_id_token(
|
||||
&self,
|
||||
subject: impl Into<String>,
|
||||
issuer: impl Into<String>,
|
||||
audience: impl Into<String>,
|
||||
claims: StandardClaims,
|
||||
) -> trc::Result<String> {
|
||||
let now = now() as i64;
|
||||
|
||||
jsonwebtoken::encode(
|
||||
&Header {
|
||||
kid: Some("default".into()),
|
||||
..Header::new(self.core.oauth.oidc_signature_algorithm)
|
||||
},
|
||||
&IdTokenClaims {
|
||||
iss: issuer.into(),
|
||||
sub: subject.into(),
|
||||
aud: audience.into(),
|
||||
nbf: now,
|
||||
iat: now,
|
||||
exp: now + self.core.oauth.oidc_expiry_id_token as i64,
|
||||
private: claims,
|
||||
},
|
||||
&self.core.oauth.oidc_signing_secret,
|
||||
)
|
||||
.map_err(|err| {
|
||||
trc::AuthEvent::Error
|
||||
.into_err()
|
||||
.reason(err)
|
||||
.details("Failed to encode ID token")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn any_bool<'de, D>(deserializer: D) -> Result<bool, D::Error>
|
||||
where
|
||||
D: Deserializer<'de>,
|
||||
{
|
||||
struct AnyBoolVisitor;
|
||||
|
||||
impl Visitor<'_> for AnyBoolVisitor {
|
||||
type Value = bool;
|
||||
|
||||
fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
|
||||
formatter.write_str("a boolean value")
|
||||
}
|
||||
|
||||
fn visit_str<E>(self, value: &str) -> Result<bool, E>
|
||||
where
|
||||
E: de::Error,
|
||||
{
|
||||
match value {
|
||||
"true" => Ok(true),
|
||||
"false" => Ok(false),
|
||||
_ => Err(E::custom(format!("Unknown boolean: {value}"))),
|
||||
}
|
||||
}
|
||||
|
||||
fn visit_bool<E>(self, value: bool) -> Result<bool, E>
|
||||
where
|
||||
E: de::Error,
|
||||
{
|
||||
Ok(value)
|
||||
}
|
||||
}
|
||||
|
||||
deserializer.deserialize_any(AnyBoolVisitor)
|
||||
}
|
||||
@@ -0,0 +1,295 @@
|
||||
/*
|
||||
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
|
||||
*
|
||||
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
|
||||
*/
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
|
||||
#[derive(Serialize, Deserialize, Debug, Default)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub struct ClientRegistrationRequest {
|
||||
pub redirect_uris: Vec<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub scope: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Vec::is_empty")]
|
||||
pub response_types: Vec<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Vec::is_empty")]
|
||||
pub grant_types: Vec<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub application_type: Option<ApplicationType>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Vec::is_empty")]
|
||||
pub contacts: Vec<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub client_name: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub logo_uri: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub client_uri: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub policy_uri: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tos_uri: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub jwks_uri: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub jwks: Option<serde_json::Value>, // Using serde_json::Value for flexibility
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub sector_identifier_uri: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub subject_type: Option<SubjectType>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub id_token_signed_response_alg: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub id_token_encrypted_response_alg: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub id_token_encrypted_response_enc: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub userinfo_signed_response_alg: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub userinfo_encrypted_response_alg: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub userinfo_encrypted_response_enc: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub request_object_signing_alg: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub request_object_encryption_alg: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub request_object_encryption_enc: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub token_endpoint_auth_method: Option<TokenEndpointAuthMethod>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub token_endpoint_auth_signing_alg: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub default_max_age: Option<u64>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub require_auth_time: Option<bool>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Vec::is_empty")]
|
||||
pub default_acr_values: Vec<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub initiate_login_uri: Option<String>,
|
||||
|
||||
#[serde(default)]
|
||||
#[serde(skip_serializing_if = "Vec::is_empty")]
|
||||
pub request_uris: Vec<String>,
|
||||
|
||||
#[serde(flatten)]
|
||||
#[serde(skip_serializing_if = "HashMap::is_empty")]
|
||||
pub additional_fields: HashMap<String, serde_json::Value>,
|
||||
}
|
||||
|
||||
#[derive(Serialize, Deserialize, Debug, Default)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub struct ClientRegistrationResponse {
|
||||
// Required fields
|
||||
pub client_id: String,
|
||||
|
||||
// Optional fields specific to the response
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub client_secret: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub registration_access_token: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub registration_client_uri: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub client_id_issued_at: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub client_secret_expires_at: Option<u64>,
|
||||
|
||||
// Echo back the request
|
||||
#[serde(flatten)]
|
||||
pub request: ClientRegistrationRequest,
|
||||
}
|
||||
|
||||
#[derive(Serialize, Deserialize, Debug)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub enum ApplicationType {
|
||||
Web,
|
||||
Native,
|
||||
}
|
||||
|
||||
#[derive(Serialize, Deserialize, Debug)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub enum SubjectType {
|
||||
Pairwise,
|
||||
Public,
|
||||
}
|
||||
|
||||
#[derive(Serialize, Deserialize, Debug, Clone, Copy, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum TokenEndpointAuthMethod {
|
||||
ClientSecretPost,
|
||||
ClientSecretBasic,
|
||||
ClientSecretJwt,
|
||||
PrivateKeyJwt,
|
||||
None,
|
||||
}
|
||||
|
||||
#[derive(Serialize, Debug)]
|
||||
pub struct ClientRegistrationError {
|
||||
pub error: &'static str,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub error_description: Option<&'static str>,
|
||||
}
|
||||
|
||||
impl ClientRegistrationError {
|
||||
pub fn invalid_redirect_uri(description: &'static str) -> Self {
|
||||
ClientRegistrationError {
|
||||
error: "invalid_redirect_uri",
|
||||
error_description: Some(description),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn invalid_client_metadata(description: &'static str) -> Self {
|
||||
ClientRegistrationError {
|
||||
error: "invalid_client_metadata",
|
||||
error_description: Some(description),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn loopback_redirect_parts(uri: &str) -> Option<(&str, &str)> {
|
||||
let uri = uri.strip_prefix("http://")?;
|
||||
|
||||
for host in ["127.0.0.1", "[::1]"] {
|
||||
if let Some(rest) = uri.strip_prefix(host) {
|
||||
if let Some(path) = rest.strip_prefix('/') {
|
||||
return Some((host, path));
|
||||
} else if let Some(after_colon) = rest.strip_prefix(':')
|
||||
&& let Some((port, path)) = after_colon.split_once('/')
|
||||
&& !port.is_empty()
|
||||
&& port.bytes().all(|b| b.is_ascii_digit())
|
||||
{
|
||||
return Some((host, path));
|
||||
}
|
||||
}
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
pub fn redirect_uri_matches(registered: &str, presented: &str) -> bool {
|
||||
registered == presented
|
||||
|| matches!(
|
||||
(
|
||||
loopback_redirect_parts(registered),
|
||||
loopback_redirect_parts(presented),
|
||||
),
|
||||
(Some(reg), Some(pres)) if reg == pres
|
||||
)
|
||||
}
|
||||
|
||||
pub fn validate_redirect_uri(uri: &str) -> Result<(), ClientRegistrationError> {
|
||||
if uri.contains('#') {
|
||||
return Err(ClientRegistrationError::invalid_redirect_uri(
|
||||
"Redirect URI must not contain a fragment.",
|
||||
));
|
||||
} else if uri.contains("..") {
|
||||
return Err(ClientRegistrationError::invalid_redirect_uri(
|
||||
"Redirect URI must not contain consecutive dots.",
|
||||
));
|
||||
} else if uri.starts_with("https://") || loopback_redirect_parts(uri).is_some() {
|
||||
return Ok(());
|
||||
} else if let Some((scheme, _)) = uri.split_once(':')
|
||||
&& scheme.contains('.')
|
||||
&& scheme
|
||||
.as_bytes()
|
||||
.first()
|
||||
.is_some_and(u8::is_ascii_alphabetic)
|
||||
&& scheme
|
||||
.bytes()
|
||||
.all(|b| b.is_ascii_alphanumeric() || matches!(b, b'.' | b'-' | b'+'))
|
||||
{
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
Err(ClientRegistrationError::invalid_redirect_uri(
|
||||
"Redirect URI must be an https URL, a loopback (http://127.0.0.1/, http://[::1]/) or a private-use scheme URI.",
|
||||
))
|
||||
}
|
||||
|
||||
pub fn validate_grant_metadata(
|
||||
request: &ClientRegistrationRequest,
|
||||
) -> Result<(), ClientRegistrationError> {
|
||||
if !request.response_types.is_empty() && !request.response_types.iter().any(|t| t == "code") {
|
||||
return Err(ClientRegistrationError::invalid_client_metadata(
|
||||
"response_types must include \"code\".",
|
||||
));
|
||||
}
|
||||
if !request.grant_types.is_empty() {
|
||||
if !request
|
||||
.grant_types
|
||||
.iter()
|
||||
.any(|t| t == "authorization_code")
|
||||
{
|
||||
return Err(ClientRegistrationError::invalid_client_metadata(
|
||||
"grant_types must include \"authorization_code\".",
|
||||
));
|
||||
}
|
||||
if !request.grant_types.iter().any(|t| t == "refresh_token") {
|
||||
return Err(ClientRegistrationError::invalid_client_metadata(
|
||||
"grant_types must include \"refresh_token\".",
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,418 @@
|
||||
/*
|
||||
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
|
||||
*
|
||||
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
|
||||
*/
|
||||
|
||||
use super::{GrantType, crypto::SymmetricEncrypt};
|
||||
use crate::Server;
|
||||
use base64::{Engine, engine::general_purpose};
|
||||
use std::time::SystemTime;
|
||||
use store::rand::{RngExt, rng};
|
||||
use utils::codec::leb128::{Leb128Iterator, Leb128Vec};
|
||||
|
||||
pub const FAILED_TO_DECODE_TOKEN: &str = concat!(
|
||||
"Failed to decode token. If you are using an ",
|
||||
"external OIDC provider, make sure it is configured as the default directory under ",
|
||||
"the Authentication object."
|
||||
);
|
||||
|
||||
const TOKEN_HEADER: &str = "sw1.";
|
||||
const TOKEN_KEY_CONTEXT: &str = "stalwart-oauth-token-sw1";
|
||||
const OAUTH_EPOCH: u64 = 946684800; // Jan 1, 2000
|
||||
|
||||
pub struct TokenInfo {
|
||||
pub grant_type: GrantType,
|
||||
pub account_id: u32,
|
||||
pub claims: Option<String>,
|
||||
pub expiry: u64,
|
||||
pub issued_at: u64,
|
||||
pub expires_in: u64,
|
||||
}
|
||||
|
||||
struct RawToken {
|
||||
grant_type: GrantType,
|
||||
account_id: u32,
|
||||
claims: Option<String>,
|
||||
issued_at: u64,
|
||||
expiry: u64,
|
||||
credential_version: u64,
|
||||
}
|
||||
|
||||
impl Server {
|
||||
pub async fn encode_access_token(
|
||||
&self,
|
||||
grant_type: GrantType,
|
||||
account_id: u32,
|
||||
account_name: &str,
|
||||
expiry_in: u64,
|
||||
claims: Option<&str>,
|
||||
credential_version: Option<u64>,
|
||||
) -> trc::Result<String> {
|
||||
let issued_at = seconds_since_oauth_epoch();
|
||||
let raw = RawToken {
|
||||
grant_type,
|
||||
account_id,
|
||||
claims: claims.map(|claims| claims.to_string()),
|
||||
issued_at,
|
||||
expiry: issued_at + expiry_in,
|
||||
credential_version: credential_version
|
||||
.filter(|_| !matches!(grant_type, GrantType::Rsvp))
|
||||
.unwrap_or_default(),
|
||||
};
|
||||
|
||||
seal_token(
|
||||
self.core.oauth.oauth_key.as_bytes(),
|
||||
&raw,
|
||||
account_name.as_bytes(),
|
||||
)
|
||||
.map_err(|err| {
|
||||
trc::AuthEvent::Error
|
||||
.into_err()
|
||||
.ctx(trc::Key::Reason, "Failed to encrypt token")
|
||||
.reason(err)
|
||||
.caused_by(trc::location!())
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn validate_access_token(
|
||||
&self,
|
||||
expected_grant_type: Option<GrantType>,
|
||||
token_: &str,
|
||||
) -> trc::Result<TokenInfo> {
|
||||
let token = open_token(self.core.oauth.oauth_key.as_bytes(), token_).map_err(|_| {
|
||||
trc::AuthEvent::Error
|
||||
.into_err()
|
||||
.ctx(trc::Key::Reason, FAILED_TO_DECODE_TOKEN)
|
||||
.caused_by(trc::location!())
|
||||
.details(token_.to_string())
|
||||
})?;
|
||||
|
||||
// Validate expiration
|
||||
let now = seconds_since_oauth_epoch();
|
||||
if token.expiry <= now || token.issued_at > now {
|
||||
return Err(trc::AuthEvent::TokenExpired.into_err());
|
||||
}
|
||||
|
||||
// Validate grant type
|
||||
if expected_grant_type.is_some_and(|g| g != token.grant_type) {
|
||||
return Err(trc::AuthEvent::Error
|
||||
.into_err()
|
||||
.details("Invalid grant type"));
|
||||
}
|
||||
|
||||
// Enforce credential revocation for long lived tokens
|
||||
if token.credential_version != 0 {
|
||||
let current = self
|
||||
.access_token(token.account_id)
|
||||
.await
|
||||
.map_err(|err| trc::AuthEvent::Error.into_err().ctx(trc::Key::Details, err))?
|
||||
.credential_version();
|
||||
if current != token.credential_version {
|
||||
return Err(trc::AuthEvent::TokenExpired
|
||||
.into_err()
|
||||
.details("Token revoked"));
|
||||
}
|
||||
}
|
||||
|
||||
Ok(TokenInfo {
|
||||
grant_type: token.grant_type,
|
||||
account_id: token.account_id,
|
||||
claims: token.claims,
|
||||
expiry: token.expiry + OAUTH_EPOCH,
|
||||
issued_at: token.issued_at + OAUTH_EPOCH,
|
||||
expires_in: token.expiry - now,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn seal_token(key: &[u8], token: &RawToken, footer: &[u8]) -> Result<String, String> {
|
||||
let mut payload = Vec::with_capacity(32);
|
||||
payload.push_leb128(token.account_id);
|
||||
payload.push(token.grant_type.id());
|
||||
payload.push_leb128(token.issued_at);
|
||||
payload.push_leb128(token.expiry);
|
||||
payload.push_leb128(token.credential_version);
|
||||
if let Some(claims) = token.claims.as_deref().filter(|claims| !claims.is_empty()) {
|
||||
payload.extend_from_slice(claims.as_bytes());
|
||||
}
|
||||
|
||||
let nonce = rng().random::<[u8; SymmetricEncrypt::NONCE_LEN]>();
|
||||
let ciphertext =
|
||||
SymmetricEncrypt::new(key, TOKEN_KEY_CONTEXT).encrypt_with_aad(&payload, &nonce, footer)?;
|
||||
|
||||
let mut body = Vec::with_capacity(nonce.len() + ciphertext.len());
|
||||
body.extend_from_slice(&nonce);
|
||||
body.extend_from_slice(&ciphertext);
|
||||
|
||||
let mut out = String::with_capacity(TOKEN_HEADER.len() + (body.len() + footer.len()) * 2);
|
||||
out.push_str(TOKEN_HEADER);
|
||||
general_purpose::URL_SAFE_NO_PAD.encode_string(&body, &mut out);
|
||||
if !footer.is_empty() {
|
||||
out.push('.');
|
||||
general_purpose::URL_SAFE_NO_PAD.encode_string(footer, &mut out);
|
||||
}
|
||||
|
||||
Ok(out)
|
||||
}
|
||||
|
||||
fn open_token(key: &[u8], token: &str) -> Result<RawToken, ()> {
|
||||
let rest = token.strip_prefix(TOKEN_HEADER).ok_or(())?;
|
||||
let (body, footer) = match rest.split_once('.') {
|
||||
Some((body, footer)) => (
|
||||
body,
|
||||
general_purpose::URL_SAFE_NO_PAD
|
||||
.decode(footer.as_bytes())
|
||||
.map_err(|_| ())?,
|
||||
),
|
||||
None => (rest, Vec::new()),
|
||||
};
|
||||
let body = general_purpose::URL_SAFE_NO_PAD
|
||||
.decode(body.as_bytes())
|
||||
.map_err(|_| ())?;
|
||||
if body.len() < SymmetricEncrypt::NONCE_LEN + SymmetricEncrypt::ENCRYPT_TAG_LEN {
|
||||
return Err(());
|
||||
}
|
||||
let (nonce, ciphertext) = body.split_at(SymmetricEncrypt::NONCE_LEN);
|
||||
|
||||
let payload = SymmetricEncrypt::new(key, TOKEN_KEY_CONTEXT)
|
||||
.decrypt_with_aad(ciphertext, nonce, &footer)
|
||||
.map_err(|_| ())?;
|
||||
|
||||
let mut bytes = payload.iter();
|
||||
let account_id: u32 = bytes.next_leb128().ok_or(())?;
|
||||
let grant_type = GrantType::from_id(bytes.next().copied().ok_or(())?).ok_or(())?;
|
||||
let issued_at: u64 = bytes.next_leb128().ok_or(())?;
|
||||
let expiry: u64 = bytes.next_leb128().ok_or(())?;
|
||||
let credential_version: u64 = bytes.next_leb128().ok_or(())?;
|
||||
let bytes = bytes.as_slice();
|
||||
let claims = if bytes.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(String::from_utf8(bytes.to_vec()).map_err(|_| ())?)
|
||||
};
|
||||
|
||||
Ok(RawToken {
|
||||
grant_type,
|
||||
account_id,
|
||||
claims,
|
||||
issued_at,
|
||||
expiry,
|
||||
credential_version,
|
||||
})
|
||||
}
|
||||
|
||||
#[inline(always)]
|
||||
fn seconds_since_oauth_epoch() -> u64 {
|
||||
SystemTime::now()
|
||||
.duration_since(SystemTime::UNIX_EPOCH)
|
||||
.map_or(0, |d| d.as_secs())
|
||||
.saturating_sub(OAUTH_EPOCH)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
const KEY: &[u8] = b"a-test-encryption-key-of-some-length";
|
||||
const NAME: &[u8] = b"[email protected]";
|
||||
|
||||
fn sample(grant_type: GrantType, claims: Option<&str>, cv: u64) -> RawToken {
|
||||
RawToken {
|
||||
grant_type,
|
||||
account_id: 42,
|
||||
claims: claims.map(|c| c.to_string()),
|
||||
issued_at: 1_000,
|
||||
expiry: 2_000,
|
||||
credential_version: cv,
|
||||
}
|
||||
}
|
||||
|
||||
fn assert_eq_fields(a: &RawToken, b: &RawToken) {
|
||||
assert_eq!(a.account_id, b.account_id);
|
||||
assert_eq!(a.grant_type, b.grant_type);
|
||||
assert_eq!(a.claims, b.claims);
|
||||
assert_eq!(a.issued_at, b.issued_at);
|
||||
assert_eq!(a.expiry, b.expiry);
|
||||
assert_eq!(a.credential_version, b.credential_version);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn round_trip_preserves_all_fields() {
|
||||
for (raw, footer) in [
|
||||
(sample(GrantType::AccessToken, None, 0), NAME),
|
||||
(
|
||||
sample(GrantType::RefreshToken, None, 0xdead_beef_cafe),
|
||||
NAME,
|
||||
),
|
||||
(
|
||||
sample(GrantType::Rsvp, Some("[email protected];7"), 0),
|
||||
b"[email protected]",
|
||||
),
|
||||
(sample(GrantType::AccessToken, None, 0), b""),
|
||||
(
|
||||
RawToken {
|
||||
account_id: u32::MAX,
|
||||
credential_version: u64::MAX,
|
||||
..sample(GrantType::AccessToken, Some("名前;1"), 1)
|
||||
},
|
||||
"名字@example.org".as_bytes(),
|
||||
),
|
||||
] {
|
||||
let token = seal_token(KEY, &raw, footer).unwrap();
|
||||
assert!(token.starts_with(TOKEN_HEADER));
|
||||
let opened = open_token(KEY, &token).unwrap();
|
||||
assert_eq_fields(&raw, &opened);
|
||||
|
||||
// The footer (account name) round-trips in clear text for proxies
|
||||
if footer.is_empty() {
|
||||
assert!(!token[TOKEN_HEADER.len()..].contains('.'));
|
||||
} else {
|
||||
let segment = token.rsplit_once('.').unwrap().1;
|
||||
assert_eq!(
|
||||
general_purpose::URL_SAFE_NO_PAD.decode(segment).unwrap(),
|
||||
footer
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn account_name_is_readable_in_clear_text_footer() {
|
||||
let token = seal_token(
|
||||
KEY,
|
||||
&sample(GrantType::AccessToken, None, 0),
|
||||
b"[email protected]",
|
||||
)
|
||||
.unwrap();
|
||||
let footer = token.rsplit_once('.').unwrap().1;
|
||||
let decoded = general_purpose::URL_SAFE_NO_PAD.decode(footer).unwrap();
|
||||
assert_eq!(decoded, b"[email protected]");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn wrong_key_is_rejected() {
|
||||
let token = seal_token(KEY, &sample(GrantType::AccessToken, None, 0), NAME).unwrap();
|
||||
assert!(open_token(b"a-different-encryption-key-entirely!", &token).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tampering_with_ciphertext_is_rejected() {
|
||||
let raw = sample(GrantType::AccessToken, None, 0);
|
||||
let token = seal_token(KEY, &raw, NAME).unwrap();
|
||||
let (header, rest) = token.split_at(TOKEN_HEADER.len());
|
||||
let (body_b64, footer) = match rest.split_once('.') {
|
||||
Some((b, f)) => (b.to_string(), Some(f.to_string())),
|
||||
None => (rest.to_string(), None),
|
||||
};
|
||||
let mut body = general_purpose::URL_SAFE_NO_PAD.decode(&body_b64).unwrap();
|
||||
|
||||
for idx in 0..body.len() {
|
||||
let mut tampered = body.clone();
|
||||
tampered[idx] ^= 0x01;
|
||||
let mut rebuilt = String::from(header);
|
||||
rebuilt.push_str(&general_purpose::URL_SAFE_NO_PAD.encode(&tampered));
|
||||
if let Some(footer) = &footer {
|
||||
rebuilt.push('.');
|
||||
rebuilt.push_str(footer);
|
||||
}
|
||||
assert!(
|
||||
open_token(KEY, &rebuilt).is_err(),
|
||||
"flipping byte {idx} of the body must invalidate the token"
|
||||
);
|
||||
}
|
||||
|
||||
// Sanity: the untampered token still opens
|
||||
body[0] ^= 0x00;
|
||||
assert!(open_token(KEY, &token).is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tampering_with_clear_text_footer_is_rejected() {
|
||||
let raw = sample(GrantType::AccessToken, None, 0);
|
||||
let token = seal_token(KEY, &raw, b"[email protected]").unwrap();
|
||||
let (body, _) = token.rsplit_once('.').unwrap();
|
||||
|
||||
// An attacker rewrites the clear-text account name to impersonate another account
|
||||
let forged_footer = general_purpose::URL_SAFE_NO_PAD.encode(b"[email protected]");
|
||||
let forged = format!("{body}.{forged_footer}");
|
||||
assert!(
|
||||
open_token(KEY, &forged).is_err(),
|
||||
"the footer is bound through the associated data and must be authenticated"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn swapping_footers_between_tokens_is_rejected() {
|
||||
let a = seal_token(
|
||||
KEY,
|
||||
&sample(GrantType::AccessToken, None, 0),
|
||||
b"[email protected]",
|
||||
)
|
||||
.unwrap();
|
||||
let b = seal_token(
|
||||
KEY,
|
||||
&sample(GrantType::AccessToken, None, 0),
|
||||
b"[email protected]",
|
||||
)
|
||||
.unwrap();
|
||||
let a_body = a.rsplit_once('.').unwrap().0;
|
||||
let b_footer = b.rsplit_once('.').unwrap().1;
|
||||
let frankentoken = format!("{a_body}.{b_footer}");
|
||||
assert!(open_token(KEY, &frankentoken).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn malformed_input_never_panics_and_is_rejected() {
|
||||
let valid = seal_token(KEY, &sample(GrantType::AccessToken, None, 0), NAME).unwrap();
|
||||
let cases = [
|
||||
String::new(),
|
||||
"sw1.".to_string(),
|
||||
"sw1.!!!not-base64!!!".to_string(),
|
||||
"sw1...".to_string(),
|
||||
"wrong-prefix.".to_string(),
|
||||
"sw1.AAAA".to_string(),
|
||||
"sw1.AAAA.BBBB".to_string(),
|
||||
valid.replace("sw1.", "sw2."),
|
||||
valid[..valid.len() / 2].to_string(),
|
||||
format!("sw1.{}", "A".repeat(10_000)),
|
||||
"\u{0}\u{0}\u{0}".to_string(),
|
||||
];
|
||||
for case in cases {
|
||||
assert!(open_token(KEY, &case).is_err(), "must reject {case:?}");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn truncating_the_body_is_rejected() {
|
||||
let token = seal_token(KEY, &sample(GrantType::AccessToken, None, 0), NAME).unwrap();
|
||||
let (header, rest) = token.split_at(TOKEN_HEADER.len());
|
||||
let body_b64 = rest.split_once('.').map(|(b, _)| b).unwrap_or(rest);
|
||||
let body = general_purpose::URL_SAFE_NO_PAD.decode(body_b64).unwrap();
|
||||
for len in 0..body.len() {
|
||||
let mut rebuilt = String::from(header);
|
||||
rebuilt.push_str(&general_purpose::URL_SAFE_NO_PAD.encode(&body[..len]));
|
||||
assert!(
|
||||
open_token(KEY, &rebuilt).is_err(),
|
||||
"truncation to {len} must be rejected"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn identical_input_produces_distinct_tokens() {
|
||||
let raw = sample(GrantType::AccessToken, None, 7);
|
||||
let a = seal_token(KEY, &raw, NAME).unwrap();
|
||||
let b = seal_token(KEY, &raw, NAME).unwrap();
|
||||
assert_ne!(a, b, "a random nonce must make each token unique");
|
||||
assert_eq_fields(&open_token(KEY, &a).unwrap(), &open_token(KEY, &b).unwrap());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn claims_with_separators_round_trip_exactly() {
|
||||
let raw = sample(GrantType::Rsvp, Some("a;b;c;[email protected];999"), 0);
|
||||
let token = seal_token(KEY, &raw, b"[email protected]").unwrap();
|
||||
let opened = open_token(KEY, &token).unwrap();
|
||||
assert_eq!(opened.claims.as_deref(), Some("a;b;c;[email protected];999"));
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user