1use crate::proximity_keywords::compile_keywords_proximity_config;
2use crate::scanner::config::RuleConfig;
3use crate::scanner::metrics::RuleMetrics;
4use crate::scanner::regex_rule::compiled::RegexCompiledRule;
5use crate::scanner::regex_rule::regex_store::get_memoized_regex;
6use crate::validation::{
7 RegexPatternCaptureGroupsValidationError, validate_and_create_regex,
8 validate_named_capture_group_minimum_length,
9};
10use crate::{CompiledRule, CreateScannerError, Labels};
11use regex_automata::util::captures::GroupInfo;
12use serde::{Deserialize, Serialize};
13use serde_with::DefaultOnNull;
14use serde_with::serde_as;
15use std::sync::Arc;
16use strum::{AsRefStr, EnumIter};
17
18pub const DEFAULT_KEYWORD_LOOKAHEAD: usize = 30;
19
20#[serde_as]
21#[derive(Serialize, Deserialize, Clone, Debug, PartialEq)]
22pub struct RegexRuleConfig {
23 pub pattern: String,
24 pub proximity_keywords: Option<ProximityKeywordsConfig>,
25 pub validator: Option<SecondaryValidator>,
26 #[serde_as(deserialize_as = "DefaultOnNull")]
27 #[serde(default)]
28 pub labels: Labels,
29 pub pattern_capture_groups: Option<Vec<String>>,
30}
31
32impl RegexRuleConfig {
33 pub fn new(pattern: &str) -> Self {
34 #[allow(deprecated)]
35 Self {
36 pattern: pattern.to_owned(),
37 proximity_keywords: None,
38 validator: None,
39 labels: Labels::default(),
40 pattern_capture_groups: None,
41 }
42 }
43
44 pub fn with_pattern(&self, pattern: &str) -> Self {
45 self.mutate_clone(|x| x.pattern = pattern.to_string())
46 }
47
48 pub fn with_proximity_keywords(&self, proximity_keywords: ProximityKeywordsConfig) -> Self {
49 self.mutate_clone(|x| x.proximity_keywords = Some(proximity_keywords))
50 }
51
52 pub fn with_labels(&self, labels: Labels) -> Self {
53 self.mutate_clone(|x| x.labels = labels)
54 }
55
56 pub fn with_pattern_capture_groups(&self, pattern_capture_groups: Vec<String>) -> Self {
57 self.mutate_clone(|x| x.pattern_capture_groups = Some(pattern_capture_groups))
58 }
59
60 pub fn with_pattern_capture_group(&self, pattern_capture_group: &str) -> Self {
61 self.mutate_clone(|x| match x.pattern_capture_groups {
62 Some(ref mut pattern_capture_groups) => {
63 pattern_capture_groups.push(pattern_capture_group.to_string());
64 }
65 None => {
66 x.pattern_capture_groups = Some(vec![pattern_capture_group.to_string()]);
67 }
68 })
69 }
70
71 pub fn build(&self) -> Arc<dyn RuleConfig> {
72 Arc::new(self.clone())
73 }
74
75 fn mutate_clone(&self, modify: impl FnOnce(&mut Self)) -> Self {
76 let mut clone = self.clone();
77 modify(&mut clone);
78 clone
79 }
80
81 pub fn with_included_keywords(
82 &self,
83 keywords: impl IntoIterator<Item = impl AsRef<str>>,
84 ) -> Self {
85 let mut this = self.clone();
86 let mut config = self.get_or_create_proximity_keywords_config();
87 config.included_keywords = keywords
88 .into_iter()
89 .map(|x| x.as_ref().to_string())
90 .collect::<Vec<_>>();
91 this.proximity_keywords = Some(config);
92 this
93 }
94
95 pub fn with_excluded_keywords(
96 &self,
97 keywords: impl IntoIterator<Item = impl AsRef<str>>,
98 ) -> Self {
99 let mut this = self.clone();
100 let mut config = self.get_or_create_proximity_keywords_config();
101 config.excluded_keywords = keywords
102 .into_iter()
103 .map(|x| x.as_ref().to_string())
104 .collect::<Vec<_>>();
105 this.proximity_keywords = Some(config);
106 this
107 }
108
109 pub fn with_validator(&self, validator: Option<SecondaryValidator>) -> Self {
110 let mut this = self.clone();
111 this.validator = validator;
112 this
113 }
114
115 fn get_or_create_proximity_keywords_config(&self) -> ProximityKeywordsConfig {
116 self.proximity_keywords
117 .clone()
118 .unwrap_or_else(|| ProximityKeywordsConfig {
119 look_ahead_character_count: DEFAULT_KEYWORD_LOOKAHEAD,
120 included_keywords: vec![],
121 excluded_keywords: vec![],
122 })
123 }
124}
125
126fn is_pattern_capture_groups_valid(
127 pattern: &str,
128 pattern_capture_groups: &Option<Vec<String>>,
129 group_info: &GroupInfo,
130) -> Result<(), RegexPatternCaptureGroupsValidationError> {
131 if pattern_capture_groups.is_none() {
132 return Ok(());
133 }
134 let pattern_capture_groups = pattern_capture_groups.as_ref().unwrap();
135 if pattern_capture_groups.len() != 1 {
136 return Err(
138 RegexPatternCaptureGroupsValidationError::TooManyCaptureGroups(
139 pattern_capture_groups.len(),
140 ),
141 );
142 }
143 let pattern_capture_group = pattern_capture_groups.first().unwrap();
144 if !group_info
145 .all_names()
146 .filter(|(_, _, name)| name.is_some())
147 .map(|(_, _, name)| name.unwrap())
148 .any(|name| name == pattern_capture_group)
149 {
150 return Err(
151 RegexPatternCaptureGroupsValidationError::CaptureGroupNotPresent(
152 pattern_capture_group.clone(),
153 ),
154 );
155 }
156 if pattern_capture_group != "sds_match" {
159 return Err(RegexPatternCaptureGroupsValidationError::TargetedCaptureGroupMustBeSdsMatch);
160 }
161 validate_named_capture_group_minimum_length(pattern, pattern_capture_group)?;
162 Ok(())
163}
164
165impl RuleConfig for RegexRuleConfig {
166 fn convert_to_compiled_rule(
167 &self,
168 rule_index: usize,
169 scanner_labels: Labels,
170 ) -> Result<Box<dyn CompiledRule>, CreateScannerError> {
171 let regex = get_memoized_regex(&self.pattern, validate_and_create_regex)?;
172
173 let rule_labels = scanner_labels.clone_with_labels(self.labels.clone());
174
175 let (included_keywords, excluded_keywords) = self
176 .proximity_keywords
177 .as_ref()
178 .map(|config| compile_keywords_proximity_config(config, &rule_labels))
179 .unwrap_or(Ok((None, None)))?;
180
181 is_pattern_capture_groups_valid(
182 &self.pattern,
183 &self.pattern_capture_groups,
184 regex.group_info(),
185 )?;
186
187 Ok(Box::new(RegexCompiledRule {
188 rule_index,
189 regex,
190 included_keywords,
191 excluded_keywords,
192 validator: self.validator.clone().map(|x| x.compile()),
193 metrics: RuleMetrics::new(&rule_labels),
194 pattern_capture_groups: self.pattern_capture_groups.clone(),
195 }))
196 }
197
198 fn as_regex_rule(&self) -> Option<&RegexRuleConfig> {
199 Some(self)
200 }
201}
202
203#[serde_as]
204#[derive(Serialize, Deserialize, Clone, Debug, PartialEq)]
205pub struct ProximityKeywordsConfig {
206 pub look_ahead_character_count: usize,
207
208 #[serde_as(deserialize_as = "DefaultOnNull")]
209 #[serde(default)]
210 pub included_keywords: Vec<String>,
211
212 #[serde_as(deserialize_as = "DefaultOnNull")]
213 #[serde(default)]
214 pub excluded_keywords: Vec<String>,
215}
216
217#[derive(Serialize, Deserialize, Clone, Debug, PartialEq, EnumIter, AsRefStr)]
218#[serde(tag = "type")]
219pub enum SecondaryValidator {
220 AbaRtnChecksum,
221 AtlassianTokenChecksum,
222 AustralianMedicareChecksum,
223 AustralianTfnChecksum,
224 AustrianSSNChecksum,
225 BelgiumNationalRegisterChecksum,
226 BrazilianCnpjChecksum,
227 BrazilianCpfChecksum,
228 BtcChecksum,
229 BulgarianEGNChecksum,
230 ChineseIdChecksum,
231 CoordinationNumberChecksum,
232 CzechPersonalIdentificationNumberChecksum,
233 CzechTaxIdentificationNumberChecksum,
234 DutchBsnChecksum,
235 DutchPassportChecksum,
236 EntropyCheck,
237 EstoniaPersonalCodeChecksum,
238 EthereumChecksum,
239 FinnishHetuChecksum,
240 FranceNifChecksum,
241 FranceSsnChecksum,
242 GermanIdsChecksum,
243 GermanSvnrChecksum,
244 GithubTokenChecksum,
245 GreeceAmkaChecksum,
246 GreekTinChecksum,
247 HungarianTinChecksum,
248 IbanChecker,
249 IrishPpsChecksum,
250 ItalianNationalIdChecksum,
251 JwtClaimsValidator { config: JwtClaimsValidatorConfig },
252 JwtExpirationChecker,
253 LatviaNationalIdChecksum,
254 LithuanianPersonalIdentificationNumberChecksum,
255 LuhnChecksum,
256 LuxembourgIndividualNINChecksum,
257 Mod11_10checksum,
258 Mod11_2checksum,
259 Mod1271_36Checksum,
260 Mod27_26checksum,
261 Mod37_2checksum,
262 Mod37_36checksum,
263 Mod661_26checksum,
264 Mod97_10checksum,
265 MoneroAddress,
266 NhsCheckDigit,
267 NirChecksum,
268 NonHexChecker,
269 PolishNationalIdChecksum,
270 PolishNipChecksum,
271 PortugueseTaxIdChecksum,
272 RodneCisloNumberChecksum,
273 RomanianPersonalNumericCode,
274 SingaporeNricChecksum,
275 SloveniaTinChecksum,
276 SlovenianPINChecksum,
277 SpanishDniChecksum,
278 SpanishNussChecksum,
279 SwedenPINChecksum,
280 UkNinoFormatCheck,
281 UkTrnChecksum,
282 UsDeaChecksum,
283 UsNpiChecksum,
284 VerhoeffChecksum,
285}
286
287#[derive(Serialize, Deserialize, Clone, Debug, PartialEq)]
288#[serde(tag = "type", content = "config")]
289pub enum ClaimRequirement {
290 Present,
292 NotExpired,
294 ExactValue(String),
296 RegexMatch(String),
298}
299
300#[derive(Serialize, Deserialize, Default, Clone, Debug, PartialEq)]
301pub struct JwtClaimsValidatorConfig {
302 #[serde(default)]
303 pub required_headers: std::collections::BTreeMap<String, ClaimRequirement>,
304 #[serde(default)]
305 pub required_claims: std::collections::BTreeMap<String, ClaimRequirement>,
306}
307
308#[cfg(test)]
309mod test {
310 use crate::{AwsType, CustomHttpConfig, MatchAction, MatchValidationType, RootRuleConfig};
311 use std::collections::BTreeMap;
312 use strum::IntoEnumIterator;
313
314 use super::*;
315
316 #[test]
317 fn should_override_pattern() {
318 let rule_config = RegexRuleConfig::new("123").with_pattern("456");
319 assert_eq!(rule_config.pattern, "456");
320 }
321
322 #[test]
323 #[allow(deprecated)]
324 fn should_have_default() {
325 let rule_config = RegexRuleConfig::new("123");
326 assert_eq!(
327 rule_config,
328 RegexRuleConfig {
329 pattern: "123".to_string(),
330 proximity_keywords: None,
331 validator: None,
332 labels: Labels::empty(),
333 pattern_capture_groups: None,
334 }
335 );
336 }
337
338 #[test]
339 fn should_use_capture_group() {
340 let rule_config = RegexRuleConfig::new("hey (?<capture_group>world)")
341 .with_pattern_capture_groups(vec!["capture_group".to_string()]);
342 assert_eq!(
343 rule_config,
344 RegexRuleConfig {
345 pattern: "hey (?<capture_group>world)".to_string(),
346 proximity_keywords: None,
347 validator: None,
348 labels: Labels::empty(),
349 pattern_capture_groups: Some(vec!["capture_group".to_string()]),
350 }
351 );
352 }
353
354 #[test]
355 fn match_action_should_default_to_none_on_deserialization() {
356 let config: RootRuleConfig<RegexRuleConfig> =
357 serde_json::from_str(r#"{"pattern":"hello"}"#).unwrap();
358 assert_eq!(config.match_action, MatchAction::None);
359 }
360
361 #[test]
362 fn proximity_keywords_should_have_default() {
363 let json_config = r#"{"look_ahead_character_count": 0}"#;
364 let test: ProximityKeywordsConfig = serde_json::from_str(json_config).unwrap();
365 assert_eq!(
366 test,
367 ProximityKeywordsConfig {
368 look_ahead_character_count: 0,
369 included_keywords: vec![],
370 excluded_keywords: vec![]
371 }
372 );
373
374 let json_config = r#"{"look_ahead_character_count": 0, "excluded_keywords": null, "included_keywords": null}"#;
375 let test: ProximityKeywordsConfig = serde_json::from_str(json_config).unwrap();
376 assert_eq!(
377 test,
378 ProximityKeywordsConfig {
379 look_ahead_character_count: 0,
380 included_keywords: vec![],
381 excluded_keywords: vec![]
382 }
383 );
384 }
385
386 #[test]
387 #[allow(deprecated)]
388 fn test_third_party_active_checker() {
389 let http_config = CustomHttpConfig::default().with_endpoint("http://test.com".to_string());
391 let validation_type = MatchValidationType::CustomHttp(http_config.clone());
392 let rule_config = RootRuleConfig::new(RegexRuleConfig::new("123"))
393 .third_party_active_checker(validation_type.clone());
394
395 assert_eq!(
396 rule_config.third_party_active_checker,
397 Some(validation_type.clone())
398 );
399 assert_eq!(rule_config.match_validation_type, None);
400 assert_eq!(
401 rule_config.get_third_party_active_checker(),
402 Some(&validation_type)
403 );
404
405 let aws_type = AwsType::AwsId;
407 let validation_type2 = MatchValidationType::Aws(aws_type);
408 let rule_config = RootRuleConfig::new(RegexRuleConfig::new("123"))
409 .third_party_active_checker(validation_type2.clone());
410
411 assert_eq!(
412 rule_config.third_party_active_checker,
413 Some(validation_type2.clone())
414 );
415 assert_eq!(
416 rule_config.get_third_party_active_checker(),
417 Some(&validation_type2)
418 );
419
420 let rule_config = RootRuleConfig::new(RegexRuleConfig::new("123"))
422 .third_party_active_checker(MatchValidationType::CustomHttp(http_config.clone()));
423
424 assert_eq!(
425 rule_config.get_third_party_active_checker(),
426 Some(&MatchValidationType::CustomHttp(http_config.clone()))
427 );
428 }
429
430 #[test]
431 fn test_secondary_validator_enum_iter() {
432 let validators: Vec<SecondaryValidator> = SecondaryValidator::iter().collect();
434 assert!(validators.contains(&SecondaryValidator::GithubTokenChecksum));
436 assert!(validators.contains(&SecondaryValidator::JwtExpirationChecker));
437 }
438
439 #[test]
440 fn test_secondary_validator_are_sorted() {
441 let validator_names: Vec<String> = SecondaryValidator::iter()
442 .map(|a| a.as_ref().to_string())
443 .collect();
444 let mut sorted_validator_names = validator_names.clone();
445 sorted_validator_names.sort();
446 assert_eq!(
447 sorted_validator_names, validator_names,
448 "Secondary validators should be sorted by alphabetical order, but it's not the case, expected order:"
449 );
450 }
451
452 #[test]
454 fn test_jwt_claims_validator_config_serialization_order() {
455 let mut required_claims = BTreeMap::new();
457 required_claims.insert("zzz".to_string(), ClaimRequirement::Present);
458 required_claims.insert("exp".to_string(), ClaimRequirement::NotExpired);
459 required_claims.insert(
460 "aaa".to_string(),
461 ClaimRequirement::ExactValue("test".to_string()),
462 );
463 required_claims.insert(
464 "mmm".to_string(),
465 ClaimRequirement::RegexMatch(r"^test.*".to_string()),
466 );
467
468 let config = JwtClaimsValidatorConfig {
469 required_claims,
470 required_headers: std::collections::BTreeMap::new(),
471 };
472
473 let serialized1 = serde_json::to_string(&config).unwrap();
475 let serialized2 = serde_json::to_string(&config).unwrap();
476
477 assert_eq!(serialized1, serialized2, "Serialization should be stable");
479
480 assert!(serialized1.find("aaa").unwrap() < serialized1.find("exp").unwrap());
482 assert!(serialized1.find("exp").unwrap() < serialized1.find("mmm").unwrap());
483 assert!(serialized1.find("mmm").unwrap() < serialized1.find("zzz").unwrap());
484 }
485
486 #[test]
487 fn test_capture_groups_validation() {
488 let test_cases: Vec<(
489 &str,
490 Vec<String>,
491 Result<(), RegexPatternCaptureGroupsValidationError>,
492 )> = vec![
493 (
494 "hello (?<sds_match>world)",
495 vec!["sds_match".to_string()],
496 Ok(()),
497 ),
498 (
499 "hello (?<capture_group>world)",
500 vec!["capture_group".to_string()],
501 Err(RegexPatternCaptureGroupsValidationError::TargetedCaptureGroupMustBeSdsMatch),
502 ),
503 (
504 "hello (?<sds_match>world) and (?<another_group>world)",
505 vec!["sds_match".to_string()],
506 Ok(()),
507 ),
508 (
509 "hello (?<capture_grou>world)",
510 vec!["capture_group".to_string()],
511 Err(
512 RegexPatternCaptureGroupsValidationError::CaptureGroupNotPresent(
513 "capture_group".to_string(),
514 ),
515 ),
516 ),
517 (
518 "hello (?<sds_match>d*)",
519 vec!["sds_match".to_string()],
520 Err(RegexPatternCaptureGroupsValidationError::CaptureGroupMatchesEmptyString),
521 ),
522 (
523 "hello (?<sds_match>world)",
524 vec!["sds_match".to_string(), "sds_match2".to_string()],
525 Err(RegexPatternCaptureGroupsValidationError::TooManyCaptureGroups(2)),
526 ),
527 ];
528 for (pattern, capture_groups, expected_result) in test_cases {
529 let rule_config =
530 RegexRuleConfig::new(pattern).with_pattern_capture_groups(capture_groups);
531 assert_eq!(
532 is_pattern_capture_groups_valid(
533 &rule_config.pattern,
534 &rule_config.pattern_capture_groups,
535 &get_memoized_regex(pattern, validate_and_create_regex)
536 .unwrap()
537 .group_info()
538 ),
539 expected_result
540 );
541 }
542 }
543}