Skip to main content

tokio_prompt_orchestrator/
prompt_safety.rs

1//! Content moderation, toxicity detection, and PII detection.
2
3use std::collections::HashSet;
4
5/// Categories of unsafe content.
6#[derive(Debug, Clone, PartialEq, Eq, Hash)]
7pub enum SafetyCategory {
8    Hate,
9    Violence,
10    SelfHarm,
11    SexualContent,
12    Harassment,
13    Spam,
14    PII,
15    Safe,
16}
17
18/// A score for a single safety category.
19#[derive(Debug, Clone)]
20pub struct SafetyScore {
21    pub category: SafetyCategory,
22    pub confidence: f64,
23    pub triggered_phrases: Vec<String>,
24}
25
26/// A detected PII match within the input text.
27#[derive(Debug, Clone)]
28pub struct PiiMatch {
29    pub pii_type: String,
30    pub value: String,
31    pub start: usize,
32    pub end: usize,
33}
34
35/// Full result of a safety analysis.
36#[derive(Debug, Clone)]
37pub struct SafetyResult {
38    pub input: String,
39    pub scores: Vec<SafetyScore>,
40    pub overall_risk: f64,
41    pub blocked: bool,
42    pub pii_detected: Vec<PiiMatch>,
43}
44
45/// Configuration for the content moderator.
46#[derive(Debug, Clone)]
47pub struct SafetyConfig {
48    pub block_threshold: f64,
49    pub categories_enabled: Vec<SafetyCategory>,
50    pub redact_pii: bool,
51}
52
53impl Default for SafetyConfig {
54    fn default() -> Self {
55        Self {
56            block_threshold: 0.5,
57            categories_enabled: vec![
58                SafetyCategory::Hate,
59                SafetyCategory::Violence,
60                SafetyCategory::SelfHarm,
61                SafetyCategory::SexualContent,
62                SafetyCategory::Harassment,
63                SafetyCategory::Spam,
64                SafetyCategory::PII,
65            ],
66            redact_pii: true,
67        }
68    }
69}
70
71/// Content moderation engine.
72pub struct ContentModerator {
73    pub config: SafetyConfig,
74    pub hate_phrases: Vec<String>,
75    pub violence_phrases: Vec<String>,
76    pub spam_indicators: Vec<String>,
77}
78
79impl ContentModerator {
80    /// Create a new `ContentModerator` with pre-populated phrase lists.
81    pub fn new(config: SafetyConfig) -> Self {
82        Self {
83            config,
84            hate_phrases: vec![
85                "kill all".to_string(),
86                "hate you".to_string(),
87                "die".to_string(),
88                "worthless".to_string(),
89            ],
90            violence_phrases: vec![
91                "bomb".to_string(),
92                "explode".to_string(),
93                "shoot".to_string(),
94                "attack".to_string(),
95            ],
96            spam_indicators: vec![
97                "click here".to_string(),
98                "buy now".to_string(),
99                "limited offer".to_string(),
100                "act now".to_string(),
101                "free money".to_string(),
102            ],
103        }
104    }
105
106    /// Scan text for toxic content per category. Confidence = hit_count / word_count.
107    pub fn detect_toxicity(&self, text: &str) -> Vec<SafetyScore> {
108        let lower = text.to_lowercase();
109        let word_count = text.split_whitespace().count().max(1);
110        let mut scores = Vec::new();
111
112        // Hate
113        {
114            let mut triggered = Vec::new();
115            for phrase in &self.hate_phrases {
116                if lower.contains(phrase.as_str()) {
117                    triggered.push(phrase.clone());
118                }
119            }
120            if !triggered.is_empty() {
121                let confidence = (triggered.len() as f64 / word_count as f64).min(1.0);
122                scores.push(SafetyScore {
123                    category: SafetyCategory::Hate,
124                    confidence,
125                    triggered_phrases: triggered,
126                });
127            }
128        }
129
130        // Violence
131        {
132            let mut triggered = Vec::new();
133            for phrase in &self.violence_phrases {
134                if lower.contains(phrase.as_str()) {
135                    triggered.push(phrase.clone());
136                }
137            }
138            if !triggered.is_empty() {
139                let confidence = (triggered.len() as f64 / word_count as f64).min(1.0);
140                scores.push(SafetyScore {
141                    category: SafetyCategory::Violence,
142                    confidence,
143                    triggered_phrases: triggered,
144                });
145            }
146        }
147
148        // Spam
149        {
150            let mut triggered = Vec::new();
151            for phrase in &self.spam_indicators {
152                if lower.contains(phrase.as_str()) {
153                    triggered.push(phrase.clone());
154                }
155            }
156            if !triggered.is_empty() {
157                let confidence = (triggered.len() as f64 / word_count as f64).min(1.0);
158                scores.push(SafetyScore {
159                    category: SafetyCategory::Spam,
160                    confidence,
161                    triggered_phrases: triggered,
162                });
163            }
164        }
165
166        scores
167    }
168
169    /// Detect PII in the given text using pattern matching.
170    pub fn detect_pii(text: &str) -> Vec<PiiMatch> {
171        let mut matches = Vec::new();
172
173        // Email: contains '@' and '.'
174        for (start, word) in Self::word_spans(text) {
175            let w = word.trim_matches(|c: char| !c.is_alphanumeric() && c != '@' && c != '.' && c != '-' && c != '_');
176            if w.contains('@') && w.contains('.') && w.len() > 3 {
177                let end = start + word.len();
178                matches.push(PiiMatch {
179                    pii_type: "email".to_string(),
180                    value: w.to_string(),
181                    start,
182                    end,
183                });
184            }
185        }
186
187        // Phone: 10+ consecutive digit sequence (may include spaces/dashes)
188        {
189            let bytes = text.as_bytes();
190            let mut i = 0;
191            while i < bytes.len() {
192                if bytes[i].is_ascii_digit() {
193                    let start = i;
194                    let mut count = 0;
195                    let mut end = i;
196                    while end < bytes.len() && (bytes[end].is_ascii_digit() || bytes[end] == b'-' || bytes[end] == b' ') {
197                        if bytes[end].is_ascii_digit() {
198                            count += 1;
199                        }
200                        end += 1;
201                    }
202                    if count >= 10 {
203                        let value = text[start..end].trim().to_string();
204                        matches.push(PiiMatch {
205                            pii_type: "phone".to_string(),
206                            value,
207                            start,
208                            end,
209                        });
210                        i = end;
211                        continue;
212                    }
213                }
214                i += 1;
215            }
216        }
217
218        // SSN: XXX-XX-XXXX pattern
219        {
220            let chars: Vec<char> = text.chars().collect();
221            let s: String = chars.iter().collect();
222            let mut search_start = 0;
223            while search_start < s.len() {
224                if let Some(pos) = Self::find_ssn(&s[search_start..]) {
225                    let abs_pos = search_start + pos.0;
226                    matches.push(PiiMatch {
227                        pii_type: "ssn".to_string(),
228                        value: pos.1.clone(),
229                        start: abs_pos,
230                        end: abs_pos + pos.1.len(),
231                    });
232                    search_start = abs_pos + pos.1.len();
233                } else {
234                    break;
235                }
236            }
237        }
238
239        // Credit card: 16-digit groups (4x4 separated by spaces or dashes)
240        {
241            let mut search = text;
242            let mut offset = 0;
243            while let Some((pos, val)) = Self::find_credit_card(search) {
244                matches.push(PiiMatch {
245                    pii_type: "credit_card".to_string(),
246                    value: val.clone(),
247                    start: offset + pos,
248                    end: offset + pos + val.len(),
249                });
250                let advance = pos + val.len();
251                offset += advance;
252                search = &search[advance..];
253            }
254        }
255
256        // IP address: X.X.X.X pattern
257        {
258            let mut search = text;
259            let mut offset = 0;
260            while let Some((pos, val)) = Self::find_ip(search) {
261                matches.push(PiiMatch {
262                    pii_type: "ip_address".to_string(),
263                    value: val.clone(),
264                    start: offset + pos,
265                    end: offset + pos + val.len(),
266                });
267                let advance = pos + val.len();
268                offset += advance;
269                search = &search[advance..];
270            }
271        }
272
273        matches
274    }
275
276    /// Replace matched PII spans with "`REDACTED`".
277    pub fn redact_pii(text: &str, matches: &[PiiMatch]) -> String {
278        if matches.is_empty() {
279            return text.to_string();
280        }
281        // Sort by start position descending so we can replace without shifting indices
282        let mut sorted: Vec<&PiiMatch> = matches.iter().collect();
283        sorted.sort_by_key(|m| std::cmp::Reverse(m.start));
284
285        let mut result = text.to_string();
286        // Deduplicate overlapping ranges
287        let mut seen: HashSet<(usize, usize)> = HashSet::new();
288        for m in sorted {
289            if seen.contains(&(m.start, m.end)) {
290                continue;
291            }
292            if m.end <= result.len() {
293                seen.insert((m.start, m.end));
294                result.replace_range(m.start..m.end, "[REDACTED]");
295            }
296        }
297        result
298    }
299
300    /// Run all detectors and produce a `SafetyResult`.
301    pub fn analyze(&self, text: &str) -> SafetyResult {
302        let scores = self.detect_toxicity(text);
303        let pii_detected = Self::detect_pii(text);
304        let overall_risk = Self::overall_risk(&scores);
305        let blocked = overall_risk >= self.config.block_threshold;
306        SafetyResult {
307            input: text.to_string(),
308            scores,
309            overall_risk,
310            blocked,
311            pii_detected,
312        }
313    }
314
315    /// Compute overall risk as the maximum confidence across all triggered categories.
316    pub fn overall_risk(scores: &[SafetyScore]) -> f64 {
317        scores.iter().map(|s| s.confidence).fold(0.0_f64, f64::max)
318    }
319
320    // --- private helpers ---
321
322    fn word_spans(text: &str) -> Vec<(usize, &str)> {
323        let mut spans = Vec::new();
324        let mut start = 0;
325        let mut in_word = false;
326        for (i, c) in text.char_indices() {
327            if c.is_whitespace() {
328                if in_word {
329                    spans.push((start, &text[start..i]));
330                    in_word = false;
331                }
332                start = i + c.len_utf8();
333            } else if !in_word {
334                start = i;
335                in_word = true;
336            }
337        }
338        if in_word {
339            spans.push((start, &text[start..]));
340        }
341        spans
342    }
343
344    fn find_ssn(text: &str) -> Option<(usize, String)> {
345        let bytes = text.as_bytes();
346        for i in 0..bytes.len() {
347            if i + 11 <= bytes.len() {
348                let slice = &bytes[i..i + 11];
349                // Pattern: DDD-DD-DDDD
350                let ok = slice[0].is_ascii_digit()
351                    && slice[1].is_ascii_digit()
352                    && slice[2].is_ascii_digit()
353                    && slice[3] == b'-'
354                    && slice[4].is_ascii_digit()
355                    && slice[5].is_ascii_digit()
356                    && slice[6] == b'-'
357                    && slice[7].is_ascii_digit()
358                    && slice[8].is_ascii_digit()
359                    && slice[9].is_ascii_digit()
360                    && slice[10].is_ascii_digit();
361                if ok {
362                    // Ensure not part of a longer digit sequence
363                    let before_ok = i == 0 || !bytes[i - 1].is_ascii_digit();
364                    let after_ok = i + 11 >= bytes.len() || !bytes[i + 11].is_ascii_digit();
365                    if before_ok && after_ok {
366                        let val = String::from_utf8_lossy(&bytes[i..i + 11]).into_owned();
367                        return Some((i, val));
368                    }
369                }
370            }
371        }
372        None
373    }
374
375    fn find_credit_card(text: &str) -> Option<(usize, String)> {
376        let bytes = text.as_bytes();
377        // Look for 16 consecutive digits, possibly split as 4-4-4-4 with space/dash
378        for i in 0..bytes.len() {
379            // Try compact 16 digits
380            if i + 16 <= bytes.len() {
381                let slice = &bytes[i..i + 16];
382                if slice.iter().all(|b| b.is_ascii_digit()) {
383                    let before_ok = i == 0 || !bytes[i - 1].is_ascii_digit();
384                    let after_ok = i + 16 >= bytes.len() || !bytes[i + 16].is_ascii_digit();
385                    if before_ok && after_ok {
386                        let val = String::from_utf8_lossy(slice).into_owned();
387                        return Some((i, val));
388                    }
389                }
390            }
391            // Try 4-4-4-4 with separator
392            if i + 19 <= bytes.len() {
393                let sep = bytes[i + 4];
394                if (sep == b' ' || sep == b'-')
395                    && bytes[i..i + 4].iter().all(|b| b.is_ascii_digit())
396                    && bytes[i + 5..i + 9].iter().all(|b| b.is_ascii_digit())
397                    && bytes[i + 9] == sep
398                    && bytes[i + 10..i + 14].iter().all(|b| b.is_ascii_digit())
399                    && bytes[i + 14] == sep
400                    && bytes[i + 15..i + 19].iter().all(|b| b.is_ascii_digit())
401                {
402                    let before_ok = i == 0 || !bytes[i - 1].is_ascii_digit();
403                    let after_ok = i + 19 >= bytes.len() || !bytes[i + 19].is_ascii_digit();
404                    if before_ok && after_ok {
405                        let val = String::from_utf8_lossy(&bytes[i..i + 19]).into_owned();
406                        return Some((i, val));
407                    }
408                }
409            }
410        }
411        None
412    }
413
414    fn find_ip(text: &str) -> Option<(usize, String)> {
415        let bytes = text.as_bytes();
416        let mut i = 0;
417        while i < bytes.len() {
418            if bytes[i].is_ascii_digit() {
419                // Try to parse X.X.X.X
420                let start = i;
421                let mut parts: Vec<u8> = Vec::new();
422                let mut cur_num: Option<u32> = None;
423                let mut j = i;
424                let mut dot_count = 0;
425                while j < bytes.len() {
426                    if bytes[j].is_ascii_digit() {
427                        let d = (bytes[j] - b'0') as u32;
428                        cur_num = Some(cur_num.unwrap_or(0) * 10 + d);
429                        if cur_num.unwrap_or(0) > 255 {
430                            break;
431                        }
432                        j += 1;
433                    } else if bytes[j] == b'.' && dot_count < 3 {
434                        if let Some(n) = cur_num {
435                            parts.push(n as u8);
436                            cur_num = None;
437                            dot_count += 1;
438                            j += 1;
439                        } else {
440                            break;
441                        }
442                    } else {
443                        break;
444                    }
445                }
446                if let Some(n) = cur_num {
447                    parts.push(n as u8);
448                }
449                if parts.len() == 4 {
450                    let before_ok = start == 0 || !bytes[start - 1].is_ascii_digit() && bytes[start - 1] != b'.';
451                    let after_ok = j >= bytes.len() || !bytes[j].is_ascii_digit() && bytes[j] != b'.';
452                    if before_ok && after_ok {
453                        let val = text[start..j].to_string();
454                        return Some((start, val));
455                    }
456                }
457                i = j.max(start + 1);
458            } else {
459                i += 1;
460            }
461        }
462        None
463    }
464}
465
466/// A safety filter that wraps a `ContentModerator` and supports an allow-list.
467pub struct SafetyFilter {
468    pub moderator: ContentModerator,
469    pub allow_list: Vec<String>,
470}
471
472impl SafetyFilter {
473    /// Create a new `SafetyFilter`.
474    pub fn new(moderator: ContentModerator, allow_list: Vec<String>) -> Self {
475        Self { moderator, allow_list }
476    }
477
478    /// Redact PII from `text` and return the processed text along with the full `SafetyResult`.
479    pub fn check_and_redact(&self, text: &str) -> (String, SafetyResult) {
480        let result = self.moderator.analyze(text);
481        let processed = if self.moderator.config.redact_pii {
482            ContentModerator::redact_pii(text, &result.pii_detected)
483        } else {
484            text.to_string()
485        };
486        (processed, result)
487    }
488}
489
490#[cfg(test)]
491mod tests {
492    use super::*;
493
494    fn default_moderator() -> ContentModerator {
495        ContentModerator::new(SafetyConfig::default())
496    }
497
498    #[test]
499    fn test_hate_phrase_detection() {
500        let m = default_moderator();
501        let scores = m.detect_toxicity("I hate you so much, you are worthless.");
502        let hate = scores.iter().find(|s| s.category == SafetyCategory::Hate);
503        assert!(hate.is_some(), "expected hate category to be triggered");
504        assert!(hate.unwrap().confidence > 0.0);
505    }
506
507    #[test]
508    fn test_spam_detection() {
509        let m = default_moderator();
510        let scores = m.detect_toxicity("Click here and buy now for free money!");
511        let spam = scores.iter().find(|s| s.category == SafetyCategory::Spam);
512        assert!(spam.is_some(), "expected spam category to be triggered");
513    }
514
515    #[test]
516    fn test_email_pii_found() {
517        let matches = ContentModerator::detect_pii("Contact us at user@example.com for details.");
518        let email = matches.iter().find(|m| m.pii_type == "email");
519        assert!(email.is_some(), "expected email PII to be detected");
520        assert!(email.unwrap().value.contains('@'));
521    }
522
523    #[test]
524    fn test_ssn_pattern_matched() {
525        let matches = ContentModerator::detect_pii("SSN is 123-45-6789 on file.");
526        let ssn = matches.iter().find(|m| m.pii_type == "ssn");
527        assert!(ssn.is_some(), "expected SSN PII to be detected");
528        assert_eq!(ssn.unwrap().value, "123-45-6789");
529    }
530
531    #[test]
532    fn test_pii_redaction() {
533        let text = "Email user@example.com or call 1234567890 today.";
534        let matches = ContentModerator::detect_pii(text);
535        let redacted = ContentModerator::redact_pii(text, &matches);
536        assert!(!redacted.contains("user@example.com"), "email should be redacted");
537    }
538
539    #[test]
540    fn test_safe_text_passes() {
541        let m = default_moderator();
542        let result = m.analyze("The weather is nice today. I enjoyed a walk in the park.");
543        assert!(!result.blocked, "safe text should not be blocked");
544        assert_eq!(result.overall_risk, 0.0);
545    }
546}