#[derive(Debug, thiserror::Error)]
pub(crate) enum DefaultEntitiesLimitsError {
#[error(
"Cedar entity data size ({size}) for default entity '{entity_id}' exceeds maximum allowed size ({max_size})"
)]
DataSizeExceeded {
entity_id: String,
size: usize,
max_size: usize,
},
#[error("Maximum number of default entities ({max_entities}) exceeded, found {found}")]
CountExceeded { max_entities: usize, found: usize },
}
#[derive(Debug, Clone)]
pub(super) struct DefaultEntitiesLimits {
pub max_entities: usize,
pub max_entity_size: usize,
}
impl Default for DefaultEntitiesLimits {
fn default() -> Self {
Self {
max_entities: Self::DEFAULT_MAX_ENTITIES,
max_entity_size: Self::DEFAULT_MAX_ENTITY_SIZE,
}
}
}
impl DefaultEntitiesLimits {
pub(super) const DEFAULT_MAX_ENTITIES: usize = 1000;
pub(super) const DEFAULT_MAX_ENTITY_SIZE: usize = 1024 * 1024;
fn validate_default_entity_data_size(
&self,
entity_id: &str,
entity_str: &str,
) -> Result<(), DefaultEntitiesLimitsError> {
if entity_str.len() > self.max_entity_size {
Err(DefaultEntitiesLimitsError::DataSizeExceeded {
entity_id: entity_id.to_string(),
size: entity_str.len(),
max_size: self.max_entity_size,
})
} else {
Ok(())
}
}
pub(super) fn validate_default_entity(
&self,
entity_id: &str,
entity_data: &serde_json::Value,
) -> Result<(), DefaultEntitiesLimitsError> {
if let Some(entity_str) = entity_data.as_str() {
self.validate_default_entity_data_size(entity_id, entity_str)
} else {
Ok(())
}
}
pub(super) fn validate_entities_count(
&self,
entity_count: usize,
) -> Result<(), DefaultEntitiesLimitsError> {
if entity_count > self.max_entities {
Err(DefaultEntitiesLimitsError::CountExceeded {
max_entities: self.max_entities,
found: entity_count,
})
} else {
Ok(())
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn test_validate_default_entity_data_size() {
let limits = DefaultEntitiesLimits {
max_entities: 2,
max_entity_size: 100,
};
let result = limits.validate_default_entity_data_size("entity1", "dGVzdA==");
assert!(result.is_ok());
let large_base64 = "dGVzdA==".repeat(20); let result = limits.validate_default_entity_data_size("entity1", &large_base64);
assert!(result.is_err());
assert!(
result
.unwrap_err()
.to_string()
.contains("exceeds maximum allowed size")
);
}
#[test]
fn test_validate_default_entity() {
let limits = DefaultEntitiesLimits {
max_entities: 2,
max_entity_size: 100,
};
let valid_entity = json!("dGVzdA==");
let result = limits.validate_default_entity("entity1", &valid_entity);
assert!(result.is_ok());
let non_string_entity = json!({ "key": "value" });
let result = limits.validate_default_entity("entity1", &non_string_entity);
assert!(result.is_ok());
let large_base64 = "dGVzdA==".repeat(20); let large_entity = json!(large_base64);
let result = limits.validate_default_entity("entity1", &large_entity);
assert!(
result
.unwrap_err()
.to_string()
.contains("exceeds maximum allowed size")
);
}
#[test]
fn test_validate_entities_count() {
let limits = DefaultEntitiesLimits {
max_entities: 2,
max_entity_size: 100,
};
let result = limits.validate_entities_count(2);
assert!(result.is_ok());
let result = limits.validate_entities_count(3);
assert!(
result
.unwrap_err()
.to_string()
.contains("Maximum number of default entities (2) exceeded")
);
}
}