prism_mcp_rs/utils/
validation.rs1use super::{UtilError, UtilResult};
4use std::collections::HashSet;
5
6pub 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
38pub 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
60pub 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
66pub 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
82pub 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
111pub 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
128pub 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
136pub 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
160pub 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 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}