#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum TraversalOrder {
#[default]
PreOrder,
PostOrder,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum TraversalMode {
#[default]
DepthFirst,
BreadthFirst,
}
#[derive(Debug, Clone)]
pub struct TraversalConfig {
pub order: TraversalOrder,
pub mode: TraversalMode,
pub max_depth: Option<usize>,
pub follow_references: bool,
pub visit_expressions: bool,
pub visit_tensors: bool,
}
impl Default for TraversalConfig {
fn default() -> Self {
Self {
order: TraversalOrder::PreOrder,
mode: TraversalMode::DepthFirst,
max_depth: None,
follow_references: true,
visit_expressions: true,
visit_tensors: true,
}
}
}
impl TraversalConfig {
pub fn new() -> Self {
Self::default()
}
pub fn with_order(mut self, order: TraversalOrder) -> Self {
self.order = order;
self
}
pub fn with_mode(mut self, mode: TraversalMode) -> Self {
self.mode = mode;
self
}
pub fn with_max_depth(mut self, max_depth: usize) -> Self {
self.max_depth = Some(max_depth);
self
}
pub fn with_follow_references(mut self, follow: bool) -> Self {
self.follow_references = follow;
self
}
pub fn with_visit_expressions(mut self, visit: bool) -> Self {
self.visit_expressions = visit;
self
}
pub fn with_visit_tensors(mut self, visit: bool) -> Self {
self.visit_tensors = visit;
self
}
pub fn is_depth_limit_reached(&self, current_depth: usize) -> bool {
if let Some(max) = self.max_depth {
current_depth >= max
} else {
false
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_default_traversal_config() {
let config = TraversalConfig::default();
assert_eq!(config.order, TraversalOrder::PreOrder);
assert_eq!(config.mode, TraversalMode::DepthFirst);
assert!(config.max_depth.is_none());
assert!(config.follow_references);
assert!(config.visit_expressions);
assert!(config.visit_tensors);
}
#[test]
fn test_builder_pattern() {
let config = TraversalConfig::new()
.with_order(TraversalOrder::PostOrder)
.with_mode(TraversalMode::BreadthFirst)
.with_max_depth(5)
.with_follow_references(false)
.with_visit_expressions(false)
.with_visit_tensors(false);
assert_eq!(config.order, TraversalOrder::PostOrder);
assert_eq!(config.mode, TraversalMode::BreadthFirst);
assert_eq!(config.max_depth, Some(5));
assert!(!config.follow_references);
assert!(!config.visit_expressions);
assert!(!config.visit_tensors);
}
#[test]
fn test_depth_limit_check() {
let config = TraversalConfig::new().with_max_depth(3);
assert!(!config.is_depth_limit_reached(0));
assert!(!config.is_depth_limit_reached(2));
assert!(config.is_depth_limit_reached(3));
assert!(config.is_depth_limit_reached(10));
}
#[test]
fn test_no_depth_limit() {
let config = TraversalConfig::new();
assert!(!config.is_depth_limit_reached(0));
assert!(!config.is_depth_limit_reached(1000));
}
#[test]
fn test_traversal_order_default() {
assert_eq!(TraversalOrder::default(), TraversalOrder::PreOrder);
}
#[test]
fn test_traversal_mode_default() {
assert_eq!(TraversalMode::default(), TraversalMode::DepthFirst);
}
#[test]
fn test_traversal_order_equality() {
assert_eq!(TraversalOrder::PreOrder, TraversalOrder::PreOrder);
assert_ne!(TraversalOrder::PreOrder, TraversalOrder::PostOrder);
}
#[test]
fn test_traversal_mode_equality() {
assert_eq!(TraversalMode::DepthFirst, TraversalMode::DepthFirst);
assert_ne!(TraversalMode::DepthFirst, TraversalMode::BreadthFirst);
}
#[test]
fn test_clone() {
let config = TraversalConfig::new().with_max_depth(5);
let cloned = config.clone();
assert_eq!(cloned.max_depth, Some(5));
}
#[test]
fn test_debug() {
let config = TraversalConfig::default();
let debug = format!("{:?}", config);
assert!(debug.contains("TraversalConfig"));
}
}