use std::collections::BTreeSet;
use base64::Engine as _;
use serde_json::{Map, Value};
use crate::jsonrpc::{RawJsonAdmissionError, RawJsonAdmissionFailure, RawJsonTopLevel};
pub const MAX_OAUTH_METADATA_BYTES: usize = 64 * 1024;
pub const MAX_PROTECTED_RESOURCE_METADATA_BYTES: usize = 64 * 1024;
pub const MAX_OIDC_PROVIDER_METADATA_BYTES: usize = 128 * 1024;
pub const MAX_CLIENT_ID_METADATA_BYTES: usize = 32 * 1024;
pub const MAX_CLIENT_REGISTRATION_BYTES: usize = 32 * 1024;
pub const MAX_TOKEN_RESPONSE_BYTES: usize = 16 * 1024;
pub const MAX_JWK_SET_BYTES: usize = 64 * 1024;
pub const MAX_JWK_BYTES: usize = 8 * 1024;
pub const MAX_JWS_PROTECTED_HEADER_BYTES: usize = 2 * 1024;
pub const MAX_JWS_CLAIMS_BYTES: usize = 8 * 1024;
pub const MAX_COMPACT_JWS_ENCODED_BYTES: usize = 16 * 1024;
pub const MAX_JWS_SIGNATURE_BYTES: usize = 1024;
pub const MAX_RFC7638_CANONICAL_INPUT_BYTES: usize = 2 * 1024;
pub const MAX_JWK_KID_BYTES: usize = 256;
pub const MAX_JWK_SET_KEYS: usize = 64;
pub const MIN_RSA_MODULUS_BYTES: usize = 256;
pub const MAX_RSA_MODULUS_BYTES: usize = 512;
pub const ADMITTED_RSA_PUBLIC_EXPONENT: [u8; 3] = [0x01, 0x00, 0x01];
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum SecurityAdmissionError {
Document(RawJsonAdmissionFailure),
MalformedCompactSerialization,
InvalidBase64Url(&'static str),
NotAJsonObject(&'static str),
ProfileTypeMismatch,
UnsupportedAlgorithm,
MissingOrMalformedClaim(&'static str),
ForbiddenClaim(&'static str),
CrossProfileConfusion,
MalformedJwkMember(&'static str),
NonPublicKeyMaterial,
NotAVerificationKey,
AmbiguousKeyUsage,
UnsupportedKeyStrength,
DuplicateKeyIdentifier,
InputTooLong(&'static str),
}
impl std::fmt::Display for SecurityAdmissionError {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Document(failure) => write!(formatter, "security document refused: {failure}"),
Self::MalformedCompactSerialization => {
formatter.write_str("compact serialization must have exactly three segments")
}
Self::InvalidBase64Url(part) => {
write!(formatter, "{part} is not canonical unpadded base64url")
}
Self::NotAJsonObject(part) => write!(formatter, "{part} must be a JSON object"),
Self::ProfileTypeMismatch => {
formatter.write_str("the typ header does not match the selected JWS profile")
}
Self::UnsupportedAlgorithm => formatter.write_str("unsupported or absent alg header"),
Self::MissingOrMalformedClaim(claim) => {
write!(formatter, "claim {claim} is absent or malformed")
}
Self::ForbiddenClaim(claim) => {
write!(formatter, "claim {claim} is forbidden by this profile")
}
Self::CrossProfileConfusion => {
formatter.write_str("the claim set belongs to a different JWS profile")
}
Self::MalformedJwkMember(member) => {
write!(formatter, "JWK member {member} is absent or not canonical")
}
Self::NonPublicKeyMaterial => {
formatter.write_str("private or symmetric key material is never admitted")
}
Self::NotAVerificationKey => {
formatter.write_str("the key is not admitted for signature verification")
}
Self::AmbiguousKeyUsage => formatter.write_str("use and key_ops disagree"),
Self::UnsupportedKeyStrength => formatter.write_str("unsupported key strength"),
Self::DuplicateKeyIdentifier => formatter.write_str("duplicate kid in one key set"),
Self::InputTooLong(what) => write!(formatter, "{what} exceeds its fixed bound"),
}
}
}
impl std::error::Error for SecurityAdmissionError {}
impl From<RawJsonAdmissionFailure> for SecurityAdmissionError {
fn from(failure: RawJsonAdmissionFailure) -> Self {
Self::Document(failure)
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum SecurityDocumentKind {
AuthorizationServerMetadata,
ProtectedResourceMetadata,
OidcProviderMetadata,
ClientIdMetadataDocument,
ClientRegistration,
TokenResponse,
JsonWebKeySet,
JsonWebKey,
CompactJwsProtectedHeader,
CompactJwsClaims,
}
impl SecurityDocumentKind {
#[must_use]
pub const fn byte_limit(self) -> usize {
match self {
Self::AuthorizationServerMetadata => MAX_OAUTH_METADATA_BYTES,
Self::ProtectedResourceMetadata => MAX_PROTECTED_RESOURCE_METADATA_BYTES,
Self::OidcProviderMetadata => MAX_OIDC_PROVIDER_METADATA_BYTES,
Self::ClientIdMetadataDocument => MAX_CLIENT_ID_METADATA_BYTES,
Self::ClientRegistration => MAX_CLIENT_REGISTRATION_BYTES,
Self::TokenResponse => MAX_TOKEN_RESPONSE_BYTES,
Self::JsonWebKeySet => MAX_JWK_SET_BYTES,
Self::JsonWebKey => MAX_JWK_BYTES,
Self::CompactJwsProtectedHeader => MAX_JWS_PROTECTED_HEADER_BYTES,
Self::CompactJwsClaims => MAX_JWS_CLAIMS_BYTES,
}
}
#[must_use]
pub const fn all() -> [Self; 10] {
[
Self::AuthorizationServerMetadata,
Self::ProtectedResourceMetadata,
Self::OidcProviderMetadata,
Self::ClientIdMetadataDocument,
Self::ClientRegistration,
Self::TokenResponse,
Self::JsonWebKeySet,
Self::JsonWebKey,
Self::CompactJwsProtectedHeader,
Self::CompactJwsClaims,
]
}
}
pub fn admit_security_document(
kind: SecurityDocumentKind,
bytes: &[u8],
) -> Result<(), RawJsonAdmissionFailure> {
crate::jsonrpc::admit_raw_json_document(
bytes,
kind.byte_limit(),
RawJsonTopLevel::SecurityDocumentObject,
)
}
pub fn admit_security_document_object(
kind: SecurityDocumentKind,
bytes: &[u8],
) -> Result<Map<String, Value>, SecurityAdmissionError> {
admit_security_document(kind, bytes)?;
serde_json::from_slice::<Map<String, Value>>(bytes)
.map_err(|_| SecurityAdmissionError::NotAJsonObject("document"))
}
fn admit_jws_part(
kind: SecurityDocumentKind,
bytes: &[u8],
part: &'static str,
) -> Result<Map<String, Value>, SecurityAdmissionError> {
match admit_security_document_object(kind, bytes) {
Err(SecurityAdmissionError::Document(failure))
if failure.error() == RawJsonAdmissionError::TopLevelNotObject =>
{
Err(SecurityAdmissionError::NotAJsonObject(part))
}
other => other,
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum CompactJwsProfile {
Rfc9068AccessToken,
OidcIdToken,
IdentityAssertionJwtAuthorizationGrant,
Rfc7523ClientAssertion,
BuiltInIssuerSelfVerification,
}
impl CompactJwsProfile {
#[must_use]
pub const fn all() -> [Self; 5] {
[
Self::Rfc9068AccessToken,
Self::OidcIdToken,
Self::IdentityAssertionJwtAuthorizationGrant,
Self::Rfc7523ClientAssertion,
Self::BuiltInIssuerSelfVerification,
]
}
#[must_use]
pub const fn canonical_media_type(self) -> Option<&'static str> {
match self {
Self::Rfc9068AccessToken => Some("at+jwt"),
Self::IdentityAssertionJwtAuthorizationGrant => Some("oauth-id-jag+jwt"),
Self::OidcIdToken
| Self::Rfc7523ClientAssertion
| Self::BuiltInIssuerSelfVerification => Some("jwt"),
}
}
#[must_use]
pub const fn permits_absent_media_type(self) -> bool {
matches!(self, Self::OidcIdToken | Self::Rfc7523ClientAssertion)
}
#[must_use]
pub const fn authorization_relevant_claims(self) -> &'static [&'static str] {
match self {
Self::Rfc9068AccessToken => &["scope", "groups", "roles", "entitlements", "client_id"],
Self::OidcIdToken => &["nonce", "auth_time", "acr", "amr", "azp"],
Self::IdentityAssertionJwtAuthorizationGrant => &["scope", "client_id", "azp"],
Self::Rfc7523ClientAssertion => &["jti"],
Self::BuiltInIssuerSelfVerification => &["scope", "jti"],
}
}
#[must_use]
pub const fn forbidden_claims(self) -> &'static [&'static str] {
match self {
Self::Rfc9068AccessToken => &["nonce", "at_hash"],
Self::OidcIdToken => &["scope", "client_id"],
Self::IdentityAssertionJwtAuthorizationGrant => &["nonce", "at_hash"],
Self::Rfc7523ClientAssertion => &["scope", "nonce", "at_hash", "groups", "roles"],
Self::BuiltInIssuerSelfVerification => &["nonce", "at_hash"],
}
}
#[must_use]
pub const fn required_claims(self) -> &'static [&'static str] {
match self {
Self::Rfc9068AccessToken => &["iss", "sub", "aud", "exp", "iat", "jti"],
Self::OidcIdToken => &["iss", "sub", "aud", "exp", "iat"],
Self::IdentityAssertionJwtAuthorizationGrant => &["iss", "sub", "aud", "exp", "iat"],
Self::Rfc7523ClientAssertion => &["iss", "sub", "aud", "exp", "jti"],
Self::BuiltInIssuerSelfVerification => &["iss", "sub", "aud", "exp", "iat"],
}
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct JwsProtectedHeader {
algorithm: String,
media_type: Option<String>,
key_id: Option<String>,
}
impl JwsProtectedHeader {
#[must_use]
pub fn algorithm(&self) -> &str {
&self.algorithm
}
#[must_use]
pub fn media_type(&self) -> Option<&str> {
self.media_type.as_deref()
}
#[must_use]
pub fn key_id(&self) -> Option<&str> {
self.key_id.as_deref()
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct AdmittedCompactJws {
profile: CompactJwsProfile,
header: JwsProtectedHeader,
claims: Map<String, Value>,
signing_input: String,
signature: Vec<u8>,
}
impl AdmittedCompactJws {
#[must_use]
pub const fn profile(&self) -> CompactJwsProfile {
self.profile
}
#[must_use]
pub const fn header(&self) -> &JwsProtectedHeader {
&self.header
}
#[must_use]
pub const fn claims(&self) -> &Map<String, Value> {
&self.claims
}
#[must_use]
pub fn authorization_claim(&self, name: &str) -> Option<&Value> {
if self.profile.authorization_relevant_claims().contains(&name) {
self.claims.get(name)
} else {
None
}
}
#[must_use]
pub fn signing_input(&self) -> &str {
&self.signing_input
}
#[must_use]
pub fn signature(&self) -> &[u8] {
&self.signature
}
}
fn normalize_jose_media_type(value: &str) -> Option<String> {
if value.is_empty() || value.contains(';') || value.contains(char::is_whitespace) {
return None;
}
let lowered = value.to_ascii_lowercase();
let shortened = lowered.strip_prefix("application/").unwrap_or(&lowered);
if shortened.is_empty() || shortened.contains('/') {
return None;
}
Some(shortened.to_owned())
}
pub fn admit_compact_jws(
profile: CompactJwsProfile,
token: &str,
) -> Result<AdmittedCompactJws, SecurityAdmissionError> {
if token.len() > MAX_COMPACT_JWS_ENCODED_BYTES {
return Err(SecurityAdmissionError::InputTooLong("compact JWS"));
}
let mut segments = token.split('.');
let (Some(header_segment), Some(claims_segment), Some(signature_segment), None) = (
segments.next(),
segments.next(),
segments.next(),
segments.next(),
) else {
return Err(SecurityAdmissionError::MalformedCompactSerialization);
};
let header_bytes = decode_canonical_base64url(
header_segment,
MAX_JWS_PROTECTED_HEADER_BYTES,
"protected header",
)?;
let claims_bytes =
decode_canonical_base64url(claims_segment, MAX_JWS_CLAIMS_BYTES, "claims set")?;
let signature =
decode_canonical_base64url(signature_segment, MAX_JWS_SIGNATURE_BYTES, "signature")?;
let header_members = admit_jws_part(
SecurityDocumentKind::CompactJwsProtectedHeader,
&header_bytes,
"protected header",
)?;
let claims = admit_jws_part(
SecurityDocumentKind::CompactJwsClaims,
&claims_bytes,
"claims set",
)?;
let header = admit_protected_header(profile, &header_members)?;
admit_profile_claims(profile, &claims)?;
Ok(AdmittedCompactJws {
profile,
header,
claims,
signing_input: format!("{header_segment}.{claims_segment}"),
signature,
})
}
fn admit_protected_header(
profile: CompactJwsProfile,
members: &Map<String, Value>,
) -> Result<JwsProtectedHeader, SecurityAdmissionError> {
let Some(Value::String(algorithm)) = members.get("alg") else {
return Err(SecurityAdmissionError::UnsupportedAlgorithm);
};
if algorithm.eq_ignore_ascii_case("none") || algorithm.is_empty() {
return Err(SecurityAdmissionError::UnsupportedAlgorithm);
}
let media_type = match members.get("typ") {
None => {
if !profile.permits_absent_media_type() {
return Err(SecurityAdmissionError::ProfileTypeMismatch);
}
None
}
Some(Value::String(raw)) => {
let normalized = normalize_jose_media_type(raw)
.ok_or(SecurityAdmissionError::ProfileTypeMismatch)?;
if Some(normalized.as_str()) != profile.canonical_media_type() {
return Err(SecurityAdmissionError::ProfileTypeMismatch);
}
Some(normalized)
}
Some(_) => return Err(SecurityAdmissionError::ProfileTypeMismatch),
};
let key_id = match members.get("kid") {
None => None,
Some(Value::String(kid)) if kid.len() <= MAX_JWK_KID_BYTES && !kid.is_empty() => {
Some(kid.clone())
}
Some(_) => return Err(SecurityAdmissionError::MalformedJwkMember("kid")),
};
Ok(JwsProtectedHeader {
algorithm: algorithm.clone(),
media_type,
key_id,
})
}
fn claim_str<'a>(
claims: &'a Map<String, Value>,
name: &'static str,
) -> Result<&'a str, SecurityAdmissionError> {
match claims.get(name) {
Some(Value::String(value)) if !value.is_empty() => Ok(value),
_ => Err(SecurityAdmissionError::MissingOrMalformedClaim(name)),
}
}
fn admit_profile_claims(
profile: CompactJwsProfile,
claims: &Map<String, Value>,
) -> Result<(), SecurityAdmissionError> {
for required in profile.required_claims() {
let Some(value) = claims.get(*required) else {
return Err(SecurityAdmissionError::MissingOrMalformedClaim(required));
};
let well_typed = match *required {
"exp" | "iat" | "nbf" => {
matches!(value, Value::Number(number) if number.as_str().len() <= 32)
}
"aud" => match value {
Value::String(entry) => !entry.is_empty(),
Value::Array(entries) => {
!entries.is_empty()
&& entries
.iter()
.all(|entry| matches!(entry, Value::String(text) if !text.is_empty()))
}
_ => false,
},
_ => matches!(value, Value::String(text) if !text.is_empty()),
};
if !well_typed {
return Err(SecurityAdmissionError::MissingOrMalformedClaim(required));
}
}
for forbidden in profile.forbidden_claims() {
if claims.contains_key(*forbidden) {
return Err(SecurityAdmissionError::ForbiddenClaim(forbidden));
}
}
let issuer = claim_str(claims, "iss")?;
let subject = claim_str(claims, "sub")?;
let self_issued = issuer == subject;
let self_audienced = audience_contains(claims, issuer);
match profile {
CompactJwsProfile::OidcIdToken => {
if self_issued {
return Err(SecurityAdmissionError::CrossProfileConfusion);
}
}
CompactJwsProfile::Rfc7523ClientAssertion => {
if !self_issued || self_audienced {
return Err(SecurityAdmissionError::CrossProfileConfusion);
}
}
CompactJwsProfile::BuiltInIssuerSelfVerification => {
if !self_issued || !self_audienced {
return Err(SecurityAdmissionError::CrossProfileConfusion);
}
}
CompactJwsProfile::Rfc9068AccessToken
| CompactJwsProfile::IdentityAssertionJwtAuthorizationGrant => {
}
}
Ok(())
}
fn audience_contains(claims: &Map<String, Value>, candidate: &str) -> bool {
match claims.get("aud") {
Some(Value::String(value)) => value == candidate,
Some(Value::Array(values)) => values
.iter()
.any(|value| matches!(value, Value::String(entry) if entry == candidate)),
_ => false,
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct JwkAdmissionPolicy {
minimum_rsa_modulus_bytes: usize,
maximum_rsa_modulus_bytes: usize,
require_key_id: bool,
}
impl Default for JwkAdmissionPolicy {
fn default() -> Self {
Self {
minimum_rsa_modulus_bytes: MIN_RSA_MODULUS_BYTES,
maximum_rsa_modulus_bytes: MAX_RSA_MODULUS_BYTES,
require_key_id: true,
}
}
}
impl JwkAdmissionPolicy {
#[must_use]
pub const fn minimum_rsa_modulus_bytes(self) -> usize {
self.minimum_rsa_modulus_bytes
}
#[must_use]
pub const fn without_required_key_id(mut self) -> Self {
self.require_key_id = false;
self
}
#[must_use]
pub const fn with_minimum_rsa_modulus_bytes(mut self, bytes: usize) -> Self {
if bytes > self.minimum_rsa_modulus_bytes {
self.minimum_rsa_modulus_bytes = bytes;
}
self
}
}
const PRIVATE_OR_SYMMETRIC_MEMBERS: [&str; 8] = ["d", "p", "q", "dp", "dq", "qi", "oth", "k"];
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct AdmittedRsaPublicJwk {
modulus: Vec<u8>,
exponent: Vec<u8>,
key_id: Option<String>,
algorithm: Option<String>,
}
impl AdmittedRsaPublicJwk {
#[must_use]
pub fn modulus(&self) -> &[u8] {
&self.modulus
}
#[must_use]
pub fn exponent(&self) -> &[u8] {
&self.exponent
}
#[must_use]
pub fn key_id(&self) -> Option<&str> {
self.key_id.as_deref()
}
#[must_use]
pub fn algorithm(&self) -> Option<&str> {
self.algorithm.as_deref()
}
pub fn thumbprint_sha256(&self) -> Result<JwkThumbprintSha256, SecurityAdmissionError> {
let canonical = self.rfc7638_canonical_input();
if canonical.len() > MAX_RFC7638_CANONICAL_INPUT_BYTES {
return Err(SecurityAdmissionError::InputTooLong(
"RFC 7638 canonical input",
));
}
let digest =
fastmcp_core::sha256_bounded(canonical.as_bytes(), MAX_RFC7638_CANONICAL_INPUT_BYTES)
.map_err(|_| SecurityAdmissionError::InputTooLong("RFC 7638 canonical input"))?;
Ok(JwkThumbprintSha256(*digest.as_bytes()))
}
#[must_use]
pub fn rfc7638_canonical_input(&self) -> String {
format!(
r#"{{"e":"{}","kty":"RSA","n":"{}"}}"#,
encode_base64url(&self.exponent),
encode_base64url(&self.modulus),
)
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
pub struct JwkThumbprintSha256([u8; 32]);
impl JwkThumbprintSha256 {
#[must_use]
pub const fn as_bytes(&self) -> &[u8; 32] {
&self.0
}
#[must_use]
pub fn to_base64url(&self) -> String {
encode_base64url(&self.0)
}
}
impl std::fmt::Display for JwkThumbprintSha256 {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str(&self.to_base64url())
}
}
pub fn admit_public_rsa_jwk(
policy: JwkAdmissionPolicy,
members: &Map<String, Value>,
) -> Result<AdmittedRsaPublicJwk, SecurityAdmissionError> {
for private in PRIVATE_OR_SYMMETRIC_MEMBERS {
if members.contains_key(private) {
return Err(SecurityAdmissionError::NonPublicKeyMaterial);
}
}
match members.get("kty") {
Some(Value::String(kty)) if kty == "RSA" => {}
Some(Value::String(kty)) if kty == "oct" => {
return Err(SecurityAdmissionError::NonPublicKeyMaterial);
}
_ => return Err(SecurityAdmissionError::MalformedJwkMember("kty")),
}
let declared_use = match members.get("use") {
None => None,
Some(Value::String(value)) if value == "sig" => Some("sig"),
Some(Value::String(value)) if value == "enc" => {
return Err(SecurityAdmissionError::NotAVerificationKey);
}
_ => return Err(SecurityAdmissionError::MalformedJwkMember("use")),
};
let declared_ops = match members.get("key_ops") {
None => None,
Some(Value::Array(values)) => {
let mut operations = BTreeSet::new();
for value in values {
let Value::String(operation) = value else {
return Err(SecurityAdmissionError::MalformedJwkMember("key_ops"));
};
if !operations.insert(operation.as_str()) {
return Err(SecurityAdmissionError::MalformedJwkMember("key_ops"));
}
}
for private in ["sign", "decrypt", "unwrapKey", "deriveKey", "deriveBits"] {
if operations.contains(private) {
return Err(SecurityAdmissionError::NonPublicKeyMaterial);
}
}
if !operations.contains("verify") {
return Err(SecurityAdmissionError::NotAVerificationKey);
}
Some(operations)
}
_ => return Err(SecurityAdmissionError::MalformedJwkMember("key_ops")),
};
if let (Some(_), Some(operations)) = (declared_use, declared_ops.as_ref()) {
if operations.contains("encrypt") || operations.contains("wrapKey") {
return Err(SecurityAdmissionError::AmbiguousKeyUsage);
}
}
let modulus = decode_base64url_uint(members.get("n"), MAX_RSA_MODULUS_BYTES, "n")?;
let exponent = decode_base64url_uint(members.get("e"), 8, "e")?;
if modulus.len() < policy.minimum_rsa_modulus_bytes
|| modulus.len() > policy.maximum_rsa_modulus_bytes
{
return Err(SecurityAdmissionError::UnsupportedKeyStrength);
}
if exponent.as_slice() != ADMITTED_RSA_PUBLIC_EXPONENT {
return Err(SecurityAdmissionError::UnsupportedKeyStrength);
}
let key_id = match members.get("kid") {
None if policy.require_key_id => {
return Err(SecurityAdmissionError::MalformedJwkMember("kid"));
}
None => None,
Some(Value::String(kid)) if !kid.is_empty() && kid.len() <= MAX_JWK_KID_BYTES => {
Some(kid.clone())
}
Some(_) => return Err(SecurityAdmissionError::MalformedJwkMember("kid")),
};
let algorithm = match members.get("alg") {
None => None,
Some(Value::String(alg)) if alg.starts_with("RS") || alg.starts_with("PS") => {
Some(alg.clone())
}
_ => return Err(SecurityAdmissionError::UnsupportedAlgorithm),
};
Ok(AdmittedRsaPublicJwk {
modulus,
exponent,
key_id,
algorithm,
})
}
pub fn admit_public_jwk_set(
policy: JwkAdmissionPolicy,
bytes: &[u8],
) -> Result<Vec<AdmittedRsaPublicJwk>, SecurityAdmissionError> {
let document = admit_security_document_object(SecurityDocumentKind::JsonWebKeySet, bytes)?;
let Some(Value::Array(entries)) = document.get("keys") else {
return Err(SecurityAdmissionError::MalformedJwkMember("keys"));
};
if entries.len() > MAX_JWK_SET_KEYS {
return Err(SecurityAdmissionError::InputTooLong("JWK set"));
}
let mut admitted = Vec::with_capacity(entries.len());
let mut identifiers = BTreeSet::new();
for entry in entries {
let Value::Object(members) = entry else {
return Err(SecurityAdmissionError::MalformedJwkMember("keys"));
};
let key = admit_public_rsa_jwk(policy, members)?;
if let Some(kid) = key.key_id() {
if !identifiers.insert(kid.to_owned()) {
return Err(SecurityAdmissionError::DuplicateKeyIdentifier);
}
}
admitted.push(key);
}
Ok(admitted)
}
pub fn admit_public_rsa_components(
policy: JwkAdmissionPolicy,
modulus: &[u8],
exponent: &[u8],
key_id: Option<&str>,
) -> Result<AdmittedRsaPublicJwk, SecurityAdmissionError> {
let Some(modulus) = minimal_unsigned_octets(modulus) else {
return Err(SecurityAdmissionError::MalformedJwkMember("n"));
};
let Some(exponent) = minimal_unsigned_octets(exponent) else {
return Err(SecurityAdmissionError::MalformedJwkMember("e"));
};
if modulus.len() < policy.minimum_rsa_modulus_bytes
|| modulus.len() > policy.maximum_rsa_modulus_bytes
{
return Err(SecurityAdmissionError::UnsupportedKeyStrength);
}
if exponent.as_slice() != ADMITTED_RSA_PUBLIC_EXPONENT {
return Err(SecurityAdmissionError::UnsupportedKeyStrength);
}
let key_id = match key_id {
None if policy.require_key_id => {
return Err(SecurityAdmissionError::MalformedJwkMember("kid"));
}
None => None,
Some(kid) if !kid.is_empty() && kid.len() <= MAX_JWK_KID_BYTES => Some(kid.to_owned()),
Some(_) => return Err(SecurityAdmissionError::MalformedJwkMember("kid")),
};
Ok(AdmittedRsaPublicJwk {
modulus,
exponent,
key_id,
algorithm: None,
})
}
fn minimal_unsigned_octets(bytes: &[u8]) -> Option<Vec<u8>> {
let first_significant = bytes.iter().position(|octet| *octet != 0)?;
Some(bytes[first_significant..].to_vec())
}
fn decode_base64url_uint(
value: Option<&Value>,
max_bytes: usize,
member: &'static str,
) -> Result<Vec<u8>, SecurityAdmissionError> {
let Some(Value::String(encoded)) = value else {
return Err(SecurityAdmissionError::MalformedJwkMember(member));
};
let decoded = decode_canonical_base64url(encoded, max_bytes, member)
.map_err(|_| SecurityAdmissionError::MalformedJwkMember(member))?;
if decoded.is_empty() || decoded[0] == 0 {
return Err(SecurityAdmissionError::MalformedJwkMember(member));
}
Ok(decoded)
}
fn encode_base64url(bytes: &[u8]) -> String {
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(bytes)
}
fn decode_canonical_base64url(
encoded: &str,
max_decoded_bytes: usize,
part: &'static str,
) -> Result<Vec<u8>, SecurityAdmissionError> {
let max_encoded = max_decoded_bytes.div_ceil(3).saturating_mul(4);
if encoded.is_empty() || encoded.len() > max_encoded {
return Err(SecurityAdmissionError::InvalidBase64Url(part));
}
if !encoded
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || byte == b'-' || byte == b'_')
{
return Err(SecurityAdmissionError::InvalidBase64Url(part));
}
let decoded = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(encoded)
.map_err(|_| SecurityAdmissionError::InvalidBase64Url(part))?;
if decoded.len() > max_decoded_bytes {
return Err(SecurityAdmissionError::InputTooLong(part));
}
if encode_base64url(&decoded) != encoded {
return Err(SecurityAdmissionError::InvalidBase64Url(part));
}
Ok(decoded)
}