use cedar_policy_core::validator::types::{EntityKind, Type};
use derive_more::derive::Deref;
use smol_str::{SmolStr, ToSmolStr};
use std::collections::HashMap;
use thiserror::Error;
type EntityTypeName = SmolStr;
#[derive(Debug, PartialEq)]
pub(crate) enum AttrSrc {
JwtClaim(TknClaimAttrSrc),
EntityRef(EntityRefAttrSrc),
EntityRefSet(EntityRefSetSrc),
}
#[derive(Debug, PartialEq, Deref)]
pub(crate) struct EntityRefAttrSrc(pub EntityTypeName);
#[derive(Debug, PartialEq, Deref)]
pub(crate) struct EntityRefSetSrc(pub EntityTypeName);
#[derive(Debug, PartialEq)]
pub(crate) struct TknClaimAttrSrc {
pub claim: String,
pub expected_type: ExpectedClaimType,
}
#[derive(Debug, PartialEq)]
pub(crate) enum ExpectedClaimType {
Null,
Bool,
Number,
String,
Array(Box<Self>),
Object(HashMap<SmolStr, Self>),
Extension(SmolStr),
}
impl AttrSrc {
pub(super) fn from_type(attr_name: &str, value: &Type) -> Result<Self, BuildAttrSrcError> {
let attr_src: Self = match value {
Type::Never => {
return Err(
BuildAttrSrcErrorKind::InvalidType(value.clone()).while_building(attr_name)
);
},
Type::Bool(_) => Self::JwtClaim(TknClaimAttrSrc {
claim: attr_name.to_string(),
expected_type: ExpectedClaimType::Bool,
}),
Type::Long => Self::JwtClaim(TknClaimAttrSrc {
claim: attr_name.to_string(),
expected_type: ExpectedClaimType::Number,
}),
Type::String => Self::JwtClaim(TknClaimAttrSrc {
claim: attr_name.to_string(),
expected_type: ExpectedClaimType::String,
}),
Type::Set { element_type } => {
let element_type = element_type.as_ref().ok_or(
BuildAttrSrcErrorKind::SetElementTypeNotDefined.while_building(attr_name),
)?;
let element_type = Self::from_type(attr_name, element_type)?;
match element_type {
AttrSrc::JwtClaim(src) => Self::JwtClaim(TknClaimAttrSrc {
claim: src.claim,
expected_type: ExpectedClaimType::Array(Box::new(src.expected_type)),
}),
AttrSrc::EntityRef(entity_ref_attr_src) => {
Self::EntityRefSet(EntityRefSetSrc(entity_ref_attr_src.to_smolstr()))
},
AttrSrc::EntityRefSet(type_name) => Self::EntityRefSet(type_name),
}
},
Type::Entity(entity_kind) => Self::from_entity_kind(attr_name, entity_kind)?,
Type::Record { attrs, .. } => {
let attrs = attrs
.iter()
.map(|(name, attr_kind)| {
(
name.clone(),
ExpectedClaimType::from(attr_kind.attr_type.as_ref()),
)
})
.collect::<HashMap<SmolStr, ExpectedClaimType>>();
Self::JwtClaim(TknClaimAttrSrc {
claim: attr_name.to_string(),
expected_type: ExpectedClaimType::Object(attrs),
})
},
Type::ExtensionType { name } => Self::JwtClaim(TknClaimAttrSrc {
claim: attr_name.to_string(),
expected_type: ExpectedClaimType::Extension(name.to_smolstr()),
}),
};
Ok(attr_src)
}
fn from_entity_kind(
attr_name: &str,
entity_kind: &EntityKind,
) -> Result<Self, BuildAttrSrcError> {
match entity_kind {
EntityKind::AnyEntity => {
Err(BuildAttrSrcErrorKind::MissingEntityTypeName.while_building(attr_name))
},
EntityKind::Entity(entity_lub) => {
let entity = entity_lub.get_single_entity().ok_or_else(|| {
BuildAttrSrcErrorKind::MissingEntityTypeName.while_building(attr_name)
})?;
Ok(Self::EntityRef(EntityRefAttrSrc(
entity.name().to_smolstr(),
)))
},
}
}
}
#[derive(Debug, Error)]
#[error("failed to build attribute source for {attr_name}: {kind}")]
pub(super) struct BuildAttrSrcError {
attr_name: String,
kind: BuildAttrSrcErrorKind,
}
#[derive(Debug, Error)]
pub(super) enum BuildAttrSrcErrorKind {
#[error("can't use {0:?} as an attribute source")]
InvalidType(Type),
#[error("the type of the elements of a Set was not defined")]
SetElementTypeNotDefined,
#[error("the entity type name was not defined")]
MissingEntityTypeName,
}
impl BuildAttrSrcErrorKind {
fn while_building(self, attr_name: &str) -> BuildAttrSrcError {
BuildAttrSrcError {
attr_name: attr_name.to_string(),
kind: self,
}
}
}
impl From<&Type> for ExpectedClaimType {
fn from(value: &Type) -> Self {
match value {
Type::Never => ExpectedClaimType::Null,
Type::Bool(_) => ExpectedClaimType::Bool,
Type::Long => ExpectedClaimType::Number,
Type::String => ExpectedClaimType::String,
Type::Set { element_type } => element_type
.as_ref()
.map(|t| Self::from(&**t))
.expect("cedar-policy-core type `Type::Set` should always be `Some`, according to the documentation"),
Type::Entity(entity_kind) => {
match entity_kind {
EntityKind::AnyEntity => ExpectedClaimType::Null,
EntityKind::Entity(_entity_lub) => ExpectedClaimType::Null,
}
},
Type::Record { attrs, .. } => {
let attrs = attrs
.iter()
.map(|(name, attr_kind)| {
(name.clone(), Self::from(attr_kind.attr_type.as_ref()))
})
.collect::<HashMap<SmolStr, Self>>();
Self::Object(attrs)
},
Type::ExtensionType { name } => Self::Extension(name.to_smolstr()),
}
}
}