Skip to main content

sz_orm_core/
cycle_detection.rs

1//! 循环检测 — Eager Loading 多级关联递归安全保护(v2.2.0 B-1)
2//!
3//! 当 Eager Loading 多级关联存在循环引用(如 User→Order→User)时,
4//! [`CycleDetector`] 提供三种策略避免无限递归:
5//!
6//! - [`CyclePolicy::Error`]:检测到循环时返回错误
7//! - [`CyclePolicy::Truncate`]:检测到循环时终止递归,返回已加载部分
8//! - [`CyclePolicy::AllowWithDepthLimit`]:允许循环但限制最大深度
9
10use std::collections::HashSet;
11
12use crate::DbError;
13
14/// 循环检测策略
15///
16/// 控制 [`CycleDetector`] 在检测到循环引用时的行为。
17#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
18pub enum CyclePolicy {
19    /// 检测到循环时返回 `Err`,含循环路径描述
20    Error,
21    /// 检测到循环时终止递归,返回已加载部分(默认策略)
22    #[default]
23    Truncate,
24    /// 允许循环但限制最大递归深度,超限时终止
25    AllowWithDepthLimit(usize),
26}
27
28/// 循环检测器
29///
30/// 按 entity 类型 + 关联名联合去重,避免同类型不同关联误判。
31/// 例如 `User::manager` ≠ `User::orders`,不会误判为循环。
32pub struct CycleDetector {
33    policy: CyclePolicy,
34    visited: HashSet<String>,
35    current_depth: usize,
36    path: Vec<String>,
37}
38
39impl CycleDetector {
40    /// 创建新的循环检测器
41    pub fn new(policy: CyclePolicy) -> Self {
42        Self {
43            policy,
44            visited: HashSet::new(),
45            current_depth: 0,
46            path: Vec::new(),
47        }
48    }
49
50    /// 检查是否可以继续递归
51    ///
52    /// 返回 `Ok(true)` 表示可以继续递归,`Ok(false)` 表示应终止递归,
53    /// `Err(_)` 表示检测到循环且策略为 `Error`。
54    pub fn check(&mut self, entity_type: &str, relation_name: &str) -> Result<bool, DbError> {
55        let key = format!("{}::{}", entity_type, relation_name);
56
57        if self.visited.contains(&key) {
58            return match self.policy {
59                CyclePolicy::Error => {
60                    self.path.push(key.clone());
61                    Err(DbError::InvalidInput(format!(
62                        "检测到循环引用: {}",
63                        self.path.join(" → ")
64                    )))
65                }
66                CyclePolicy::Truncate => Ok(false),
67                CyclePolicy::AllowWithDepthLimit(max_depth) => {
68                    if self.current_depth >= max_depth {
69                        Ok(false)
70                    } else {
71                        Ok(true)
72                    }
73                }
74            };
75        }
76
77        if let CyclePolicy::AllowWithDepthLimit(max_depth) = self.policy {
78            if self.current_depth >= max_depth {
79                return Ok(false);
80            }
81        }
82
83        Ok(true)
84    }
85
86    /// 进入一个关联层级
87    pub fn enter(&mut self, entity_type: &str, relation_name: &str) {
88        let key = format!("{}::{}", entity_type, relation_name);
89        self.visited.insert(key.clone());
90        self.path.push(key);
91        self.current_depth += 1;
92    }
93
94    /// 离开一个关联层级
95    pub fn leave(&mut self) {
96        self.path.pop();
97        self.current_depth = self.current_depth.saturating_sub(1);
98    }
99
100    /// 当前递归深度
101    pub fn depth(&self) -> usize {
102        self.current_depth
103    }
104}
105
106#[cfg(test)]
107mod tests {
108    use super::*;
109
110    #[test]
111    fn test_cycle_policy_default() {
112        assert_eq!(CyclePolicy::default(), CyclePolicy::Truncate);
113    }
114
115    #[test]
116    fn test_cycle_detector_no_cycle() {
117        let mut detector = CycleDetector::new(CyclePolicy::Error);
118        assert!(detector.check("User", "orders").unwrap());
119        detector.enter("User", "orders");
120        assert!(detector.check("Order", "items").unwrap());
121        detector.enter("Order", "items");
122        detector.leave();
123        detector.leave();
124    }
125
126    #[test]
127    fn test_cycle_detector_error_policy() {
128        let mut detector = CycleDetector::new(CyclePolicy::Error);
129        detector.enter("User", "orders");
130        assert!(detector.check("Order", "user").unwrap());
131        detector.enter("Order", "user");
132        let result = detector.check("User", "orders");
133        assert!(result.is_err());
134    }
135
136    #[test]
137    fn test_cycle_detector_truncate_policy() {
138        let mut detector = CycleDetector::new(CyclePolicy::Truncate);
139        detector.enter("User", "orders");
140        detector.enter("Order", "user");
141        let result = detector.check("User", "orders").unwrap();
142        assert!(!result);
143    }
144
145    #[test]
146    fn test_cycle_detector_depth_limit() {
147        let mut detector = CycleDetector::new(CyclePolicy::AllowWithDepthLimit(3));
148        detector.enter("User", "orders");
149        assert_eq!(detector.depth(), 1);
150        detector.enter("Order", "items");
151        assert_eq!(detector.depth(), 2);
152        detector.enter("OrderItem", "product");
153        assert_eq!(detector.depth(), 3);
154        let result = detector.check("Product", "category").unwrap();
155        assert!(!result);
156    }
157
158    #[test]
159    fn test_cycle_detector_different_relation_no_false_positive() {
160        let mut detector = CycleDetector::new(CyclePolicy::Error);
161        detector.enter("User", "orders");
162        let result = detector.check("User", "manager");
163        assert!(result.unwrap());
164    }
165}