use crate::document::{Document, Item, MatrixList, Node};
use crate::error::{HedlError, HedlResult};
use crate::limits::Limits;
use crate::value::Value;
use std::collections::{BTreeMap, HashMap};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
pub enum ReferenceMode {
#[default]
Strict,
Lenient,
}
impl ReferenceMode {
#[inline]
pub fn is_strict(self) -> bool {
matches!(self, ReferenceMode::Strict)
}
#[inline]
pub fn is_lenient(self) -> bool {
matches!(self, ReferenceMode::Lenient)
}
}
impl From<bool> for ReferenceMode {
fn from(strict: bool) -> Self {
if strict {
ReferenceMode::Strict
} else {
ReferenceMode::Lenient
}
}
}
pub struct TypeRegistry {
by_type: BTreeMap<String, BTreeMap<String, usize>>,
by_id: HashMap<String, Vec<String>>,
total_ids: usize,
}
impl TypeRegistry {
pub fn new() -> Self {
Self {
by_type: BTreeMap::new(),
by_id: HashMap::new(),
total_ids: 0,
}
}
pub fn register(
&mut self,
type_name: &str,
id: &str,
line_num: usize,
limits: &Limits,
) -> HedlResult<()> {
let type_registry = self.by_type.entry(type_name.to_string()).or_default();
if let Some(&prev_line) = type_registry.get(id) {
return Err(HedlError::collision(
format!(
"duplicate ID '{}' in type '{}', previously defined at line {}",
id, type_name, prev_line
),
line_num,
));
}
if self.total_ids >= limits.max_total_ids {
return Err(HedlError::security(
format!(
"total ID registrations {} exceeds limit {}",
self.total_ids, limits.max_total_ids
),
line_num,
));
}
type_registry.insert(id.to_string(), line_num);
self.by_id
.entry(id.to_string())
.or_default()
.push(type_name.to_string());
self.total_ids += 1;
Ok(())
}
pub fn contains_in_type(&self, type_name: &str, id: &str) -> bool {
self.by_type
.get(type_name)
.map(|r| r.contains_key(id))
.unwrap_or(false)
}
pub fn lookup_unqualified(&self, id: &str) -> Option<&[String]> {
self.by_id.get(id).map(|v| v.as_slice())
}
pub fn by_id_iter(&self) -> impl Iterator<Item = (&String, &Vec<String>)> {
self.by_id.iter()
}
}
impl Default for TypeRegistry {
fn default() -> Self {
Self::new()
}
}
fn check_nest_depth(depth: usize, max_depth: usize) -> HedlResult<()> {
if depth > max_depth {
return Err(HedlError::security(
format!(
"NEST hierarchy depth {} exceeds maximum allowed depth {}",
depth, max_depth
),
0,
));
}
Ok(())
}
pub fn register_node(
registries: &mut TypeRegistry,
type_name: &str,
id: &str,
line_num: usize,
limits: &Limits,
) -> HedlResult<()> {
registries.register(type_name, id, line_num, limits)
}
pub fn resolve_references(doc: &Document, mode: ReferenceMode) -> HedlResult<()> {
resolve_references_with_limits(doc, mode, &Limits::default())
}
pub fn resolve_references_with_limits(
doc: &Document,
mode: ReferenceMode,
limits: &Limits,
) -> HedlResult<()> {
let mut registries = TypeRegistry::new();
collect_node_ids(&doc.root, &mut registries, 0, limits)?;
validate_references(&doc.root, ®istries, mode, None, 0, limits.max_nest_depth)
}
fn collect_node_ids(
items: &BTreeMap<String, Item>,
registries: &mut TypeRegistry,
depth: usize,
limits: &Limits,
) -> HedlResult<()> {
check_nest_depth(depth, limits.max_nest_depth)?;
for item in items.values() {
match item {
Item::List(list) => {
collect_list_ids(list, registries, depth, limits)?;
}
Item::Object(obj) => {
collect_node_ids(obj, registries, depth + 1, limits)?;
}
Item::Scalar(_) => {}
}
}
Ok(())
}
fn collect_list_ids(
list: &MatrixList,
registries: &mut TypeRegistry,
depth: usize,
limits: &Limits,
) -> HedlResult<()> {
for node in &list.rows {
registries.register(&list.type_name, &node.id, 0, limits)?; }
for node in &list.rows {
if let Some(children) = node.children() {
for child_list in children.values() {
for child in child_list {
collect_list_ids_from_node(child, registries, depth + 1, limits)?;
}
}
}
}
Ok(())
}
fn collect_list_ids_from_node(
node: &Node,
registries: &mut TypeRegistry,
depth: usize,
limits: &Limits,
) -> HedlResult<()> {
check_nest_depth(depth, limits.max_nest_depth)?;
registries.register(&node.type_name, &node.id, 0, limits)?;
if let Some(children) = node.children() {
for child_list in children.values() {
for child in child_list {
collect_list_ids_from_node(child, registries, depth + 1, limits)?;
}
}
}
Ok(())
}
fn validate_references(
items: &BTreeMap<String, Item>,
registries: &TypeRegistry,
mode: ReferenceMode,
current_type: Option<&str>,
depth: usize,
max_depth: usize,
) -> HedlResult<()> {
check_nest_depth(depth, max_depth)?;
for item in items.values() {
match item {
Item::Scalar(value) => {
validate_value_reference(value, registries, mode, current_type)?;
}
Item::List(list) => {
for node in &list.rows {
validate_node_references(node, registries, mode, depth, max_depth)?;
}
}
Item::Object(obj) => {
validate_references(obj, registries, mode, current_type, depth + 1, max_depth)?;
}
}
}
Ok(())
}
fn validate_node_references(
node: &Node,
registries: &TypeRegistry,
mode: ReferenceMode,
depth: usize,
max_depth: usize,
) -> HedlResult<()> {
check_nest_depth(depth, max_depth)?;
for value in &node.fields {
validate_value_reference(value, registries, mode, Some(&node.type_name))?;
}
if let Some(children) = node.children() {
for child_list in children.values() {
for child in child_list {
validate_node_references(child, registries, mode, depth + 1, max_depth)?;
}
}
}
Ok(())
}
fn validate_value_reference(
value: &Value,
registries: &TypeRegistry,
mode: ReferenceMode,
current_type: Option<&str>,
) -> HedlResult<()> {
if let Value::Reference(ref_val) = value {
let resolved = match &ref_val.type_name {
Some(t) => registries.contains_in_type(t, &ref_val.id),
None => {
match current_type {
Some(type_name) => registries.contains_in_type(type_name, &ref_val.id),
None => {
let matching_types =
registries.lookup_unqualified(&ref_val.id).unwrap_or(&[]);
match matching_types.len() {
0 => false, 1 => true, _ => {
return Err(HedlError::reference(
format!(
"Ambiguous unqualified reference '@{}' matches multiple types: [{}]",
ref_val.id,
matching_types.join(", ")
),
0, ));
}
}
}
}
}
};
if !resolved && mode.is_strict() {
return Err(HedlError::reference(
format!("unresolved reference {}", ref_val.to_ref_string()),
0, ));
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_max_total_ids_limit() {
let mut registry = TypeRegistry::new();
let limits = Limits {
max_total_ids: 3,
..Default::default()
};
assert!(registry.register("Type1", "id1", 1, &limits).is_ok());
assert!(registry.register("Type2", "id2", 2, &limits).is_ok());
assert!(registry.register("Type3", "id3", 3, &limits).is_ok());
let result = registry.register("Type4", "id4", 4, &limits);
assert!(result.is_err());
let err_msg = result.unwrap_err().to_string();
assert!(
err_msg.contains("exceeds limit"),
"Expected 'exceeds limit' in error message, got: {}",
err_msg
);
}
#[test]
fn test_max_total_ids_across_types() {
let mut registry = TypeRegistry::new();
let limits = Limits {
max_total_ids: 10,
..Default::default()
};
for i in 0..5 {
assert!(registry
.register("Type1", &format!("id{}", i), i, &limits)
.is_ok());
}
for i in 0..5 {
assert!(registry
.register("Type2", &format!("id{}", i), i + 5, &limits)
.is_ok());
}
let result = registry.register("Type3", "id_extra", 10, &limits);
assert!(result.is_err());
let err_msg = result.unwrap_err().to_string();
assert!(
err_msg.contains("exceeds limit"),
"Expected 'exceeds limit' in error message, got: {}",
err_msg
);
}
#[test]
fn test_unlimited_ids() {
let mut registry = TypeRegistry::new();
let limits = Limits::unlimited();
for i in 0..10000 {
let result =
registry.register(&format!("Type{}", i % 100), &format!("id{}", i), i, &limits);
assert!(
result.is_ok(),
"Failed to register ID {} in unlimited mode",
i
);
}
}
#[test]
fn test_collision_detection_with_limits() {
let mut registry = TypeRegistry::new();
let limits = Limits::default();
assert!(registry.register("Type1", "id1", 1, &limits).is_ok());
let result = registry.register("Type1", "id1", 2, &limits);
assert!(result.is_err());
let err_msg = result.unwrap_err().to_string();
assert!(
err_msg.contains("duplicate"),
"Expected 'duplicate' in error message, got: {}",
err_msg
);
}
#[test]
fn test_max_total_ids_exact_limit() {
let mut registry = TypeRegistry::new();
let limits = Limits {
max_total_ids: 5,
..Default::default()
};
for i in 0..5 {
let result = registry.register("Type", &format!("id{}", i), i, &limits);
assert!(result.is_ok(), "Failed to register ID {} at exact limit", i);
}
let result = registry.register("Type", "id5", 5, &limits);
assert!(result.is_err());
let err_msg = result.unwrap_err().to_string();
assert!(
err_msg.contains("exceeds limit"),
"Expected 'exceeds limit' in error message, got: {}",
err_msg
);
}
#[test]
fn test_max_total_ids_just_under_limit() {
let mut registry = TypeRegistry::new();
let limits = Limits {
max_total_ids: 5,
..Default::default()
};
for i in 0..4 {
assert!(registry
.register("Type", &format!("id{}", i), i, &limits)
.is_ok());
}
assert!(registry.register("Type", "id4", 4, &limits).is_ok());
}
#[test]
fn test_max_total_ids_error_message_clarity() {
let mut registry = TypeRegistry::new();
let limits = Limits {
max_total_ids: 2,
..Default::default()
};
registry.register("Type1", "id1", 1, &limits).unwrap();
registry.register("Type2", "id2", 2, &limits).unwrap();
let result = registry.register("Type3", "id3", 3, &limits);
assert!(result.is_err());
let err_msg = result.unwrap_err().to_string();
assert!(
err_msg.contains("2"),
"Error message should contain the count"
);
assert!(
err_msg.contains("limit"),
"Error message should mention 'limit'"
);
}
#[test]
fn test_total_ids_count_tracking() {
let mut registry = TypeRegistry::new();
let limits = Limits::unlimited();
assert_eq!(registry.total_ids, 0);
registry.register("Type1", "id1", 1, &limits).unwrap();
assert_eq!(registry.total_ids, 1);
registry.register("Type2", "id2", 2, &limits).unwrap();
assert_eq!(registry.total_ids, 2);
registry.register("Type1", "id3", 3, &limits).unwrap();
assert_eq!(registry.total_ids, 3);
}
#[test]
fn test_max_total_ids_with_multiple_types() {
let mut registry = TypeRegistry::new();
let limits = Limits {
max_total_ids: 100,
..Default::default()
};
for type_idx in 0..10 {
for id_idx in 0..10 {
let result = registry.register(
&format!("Type{}", type_idx),
&format!("id{}_{}", type_idx, id_idx),
type_idx * 10 + id_idx,
&limits,
);
assert!(result.is_ok(), "Failed at type {} id {}", type_idx, id_idx);
}
}
let result = registry.register("TypeExtra", "extra", 100, &limits);
assert!(result.is_err());
}
#[test]
fn test_collision_preserves_total_count() {
let mut registry = TypeRegistry::new();
let limits = Limits::unlimited();
registry.register("Type1", "id1", 1, &limits).unwrap();
assert_eq!(registry.total_ids, 1);
let result = registry.register("Type1", "id1", 2, &limits);
assert!(result.is_err());
assert_eq!(registry.total_ids, 1);
}
#[test]
fn test_default_limits_max_total_ids() {
let limits = Limits::default();
assert_eq!(limits.max_total_ids, 10_000_000);
}
#[test]
fn test_unlimited_limits_max_total_ids() {
let limits = Limits::unlimited();
assert_eq!(limits.max_total_ids, usize::MAX);
}
#[test]
fn test_registry_new() {
let registry = TypeRegistry::new();
assert_eq!(registry.total_ids, 0);
assert!(registry.by_type.is_empty());
assert!(registry.by_id.is_empty());
}
#[test]
fn test_registry_default() {
let registry = TypeRegistry::default();
assert_eq!(registry.total_ids, 0);
}
#[test]
fn test_contains_in_type() {
let mut registry = TypeRegistry::new();
let limits = Limits::unlimited();
registry.register("User", "u1", 1, &limits).unwrap();
assert!(registry.contains_in_type("User", "u1"));
assert!(!registry.contains_in_type("User", "u2"));
assert!(!registry.contains_in_type("Post", "u1"));
}
#[test]
fn test_lookup_unqualified() {
let mut registry = TypeRegistry::new();
let limits = Limits::unlimited();
registry.register("User", "id1", 1, &limits).unwrap();
registry.register("Post", "id1", 2, &limits).unwrap();
let types = registry.lookup_unqualified("id1");
assert!(types.is_some());
let types = types.unwrap();
assert_eq!(types.len(), 2);
assert!(types.contains(&"User".to_string()));
assert!(types.contains(&"Post".to_string()));
let not_found = registry.lookup_unqualified("nonexistent");
assert!(not_found.is_none());
}
#[test]
fn test_inverted_index_maintenance() {
let mut registry = TypeRegistry::new();
let limits = Limits::unlimited();
registry.register("Type1", "shared_id", 1, &limits).unwrap();
registry.register("Type2", "shared_id", 2, &limits).unwrap();
registry.register("Type3", "shared_id", 3, &limits).unwrap();
let types = registry.lookup_unqualified("shared_id").unwrap();
assert_eq!(types.len(), 3);
}
#[test]
fn test_reference_mode_default() {
assert_eq!(ReferenceMode::default(), ReferenceMode::Strict);
}
#[test]
fn test_reference_mode_from_bool() {
assert_eq!(ReferenceMode::from(true), ReferenceMode::Strict);
assert_eq!(ReferenceMode::from(false), ReferenceMode::Lenient);
}
#[test]
fn test_reference_mode_is_strict() {
assert!(ReferenceMode::Strict.is_strict());
assert!(!ReferenceMode::Lenient.is_strict());
}
#[test]
fn test_reference_mode_is_lenient() {
assert!(ReferenceMode::Lenient.is_lenient());
assert!(!ReferenceMode::Strict.is_lenient());
}
#[test]
fn test_reference_mode_equality() {
assert_eq!(ReferenceMode::Strict, ReferenceMode::Strict);
assert_eq!(ReferenceMode::Lenient, ReferenceMode::Lenient);
assert_ne!(ReferenceMode::Strict, ReferenceMode::Lenient);
}
}