diff --git a/src/dlp.rs b/src/dlp.rs index 3b105b7..6263f25 100644 --- a/src/dlp.rs +++ b/src/dlp.rs @@ -28,14 +28,7 @@ pub struct ScanResult { } impl DlpScanner { - pub fn new(patterns: &[DlpPattern]) -> Result { - Self::with_response_scanning(patterns, true) - } - - pub fn with_response_scanning( - patterns: &[DlpPattern], - scan_responses: bool, - ) -> Result { + pub fn new(patterns: &[DlpPattern], scan_responses: bool) -> Result { let compiled = patterns .iter() .map(|p| { @@ -57,56 +50,55 @@ impl DlpScanner { }) } - /// Scans the input body for sensitive data. - /// Returns a list of pattern names that matched with action=block (without including the actual sensitive data). - #[cfg(test)] - pub fn scan(&self, body: &[u8]) -> Vec { - trace!(body_len = body.len(), "Scanning body for block patterns"); - let matches: Vec = self - .patterns + /// Returns names of patterns with action=block that match the body. + fn scan_inner(&self, body: &[u8]) -> Vec { + self.patterns .iter() .filter(|p| p.action == DlpAction::Block && p.regex.is_match(body)) .map(|p| { debug!(pattern = %p.name, "Block pattern matched"); p.name.clone() }) - .collect(); - trace!(match_count = matches.len(), "Block scan complete"); - matches + .collect() } - /// Scans body and applies both block detection and redaction. - /// Returns blocked pattern names and redacted body. - pub fn scan_and_redact(&self, body: &[u8]) -> ScanResult { - trace!( - body_len = body.len(), - "Scanning body for block+redact patterns" - ); - let blocked: Vec = self - .patterns - .iter() - .filter(|p| p.action == DlpAction::Block && p.regex.is_match(body)) - .map(|p| { - debug!(pattern = %p.name, "Block pattern matched in request"); - p.name.clone() - }) - .collect(); - + /// Replaces pattern matches in body with `[REDACTED:]`. + /// + /// - When `force` is true, all patterns are applied regardless of [`DlpAction`]. + /// - When `force` is false, only [`DlpAction::Redact`] patterns are applied. + fn redact_inner(&self, body: &[u8], force: bool) -> (Vec, Vec) { let mut redacted = body.to_vec(); - let mut was_redacted = false; + let mut redacted_names = Vec::new(); for p in &self.patterns { - if p.action == DlpAction::Redact && p.regex.is_match(&redacted) { + if !force && p.action != DlpAction::Redact { + continue; + } + if p.regex.is_match(&redacted) { debug!(pattern = %p.name, "Redact pattern matched, masking PII"); let replacement = format!("[REDACTED:{}]", p.name); redacted = p .regex .replace_all(&redacted, replacement.as_bytes()) .to_vec(); - was_redacted = true; + redacted_names.push(p.name.clone()); } } + (redacted, redacted_names) + } + + /// Scans body and applies both block detection and redaction. + /// Returns blocked pattern names and redacted body. + pub fn scan_and_redact(&self, body: &[u8]) -> ScanResult { + trace!( + body_len = body.len(), + "Scanning body for block+redact patterns" + ); + let blocked = self.scan_inner(body); + let (redacted, redacted_names) = self.redact_inner(body, false); + let was_redacted = !redacted_names.is_empty(); + trace!( blocked_count = blocked.len(), was_redacted, "Scan-and-redact complete" @@ -123,23 +115,9 @@ impl DlpScanner { /// Used for response scanning where we want to redact rather than block. pub fn redact_all(&self, body: &[u8]) -> (Vec, Vec) { trace!(body_len = body.len(), "Redacting all patterns from body"); - let mut redacted = body.to_vec(); - let mut redacted_names = Vec::new(); - - for p in &self.patterns { - if p.regex.is_match(&redacted) { - debug!(pattern = %p.name, "Pattern matched in response, redacting"); - let replacement = format!("[REDACTED:{}]", p.name); - redacted = p - .regex - .replace_all(&redacted, replacement.as_bytes()) - .to_vec(); - redacted_names.push(p.name.clone()); - } - } - - trace!(redacted_count = redacted_names.len(), "Redact-all complete"); - (redacted, redacted_names) + let result = self.redact_inner(body, true); + trace!(redacted_count = result.1.len(), "Redact-all complete"); + result } /// Whether response scanning is enabled. @@ -201,53 +179,54 @@ mod tests { #[test] fn test_detect_credit_card() { - let scanner = DlpScanner::new(&default_patterns()).unwrap(); - let matches = scanner.scan(b"My card is 4111 1111 1111 1111 please charge it"); - assert!(matches.contains(&"credit_card".to_string())); + let scanner = DlpScanner::new(&default_patterns(), false).unwrap(); + let result = scanner.scan_and_redact(b"My card is 4111 1111 1111 1111 please charge it"); + assert!(result.blocked.contains(&"credit_card".to_string())); } #[test] fn test_detect_ssn() { - let scanner = DlpScanner::new(&default_patterns()).unwrap(); - let matches = scanner.scan(b"My SSN is 123-45-6789"); - assert!(matches.contains(&"ssn".to_string())); + let scanner = DlpScanner::new(&default_patterns(), false).unwrap(); + let result = scanner.scan_and_redact(b"My SSN is 123-45-6789"); + assert!(result.blocked.contains(&"ssn".to_string())); } #[test] fn test_detect_email() { - let scanner = DlpScanner::new(&default_patterns()).unwrap(); - let matches = scanner.scan(b"Contact me at user@example.com"); - assert!(matches.contains(&"email".to_string())); + let scanner = DlpScanner::new(&default_patterns(), false).unwrap(); + let result = scanner.scan_and_redact(b"Contact me at user@example.com"); + assert!(result.blocked.contains(&"email".to_string())); } #[test] fn test_no_sensitive_data() { - let scanner = DlpScanner::new(&default_patterns()).unwrap(); - let matches = scanner.scan(b"Tell me about the weather today"); - assert!(matches.is_empty()); + let scanner = DlpScanner::new(&default_patterns(), false).unwrap(); + let result = scanner.scan_and_redact(b"Tell me about the weather today"); + assert!(result.blocked.is_empty()); } #[test] fn test_multiple_detections() { - let scanner = DlpScanner::new(&default_patterns()).unwrap(); - let matches = scanner.scan(b"Card: 4111111111111111, SSN: 123-45-6789, email: a@b.com"); - assert!(matches.contains(&"credit_card".to_string())); - assert!(matches.contains(&"ssn".to_string())); - assert!(matches.contains(&"email".to_string())); + let scanner = DlpScanner::new(&default_patterns(), false).unwrap(); + let result = + scanner.scan_and_redact(b"Card: 4111111111111111, SSN: 123-45-6789, email: a@b.com"); + assert!(result.blocked.contains(&"credit_card".to_string())); + assert!(result.blocked.contains(&"ssn".to_string())); + assert!(result.blocked.contains(&"email".to_string())); } #[test] fn test_empty_patterns() { - let scanner = DlpScanner::new(&[]).unwrap(); - let matches = scanner.scan(b"4111111111111111"); - assert!(matches.is_empty()); + let scanner = DlpScanner::new(&[], false).unwrap(); + let result = scanner.scan_and_redact(b"4111111111111111"); + assert!(result.blocked.is_empty()); } // ========== Redaction Tests ========== #[test] fn test_redact_email() { - let scanner = DlpScanner::new(&mixed_patterns()).unwrap(); + let scanner = DlpScanner::new(&mixed_patterns(), false).unwrap(); let result = scanner.scan_and_redact(b"Contact me at user@example.com"); assert!(result.blocked.is_empty()); assert!(result.was_redacted); @@ -257,7 +236,7 @@ mod tests { #[test] fn test_redact_phone() { - let scanner = DlpScanner::new(&mixed_patterns()).unwrap(); + let scanner = DlpScanner::new(&mixed_patterns(), false).unwrap(); let result = scanner.scan_and_redact(b"Call me at 555-123-4567"); assert!(result.blocked.is_empty()); assert!(result.was_redacted); @@ -267,7 +246,7 @@ mod tests { #[test] fn test_block_ssn_and_redact_email() { - let scanner = DlpScanner::new(&mixed_patterns()).unwrap(); + let scanner = DlpScanner::new(&mixed_patterns(), false).unwrap(); let result = scanner.scan_and_redact(b"SSN: 123-45-6789, email: user@example.com"); assert!(result.blocked.contains(&"ssn".to_string())); assert!(result.was_redacted); @@ -276,22 +255,22 @@ mod tests { #[test] fn test_scan_only_returns_block_patterns() { - let scanner = DlpScanner::new(&mixed_patterns()).unwrap(); - // Email is action=redact, so scan() should NOT return it - let matches = scanner.scan(b"Contact me at user@example.com"); - assert!(!matches.contains(&"email".to_string())); + let scanner = DlpScanner::new(&mixed_patterns(), false).unwrap(); + // Email is action=redact, so blocked should NOT contain it + let result = scanner.scan_and_redact(b"Contact me at user@example.com"); + assert!(!result.blocked.contains(&"email".to_string())); } #[test] fn test_scan_returns_block_patterns() { - let scanner = DlpScanner::new(&mixed_patterns()).unwrap(); - let matches = scanner.scan(b"My SSN is 123-45-6789"); - assert!(matches.contains(&"ssn".to_string())); + let scanner = DlpScanner::new(&mixed_patterns(), false).unwrap(); + let result = scanner.scan_and_redact(b"My SSN is 123-45-6789"); + assert!(result.blocked.contains(&"ssn".to_string())); } #[test] fn test_redact_all_replaces_everything() { - let scanner = DlpScanner::new(&mixed_patterns()).unwrap(); + let scanner = DlpScanner::new(&mixed_patterns(), false).unwrap(); let (redacted, names) = scanner.redact_all(b"SSN: 123-45-6789, email: user@example.com, phone: 555-123-4567"); assert!(names.contains(&"ssn".to_string())); @@ -307,7 +286,7 @@ mod tests { #[test] fn test_redact_all_clean_bytes() { - let scanner = DlpScanner::new(&mixed_patterns()).unwrap(); + let scanner = DlpScanner::new(&mixed_patterns(), false).unwrap(); let (redacted, names) = scanner.redact_all(b"Hello, how are you?"); assert!(names.is_empty()); assert_eq!(redacted, b"Hello, how are you?"); @@ -315,7 +294,7 @@ mod tests { #[test] fn test_no_redaction_when_clean() { - let scanner = DlpScanner::new(&mixed_patterns()).unwrap(); + let scanner = DlpScanner::new(&mixed_patterns(), false).unwrap(); let result = scanner.scan_and_redact(b"Hello world"); assert!(result.blocked.is_empty()); assert!(!result.was_redacted); @@ -324,9 +303,9 @@ mod tests { #[test] fn test_scan_responses_flag() { - let scanner = DlpScanner::with_response_scanning(&[], true).unwrap(); + let scanner = DlpScanner::new(&[], true).unwrap(); assert!(scanner.scan_responses()); - let scanner = DlpScanner::with_response_scanning(&[], false).unwrap(); + let scanner = DlpScanner::new(&[], false).unwrap(); assert!(!scanner.scan_responses()); } @@ -337,37 +316,12 @@ mod tests { regex: r"\b[A-Za-z0-9._%+-]+@[A-Za-z0-9.-]+\.[A-Za-z]{2,}\b".to_string(), action: DlpAction::Redact, }]; - let scanner = DlpScanner::new(&patterns).unwrap(); + let scanner = DlpScanner::new(&patterns, false).unwrap(); let result = scanner.scan_and_redact(b"a@b.com and c@d.com"); assert!(result.was_redacted); assert_eq!(result.redacted, b"[REDACTED:email] and [REDACTED:email]"); } - #[test] - fn test_detect_phone_various_formats() { - let patterns = vec![DlpPattern { - name: "phone_number".to_string(), - regex: r"\b(?:\+?1[-.\s]?)?(?:\(?\d{3}\)?[-.\s]?)?\d{3}[-.\s]?\d{4}\b".to_string(), - action: DlpAction::Redact, - }]; - let scanner = DlpScanner::new(&patterns).unwrap(); - - // Standard US format - let (redacted, names) = scanner.redact_all(b"Call 555-123-4567"); - assert!(names.contains(&"phone_number".to_string())); - assert!(!subslice(&redacted, b"555-123-4567")); - - // With parentheses - let (redacted, names) = scanner.redact_all(b"Call (555) 123-4567"); - assert!(names.contains(&"phone_number".to_string())); - assert!(!subslice(&redacted, b"(555) 123-4567")); - - // With +1 prefix - let (redacted, names) = scanner.redact_all(b"Call +1-555-123-4567"); - assert!(names.contains(&"phone_number".to_string())); - assert!(!subslice(&redacted, b"+1-555-123-4567")); - } - #[test] fn test_non_utf8_input() { let patterns = vec![DlpPattern { @@ -375,7 +329,7 @@ mod tests { regex: r"\b(?:\+?1[-.\s]?)?(?:\(?\d{3}\)?[-.\s]?)?\d{3}[-.\s]?\d{4}\b".to_string(), action: DlpAction::Redact, }]; - let scanner = DlpScanner::new(&patterns).unwrap(); + let scanner = DlpScanner::new(&patterns, false).unwrap(); let input = b"Call \xFF\xFE\xFD 555-123-4567"; std::str::from_utf8(input.as_slice()).unwrap_err(); // Confirm it's not valid UTF-8 diff --git a/src/lib.rs b/src/lib.rs index 8d935b6..b1b930b 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -47,7 +47,7 @@ impl AppState { Self { key_manager: Arc::new(KeyManager::new(config.key_map())), dlp_scanner: Arc::new( - DlpScanner::with_response_scanning(&config.dlp.patterns, config.dlp.scan_responses) + DlpScanner::new(&config.dlp.patterns, config.dlp.scan_responses) .expect("Failed to compile DLP patterns"), ), proxy_client: Arc::new(ProxyClient::with_upstream_urls( diff --git a/tests/integration.rs b/tests/integration.rs index 984e0e4..b6a1447 100644 --- a/tests/integration.rs +++ b/tests/integration.rs @@ -50,7 +50,7 @@ fn make_app(upstream_url: &str) -> axum::Router { let state = AppState { key_manager: Arc::new(KeyManager::new(key_map)), - dlp_scanner: Arc::new(DlpScanner::new(&patterns).unwrap()), + dlp_scanner: Arc::new(DlpScanner::new(&patterns, false).unwrap()), proxy_client: Arc::new(ProxyClient::with_upstream_urls( upstream_urls, "2023-06-01".to_string(), @@ -77,7 +77,7 @@ fn make_app_with_anthropic(upstream_url: &str) -> axum::Router { let state = AppState { key_manager: Arc::new(KeyManager::new(key_map)), - dlp_scanner: Arc::new(DlpScanner::new(&[]).unwrap()), + dlp_scanner: Arc::new(DlpScanner::new(&[], false).unwrap()), proxy_client: Arc::new(ProxyClient::with_upstream_urls( upstream_urls, "2023-06-01".to_string(), @@ -701,7 +701,7 @@ async fn test_proxy_error_on_unreachable_upstream() { .into_iter() .collect(), )), - dlp_scanner: Arc::new(DlpScanner::new(&[]).unwrap()), + dlp_scanner: Arc::new(DlpScanner::new(&[], false).unwrap()), proxy_client: Arc::new(ProxyClient::with_upstream_urls( { let mut urls = BTreeMap::new(); @@ -843,7 +843,7 @@ async fn test_anthropic_dlp_blocks_sensitive_data() { let state = AppState { key_manager: Arc::new(KeyManager::new(key_map)), - dlp_scanner: Arc::new(DlpScanner::new(&patterns).unwrap()), + dlp_scanner: Arc::new(DlpScanner::new(&patterns, false).unwrap()), proxy_client: Arc::new(ProxyClient::with_upstream_urls( upstream_urls, "2023-06-01".to_string(), @@ -1009,7 +1009,7 @@ fn make_app_with_redact(upstream_url: &str) -> axum::Router { let state = AppState { key_manager: Arc::new(KeyManager::new(key_map)), - dlp_scanner: Arc::new(DlpScanner::with_response_scanning(&patterns, true).unwrap()), + dlp_scanner: Arc::new(DlpScanner::new(&patterns, true).unwrap()), proxy_client: Arc::new(ProxyClient::with_upstream_urls( upstream_urls, "2023-06-01".to_string(), @@ -1206,7 +1206,7 @@ async fn test_response_dlp_disabled() { upstream_urls.insert(Provider::Anthropic, mock_server.uri()); let state = AppState { key_manager: Arc::new(KeyManager::new(key_map)), - dlp_scanner: Arc::new(DlpScanner::with_response_scanning(&patterns, false).unwrap()), + dlp_scanner: Arc::new(DlpScanner::new(&patterns, false).unwrap()), proxy_client: Arc::new(ProxyClient::with_upstream_urls( upstream_urls, "2023-06-01".to_string(),