Skip to main content

sz_rust_cli/
validator.rs

1// SPDX-License-Identifier: Apache-2.0
2// Copyright (c) 2024-2026 SZ-Rust Team
3//
4//! 输入合法性校验器
5//!
6//! 对应 design.md 第 2.2.2.6 节,校验用户输入的插件名、表名、字段定义等。
7//! 包含路径遍历防护与代码注入防护。
8
9use crate::error::CliError;
10use crate::field_parser::Field;
11
12/// 输入校验器
13pub struct InputValidator;
14
15/// 危险字符(代码注入防护)
16const DANGEROUS_CHARS: &[char] = &[';', '|', '&', '$', '`', '!', '\n', '\r', '<', '>'];
17
18impl InputValidator {
19    /// 校验插件名称(Rust crate 命名规范)
20    ///
21    /// 规则:小写字母、数字、下划线、连字符,不以数字开头
22    pub fn validate_plugin_name(name: &str) -> Result<(), CliError> {
23        if name.is_empty() {
24            return Err(CliError::InvalidPluginName(
25                "plugin name is empty".to_string(),
26            ));
27        }
28
29        if name.len() > 64 {
30            return Err(CliError::InvalidPluginName(format!(
31                "plugin name '{name}' exceeds 64 characters"
32            )));
33        }
34
35        if name
36            .chars()
37            .next()
38            .map(|c| c.is_ascii_digit())
39            .unwrap_or(false)
40        {
41            return Err(CliError::InvalidPluginName(format!(
42                "plugin name '{name}' starts with a digit"
43            )));
44        }
45
46        if !name
47            .chars()
48            .all(|c| c.is_ascii_lowercase() || c.is_ascii_digit() || c == '_' || c == '-')
49        {
50            return Err(CliError::InvalidPluginName(format!(
51                "plugin name '{name}' contains invalid characters (only lowercase letters, digits, underscores, and hyphens are allowed)"
52            )));
53        }
54
55        Ok(())
56    }
57
58    /// 校验表名(SQL 标识符规范 + 路径遍历防护)
59    pub fn validate_table_name(name: &str) -> Result<(), CliError> {
60        if name.is_empty() {
61            return Err(CliError::FieldParseError("table name is empty".to_string()));
62        }
63
64        if name.contains("..") {
65            return Err(CliError::FieldParseError(format!(
66                "table name '{name}' contains path traversal sequence '..'"
67            )));
68        }
69
70        if name.starts_with('/') || name.starts_with('\\') {
71            return Err(CliError::FieldParseError(format!(
72                "table name '{name}' is an absolute path"
73            )));
74        }
75
76        if name
77            .chars()
78            .next()
79            .map(|c| c.is_ascii_digit())
80            .unwrap_or(false)
81        {
82            return Err(CliError::FieldParseError(format!(
83                "table name '{name}' starts with a digit"
84            )));
85        }
86
87        if !name.chars().all(|c| c.is_ascii_alphanumeric() || c == '_') {
88            return Err(CliError::FieldParseError(format!(
89                "table name '{name}' contains invalid characters (only letters, digits, and underscores are allowed)"
90            )));
91        }
92
93        Ok(())
94    }
95
96    /// 校验字段定义格式 + 注入防护
97    pub fn validate_fields(fields: &str) -> Result<(), CliError> {
98        if fields.is_empty() {
99            return Err(CliError::FieldParseError(
100                "fields definition is empty".to_string(),
101            ));
102        }
103
104        for ch in DANGEROUS_CHARS {
105            if fields.contains(*ch) {
106                return Err(CliError::FieldParseError(format!(
107                    "fields definition contains dangerous character '{ch}'"
108                )));
109            }
110        }
111
112        crate::field_parser::FieldParser::parse(fields)?;
113        Ok(())
114    }
115
116    /// 校验外键字段存在于从表字段定义中
117    pub fn validate_foreign_key(fk: &str, slave_fields: &[Field]) -> Result<(), CliError> {
118        if fk.is_empty() {
119            return Err(CliError::ForeignKeyNotFound(
120                "foreign key is empty".to_string(),
121            ));
122        }
123
124        if !slave_fields.iter().any(|f| f.name == fk) {
125            return Err(CliError::ForeignKeyNotFound(fk.to_string()));
126        }
127
128        Ok(())
129    }
130}
131
132#[cfg(test)]
133mod tests {
134    use super::*;
135
136    #[test]
137    fn test_validate_plugin_name_valid() {
138        assert!(InputValidator::validate_plugin_name("my-plugin").is_ok());
139        assert!(InputValidator::validate_plugin_name("my_plugin").is_ok());
140        assert!(InputValidator::validate_plugin_name("myplugin123").is_ok());
141        assert!(InputValidator::validate_plugin_name("a").is_ok());
142    }
143
144    #[test]
145    fn test_validate_plugin_name_empty() {
146        assert!(InputValidator::validate_plugin_name("").is_err());
147    }
148
149    #[test]
150    fn test_validate_plugin_name_with_space() {
151        let result = InputValidator::validate_plugin_name("my plugin");
152        assert!(result.is_err());
153        assert!(matches!(
154            result.unwrap_err(),
155            CliError::InvalidPluginName(_)
156        ));
157    }
158
159    #[test]
160    fn test_validate_plugin_name_starts_with_digit() {
161        let result = InputValidator::validate_plugin_name("1plugin");
162        assert!(result.is_err());
163    }
164
165    #[test]
166    fn test_validate_plugin_name_uppercase() {
167        let result = InputValidator::validate_plugin_name("MyPlugin");
168        assert!(result.is_err());
169    }
170
171    #[test]
172    fn test_validate_table_name_valid() {
173        assert!(InputValidator::validate_table_name("users").is_ok());
174        assert!(InputValidator::validate_table_name("user_orders").is_ok());
175        assert!(InputValidator::validate_table_name("table123").is_ok());
176    }
177
178    #[test]
179    fn test_validate_table_name_path_traversal() {
180        let result = InputValidator::validate_table_name("../etc/evil");
181        assert!(result.is_err());
182        let err = result.unwrap_err();
183        assert!(err.to_string().contains("path traversal"));
184    }
185
186    #[test]
187    fn test_validate_table_name_absolute_path() {
188        let result = InputValidator::validate_table_name("/etc/passwd");
189        assert!(result.is_err());
190    }
191
192    #[test]
193    fn test_validate_table_name_starts_with_digit() {
194        let result = InputValidator::validate_table_name("123table");
195        assert!(result.is_err());
196    }
197
198    #[test]
199    fn test_validate_fields_valid() {
200        let result = InputValidator::validate_fields("id:i32,name:String,age:i32");
201        assert!(result.is_ok());
202    }
203
204    #[test]
205    fn test_validate_fields_empty() {
206        let result = InputValidator::validate_fields("");
207        assert!(result.is_err());
208    }
209
210    #[test]
211    fn test_validate_fields_injection_semicolon() {
212        let result = InputValidator::validate_fields("name:String;rm -rf /");
213        assert!(result.is_err());
214    }
215
216    #[test]
217    fn test_validate_fields_injection_pipe() {
218        let result = InputValidator::validate_fields("name:String|cat /etc/passwd");
219        assert!(result.is_err());
220    }
221
222    #[test]
223    fn test_validate_fields_injection_ampersand() {
224        let result = InputValidator::validate_fields("name:String&whoami");
225        assert!(result.is_err());
226    }
227
228    #[test]
229    fn test_validate_fields_injection_backtick() {
230        let result = InputValidator::validate_fields("name:String`whoami`");
231        assert!(result.is_err());
232    }
233
234    #[test]
235    fn test_validate_foreign_key_exists() {
236        let fields = vec![
237            Field {
238                name: "id".to_string(),
239                rust_type: "i32".to_string(),
240                sql_type: "INT".to_string(),
241                is_nullable: false,
242                is_primary_key: true,
243                is_indexed: false,
244            },
245            Field {
246                name: "user_id".to_string(),
247                rust_type: "i32".to_string(),
248                sql_type: "INT".to_string(),
249                is_nullable: false,
250                is_primary_key: false,
251                is_indexed: false,
252            },
253        ];
254        assert!(InputValidator::validate_foreign_key("user_id", &fields).is_ok());
255    }
256
257    #[test]
258    fn test_validate_foreign_key_not_exists() {
259        let fields = vec![Field {
260            name: "id".to_string(),
261            rust_type: "i32".to_string(),
262            sql_type: "INT".to_string(),
263            is_nullable: false,
264            is_primary_key: true,
265            is_indexed: false,
266        }];
267        let result = InputValidator::validate_foreign_key("user_id", &fields);
268        assert!(result.is_err());
269        assert!(matches!(
270            result.unwrap_err(),
271            CliError::ForeignKeyNotFound(_)
272        ));
273    }
274
275    #[test]
276    fn test_validate_foreign_key_empty() {
277        let fields = vec![];
278        let result = InputValidator::validate_foreign_key("", &fields);
279        assert!(result.is_err());
280    }
281
282    #[test]
283    fn test_validate_plugin_name_too_long() {
284        let name = "a".repeat(65);
285        let result = InputValidator::validate_plugin_name(&name);
286        assert!(result.is_err());
287        assert!(result.unwrap_err().to_string().contains("exceeds 64"));
288    }
289
290    #[test]
291    fn test_validate_plugin_name_max_length_ok() {
292        let name = "a".repeat(64);
293        let result = InputValidator::validate_plugin_name(&name);
294        assert!(result.is_ok());
295    }
296
297    #[test]
298    fn test_validate_plugin_name_with_dollar() {
299        let result = InputValidator::validate_plugin_name("bad$name");
300        assert!(result.is_err());
301    }
302
303    #[test]
304    fn test_validate_fields_with_dollar() {
305        let result = InputValidator::validate_fields("id:i32:pk$");
306        assert!(result.is_err());
307    }
308}