tokio_prompt_orchestrator/
prompt_safety.rs1use std::collections::HashSet;
4
5#[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#[derive(Debug, Clone)]
20pub struct SafetyScore {
21 pub category: SafetyCategory,
22 pub confidence: f64,
23 pub triggered_phrases: Vec<String>,
24}
25
26#[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#[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#[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
71pub 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 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 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 {
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 {
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 {
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 pub fn detect_pii(text: &str) -> Vec<PiiMatch> {
171 let mut matches = Vec::new();
172
173 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 {
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 {
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 {
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 {
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 pub fn redact_pii(text: &str, matches: &[PiiMatch]) -> String {
278 if matches.is_empty() {
279 return text.to_string();
280 }
281 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 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 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 pub fn overall_risk(scores: &[SafetyScore]) -> f64 {
317 scores.iter().map(|s| s.confidence).fold(0.0_f64, f64::max)
318 }
319
320 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 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 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 for i in 0..bytes.len() {
379 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 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 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
466pub struct SafetyFilter {
468 pub moderator: ContentModerator,
469 pub allow_list: Vec<String>,
470}
471
472impl SafetyFilter {
473 pub fn new(moderator: ContentModerator, allow_list: Vec<String>) -> Self {
475 Self { moderator, allow_list }
476 }
477
478 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}