Skip to main content

prism_mcp_rs/utils/
validation.rs

1//! Validation utilities
2
3use super::{UtilError, UtilResult};
4use std::collections::HashSet;
5
6/// Validate an email address
7pub fn validate_email(email: &str) -> UtilResult<()> {
8    if email.is_empty() {
9        return Err(UtilError::ValidationFailed(
10            "Email cannot be empty".to_string(),
11        ));
12    }
13
14    let parts: Vec<&str> = email.split('@').collect();
15    if parts.len() != 2 {
16        return Err(UtilError::ValidationFailed(
17            "Email must contain exactly one @".to_string(),
18        ));
19    }
20
21    let (local, domain) = (parts[0], parts[1]);
22
23    if local.is_empty() || local.len() > 64 {
24        return Err(UtilError::ValidationFailed(
25            "Invalid local part".to_string(),
26        ));
27    }
28
29    if domain.is_empty() || domain.len() > 255 || !domain.contains('.') {
30        return Err(UtilError::ValidationFailed(
31            "Invalid domain part".to_string(),
32        ));
33    }
34
35    Ok(())
36}
37
38/// Validate a URL
39pub fn validate_url(url: &str) -> UtilResult<()> {
40    if url.is_empty() {
41        return Err(UtilError::ValidationFailed(
42            "URL cannot be empty".to_string(),
43        ));
44    }
45
46    let valid_schemes = ["http", "https", "ftp", "ftps"];
47    let has_valid_scheme = valid_schemes
48        .iter()
49        .any(|scheme| url.starts_with(&format!("{}://", scheme)));
50
51    if !has_valid_scheme {
52        return Err(UtilError::ValidationFailed(
53            "URL must have a valid scheme".to_string(),
54        ));
55    }
56
57    Ok(())
58}
59
60/// Validate a JSON string
61pub fn validate_json(json_str: &str) -> UtilResult<serde_json::Value> {
62    serde_json::from_str(json_str)
63        .map_err(|e| UtilError::ValidationFailed(format!("Invalid JSON: {}", e)))
64}
65
66/// Validate that a string contains only allowed characters
67pub fn validate_charset(input: &str, allowed_chars: &str) -> UtilResult<()> {
68    let allowed_set: HashSet<char> = allowed_chars.chars().collect();
69
70    for c in input.chars() {
71        if !allowed_set.contains(&c) {
72            return Err(UtilError::ValidationFailed(format!(
73                "Character '{}' is not allowed",
74                c
75            )));
76        }
77    }
78
79    Ok(())
80}
81
82/// Validate string length constraints
83pub fn validate_length(
84    input: &str,
85    min_len: Option<usize>,
86    max_len: Option<usize>,
87) -> UtilResult<()> {
88    let len = input.len();
89
90    if let Some(min) = min_len {
91        if len < min {
92            return Err(UtilError::ValidationFailed(format!(
93                "String too short: {} < {}",
94                len, min
95            )));
96        }
97    }
98
99    if let Some(max) = max_len {
100        if len > max {
101            return Err(UtilError::ValidationFailed(format!(
102                "String too long: {} > {}",
103                len, max
104            )));
105        }
106    }
107
108    Ok(())
109}
110
111/// Validate that a string matches a pattern
112pub fn validate_pattern(input: &str, pattern: &str) -> UtilResult<()> {
113    use regex::Regex;
114
115    let regex = Regex::new(pattern)
116        .map_err(|e| UtilError::ValidationFailed(format!("Invalid regex pattern: {}", e)))?;
117
118    if !regex.is_match(input) {
119        return Err(UtilError::ValidationFailed(format!(
120            "String does not match pattern: {}",
121            pattern
122        )));
123    }
124
125    Ok(())
126}
127
128/// Validate a port number
129pub fn validate_port(port: u16) -> UtilResult<()> {
130    if port == 0 {
131        return Err(UtilError::ValidationFailed("Port cannot be 0".to_string()));
132    }
133    Ok(())
134}
135
136/// Validate an IPv4 address
137pub fn validate_ipv4(ip: &str) -> UtilResult<()> {
138    let parts: Vec<&str> = ip.split('.').collect();
139    if parts.len() != 4 {
140        return Err(UtilError::ValidationFailed(
141            "IPv4 must have 4 octets".to_string(),
142        ));
143    }
144
145    for part in parts {
146        let _octet: u8 = part
147            .parse()
148            .map_err(|_| UtilError::ValidationFailed("Invalid octet".to_string()))?;
149
150        if part.starts_with('0') && part.len() > 1 {
151            return Err(UtilError::ValidationFailed(
152                "Leading zeros not allowed".to_string(),
153            ));
154        }
155    }
156
157    Ok(())
158}
159
160/// Validate a semantic version string
161pub fn validate_semver(version: &str) -> UtilResult<(u32, u32, u32)> {
162    let parts: Vec<&str> = version.split('.').collect();
163    if parts.len() != 3 {
164        return Err(UtilError::ValidationFailed(
165            "Semantic version must have format MAJOR.MINOR.PATCH".to_string(),
166        ));
167    }
168
169    let major = parts[0]
170        .parse::<u32>()
171        .map_err(|_| UtilError::ValidationFailed("Invalid major version".to_string()))?;
172    let minor = parts[1]
173        .parse::<u32>()
174        .map_err(|_| UtilError::ValidationFailed("Invalid minor version".to_string()))?;
175    let patch = parts[2]
176        .parse::<u32>()
177        .map_err(|_| UtilError::ValidationFailed("Invalid patch version".to_string()))?;
178
179    // Check for leading zeros
180    for (i, part) in parts.iter().enumerate() {
181        if part.starts_with('0') && part.len() > 1 {
182            let component = ["major", "minor", "patch"][i];
183            return Err(UtilError::ValidationFailed(format!(
184                "Leading zeros not allowed in {} version",
185                component
186            )));
187        }
188    }
189
190    Ok((major, minor, patch))
191}
192
193#[cfg(test)]
194mod tests {
195    use super::*;
196
197    #[test]
198    fn test_email_validation() {
199        assert!(validate_email("test@example.com").is_ok());
200        assert!(validate_email("invalid.email").is_err());
201        assert!(validate_email("").is_err());
202        assert!(validate_email("@example.com").is_err());
203    }
204
205    #[test]
206    fn test_url_validation() {
207        assert!(validate_url("https://example.com").is_ok());
208        assert!(validate_url("http://localhost:8080").is_ok());
209        assert!(validate_url("not-a-url").is_err());
210        assert!(validate_url("").is_err());
211    }
212
213    #[test]
214    fn test_json_validation() {
215        assert!(validate_json(r#"{"key": "value"}"#).is_ok());
216        assert!(validate_json("[1, 2, 3]").is_ok());
217        assert!(validate_json("invalid json").is_err());
218    }
219
220    #[test]
221    fn test_charset_validation() {
222        assert!(validate_charset("abc123", "abcdefghijklmnopqrstuvwxyz0123456789").is_ok());
223        assert!(validate_charset("abc@123", "abcdefghijklmnopqrstuvwxyz0123456789").is_err());
224    }
225
226    #[test]
227    fn test_length_validation() {
228        assert!(validate_length("hello", Some(3), Some(10)).is_ok());
229        assert!(validate_length("hi", Some(3), Some(10)).is_err());
230        assert!(validate_length("very long string", Some(3), Some(10)).is_err());
231    }
232
233    #[test]
234    fn test_ipv4_validation() {
235        assert!(validate_ipv4("192.168.1.1").is_ok());
236        assert!(validate_ipv4("0.0.0.0").is_ok());
237        assert!(validate_ipv4("256.1.1.1").is_err());
238        assert!(validate_ipv4("192.168.1").is_err());
239    }
240
241    #[test]
242    fn test_semver_validation() {
243        assert_eq!(validate_semver("1.0.0").unwrap(), (1, 0, 0));
244        assert_eq!(validate_semver("0.1.2").unwrap(), (0, 1, 2));
245        assert!(validate_semver("1.0").is_err());
246        assert!(validate_semver("01.0.0").is_err());
247    }
248}