Skip to main content

zeph_sanitizer/
memory_validation.rs

1// SPDX-FileCopyrightText: 2026 Andrei G <bug-ops>
2// SPDX-License-Identifier: MIT OR Apache-2.0
3
4//! Memory write validation: structural checks before content reaches the memory store
5//! or the graph extractor.
6//!
7//! Configured under `[security.memory_validation]` in the agent config file.
8//! Enabled by default — guards against oversized writes, injection markers, and PII
9//! leaking into entity names.
10//!
11//! [`MemoryWriteValidator`] covers two distinct write paths:
12//!
13//! 1. **`memory_save` tool** — validates raw text before `SQLite` + Qdrant write.
14//!    Checks size limit and forbidden content patterns.
15//! 2. **Graph extraction** — validates [`ExtractionResult`]
16//!    after `GraphExtractor::extract()` returns. Checks entity count, edge count,
17//!    entity name length, fact text length, and PII in entity names.
18
19use thiserror::Error;
20use zeph_memory::graph::extractor::ExtractionResult;
21
22pub use zeph_config::MemoryWriteValidationConfig;
23
24use crate::pii::{EMAIL_RE, SSN_RE};
25
26// ---------------------------------------------------------------------------
27// Error
28// ---------------------------------------------------------------------------
29
30/// Validation failure reported by [`MemoryWriteValidator`].
31///
32/// Returned by [`validate_memory_save`](MemoryWriteValidator::validate_memory_save) and
33/// [`validate_graph_extraction`](MemoryWriteValidator::validate_graph_extraction). Callers
34/// should log the error and skip the write rather than panicking.
35#[derive(Debug, Error)]
36#[non_exhaustive]
37pub enum MemoryValidationError {
38    /// The content exceeds the configured `max_content_bytes` limit.
39    #[error("content too large: {size} bytes exceeds max {max}")]
40    ContentTooLarge { size: usize, max: usize },
41
42    /// An extracted entity name is shorter than `min_entity_name_bytes`.
43    #[error("entity name too short: '{name}' is below min {min} bytes")]
44    EntityNameTooShort { name: String, min: usize },
45
46    /// An extracted entity name exceeds `max_entity_name_bytes`.
47    #[error("entity name too long: '{name}' exceeds max {max} bytes")]
48    EntityNameTooLong { name: String, max: usize },
49
50    /// An extracted edge fact exceeds `max_fact_bytes`.
51    #[error("fact text too long: exceeds max {max} bytes")]
52    FactTooLong { max: usize },
53
54    /// The extraction produced more entities than `max_entities_per_extraction`.
55    #[error("too many entities: {count} exceeds max {max}")]
56    TooManyEntities { count: usize, max: usize },
57
58    /// The extraction produced more edges than `max_edges_per_extraction`.
59    #[error("too many edges: {count} exceeds max {max}")]
60    TooManyEdges { count: usize, max: usize },
61
62    /// The content matched one of the configured `forbidden_content_patterns`.
63    #[error("forbidden pattern detected: {pattern}")]
64    ForbiddenPattern { pattern: String },
65
66    /// An entity name contains a PII pattern (email or SSN).
67    #[error("PII detected in entity name: '{entity}'")]
68    SuspiciousPiiInEntityName { entity: String },
69}
70
71// ---------------------------------------------------------------------------
72// Validator
73// ---------------------------------------------------------------------------
74
75/// Validates content before it is written to the memory store or graph extractor.
76///
77/// Construct once from [`MemoryWriteValidationConfig`] and store on the agent.
78/// Cheap to clone.
79///
80/// # Examples
81///
82/// ```rust
83/// use zeph_sanitizer::memory_validation::MemoryWriteValidator;
84/// use zeph_config::MemoryWriteValidationConfig;
85///
86/// let validator = MemoryWriteValidator::new(MemoryWriteValidationConfig::default());
87/// assert!(validator.is_enabled());
88///
89/// // Small content passes.
90/// assert!(validator.validate_memory_save("hello world").is_ok());
91///
92/// // Content exceeding the limit is rejected.
93/// let huge = "x".repeat(10_000);
94/// assert!(validator.validate_memory_save(&huge).is_err());
95/// ```
96#[derive(Debug, Clone)]
97pub struct MemoryWriteValidator {
98    config: MemoryWriteValidationConfig,
99}
100
101impl MemoryWriteValidator {
102    /// Create a validator from the given configuration.
103    ///
104    /// # Examples
105    ///
106    /// ```rust
107    /// use zeph_sanitizer::memory_validation::MemoryWriteValidator;
108    /// use zeph_config::MemoryWriteValidationConfig;
109    ///
110    /// let validator = MemoryWriteValidator::new(MemoryWriteValidationConfig::default());
111    /// assert!(validator.is_enabled());
112    /// ```
113    #[must_use]
114    pub fn new(config: MemoryWriteValidationConfig) -> Self {
115        Self { config }
116    }
117
118    /// Validate content before it is written via the `memory_save` tool.
119    ///
120    /// # Errors
121    ///
122    /// Returns [`MemoryValidationError`] if any validation check fails.
123    pub fn validate_memory_save(&self, content: &str) -> Result<(), MemoryValidationError> {
124        if !self.config.enabled {
125            return Ok(());
126        }
127
128        let size = content.len();
129        if size > self.config.max_content_bytes {
130            return Err(MemoryValidationError::ContentTooLarge {
131                size,
132                max: self.config.max_content_bytes,
133            });
134        }
135
136        for pattern in &self.config.forbidden_content_patterns {
137            if content.contains(pattern.as_str()) {
138                return Err(MemoryValidationError::ForbiddenPattern {
139                    pattern: pattern.clone(),
140                });
141            }
142        }
143
144        Ok(())
145    }
146
147    /// Validate a graph extraction result before entities and edges are upserted.
148    ///
149    /// Called inside the spawned extraction task, after `GraphExtractor::extract()` returns.
150    ///
151    /// # Errors
152    ///
153    /// Returns [`MemoryValidationError`] if any validation check fails.
154    pub fn validate_graph_extraction(
155        &self,
156        result: &ExtractionResult,
157    ) -> Result<(), MemoryValidationError> {
158        if !self.config.enabled {
159            return Ok(());
160        }
161
162        let entity_count = result.entities.len();
163        if entity_count > self.config.max_entities_per_extraction {
164            return Err(MemoryValidationError::TooManyEntities {
165                count: entity_count,
166                max: self.config.max_entities_per_extraction,
167            });
168        }
169
170        let edge_count = result.edges.len();
171        if edge_count > self.config.max_edges_per_extraction {
172            return Err(MemoryValidationError::TooManyEdges {
173                count: edge_count,
174                max: self.config.max_edges_per_extraction,
175            });
176        }
177
178        for entity in &result.entities {
179            // Trim before length checks: both min and max apply to the trimmed form
180            // to avoid rejecting names with leading/trailing whitespace.
181            let name_len = entity.name.trim().len();
182            if name_len < self.config.min_entity_name_bytes {
183                return Err(MemoryValidationError::EntityNameTooShort {
184                    name: entity.name.clone(),
185                    min: self.config.min_entity_name_bytes,
186                });
187            }
188            if name_len > self.config.max_entity_name_bytes {
189                return Err(MemoryValidationError::EntityNameTooLong {
190                    name: entity.name.clone(),
191                    max: self.config.max_entity_name_bytes,
192                });
193            }
194            // Guard against PII leaking into entity names (email and SSN).
195            if EMAIL_RE.is_match(&entity.name) || SSN_RE.is_match(&entity.name) {
196                return Err(MemoryValidationError::SuspiciousPiiInEntityName {
197                    entity: entity.name.clone(),
198                });
199            }
200        }
201
202        for edge in &result.edges {
203            let fact_len = edge.fact.len();
204            if fact_len > self.config.max_fact_bytes {
205                return Err(MemoryValidationError::FactTooLong {
206                    max: self.config.max_fact_bytes,
207                });
208            }
209        }
210
211        for pattern in &self.config.forbidden_content_patterns {
212            let p = pattern.as_str();
213            for entity in &result.entities {
214                if entity.name.contains(p) {
215                    return Err(MemoryValidationError::ForbiddenPattern {
216                        pattern: pattern.clone(),
217                    });
218                }
219            }
220            for edge in &result.edges {
221                if edge.fact.contains(p) {
222                    return Err(MemoryValidationError::ForbiddenPattern {
223                        pattern: pattern.clone(),
224                    });
225                }
226            }
227        }
228
229        Ok(())
230    }
231
232    /// Returns `true` when validation is enabled.
233    ///
234    /// When `false`, both [`validate_memory_save`](Self::validate_memory_save) and
235    /// [`validate_graph_extraction`](Self::validate_graph_extraction) always return `Ok(())`.
236    #[must_use]
237    pub fn is_enabled(&self) -> bool {
238        self.config.enabled
239    }
240}
241
242// ---------------------------------------------------------------------------
243// Tests
244// ---------------------------------------------------------------------------
245
246#[cfg(test)]
247mod tests {
248    use std::assert_matches;
249    use zeph_memory::graph::extractor::{ExtractedEdge, ExtractedEntity};
250
251    use super::*;
252
253    fn validator() -> MemoryWriteValidator {
254        MemoryWriteValidator::new(MemoryWriteValidationConfig::default())
255    }
256
257    fn validator_disabled() -> MemoryWriteValidator {
258        MemoryWriteValidator::new(MemoryWriteValidationConfig {
259            enabled: false,
260            ..MemoryWriteValidationConfig::default()
261        })
262    }
263
264    fn entity(name: &str) -> ExtractedEntity {
265        ExtractedEntity {
266            name: name.to_owned(),
267            entity_type: "person".to_owned(),
268            summary: None,
269        }
270    }
271
272    fn edge(fact: &str) -> ExtractedEdge {
273        ExtractedEdge {
274            source: "A".to_owned(),
275            target: "B".to_owned(),
276            relation: "knows".to_owned(),
277            fact: fact.to_owned(),
278            temporal_hint: None,
279            edge_type: "semantic".to_owned(),
280            confidence: None,
281        }
282    }
283
284    fn result_with(entities: Vec<ExtractedEntity>, edges: Vec<ExtractedEdge>) -> ExtractionResult {
285        ExtractionResult { entities, edges }
286    }
287
288    // --- memory_save validation ---
289
290    #[test]
291    fn valid_content_passes() {
292        assert!(validator().validate_memory_save("hello world").is_ok());
293    }
294
295    #[test]
296    fn oversized_content_rejected() {
297        let big = "x".repeat(5000);
298        let err = validator().validate_memory_save(&big).unwrap_err();
299        assert_matches!(err, MemoryValidationError::ContentTooLarge { .. });
300    }
301
302    #[test]
303    fn forbidden_pattern_rejected() {
304        let v = MemoryWriteValidator::new(MemoryWriteValidationConfig {
305            forbidden_content_patterns: vec!["<script".to_owned()],
306            ..MemoryWriteValidationConfig::default()
307        });
308        let err = v
309            .validate_memory_save("text <script>alert(1)</script>")
310            .unwrap_err();
311        assert_matches!(err, MemoryValidationError::ForbiddenPattern { .. });
312    }
313
314    #[test]
315    fn disabled_skips_validation() {
316        let big = "x".repeat(9999);
317        assert!(validator_disabled().validate_memory_save(&big).is_ok());
318    }
319
320    // --- graph extraction validation ---
321
322    #[test]
323    fn valid_extraction_passes() {
324        let r = result_with(vec![entity("Rust"), entity("Alice")], vec![edge("fact")]);
325        assert!(validator().validate_graph_extraction(&r).is_ok());
326    }
327
328    #[test]
329    fn too_many_entities_rejected() {
330        let v = MemoryWriteValidator::new(MemoryWriteValidationConfig {
331            max_entities_per_extraction: 2,
332            ..MemoryWriteValidationConfig::default()
333        });
334        let r = result_with(vec![entity("Abc"), entity("Def"), entity("Ghi")], vec![]);
335        let err = v.validate_graph_extraction(&r).unwrap_err();
336        assert_matches!(err, MemoryValidationError::TooManyEntities { .. });
337    }
338
339    #[test]
340    fn too_many_edges_rejected() {
341        let v = MemoryWriteValidator::new(MemoryWriteValidationConfig {
342            max_edges_per_extraction: 1,
343            ..MemoryWriteValidationConfig::default()
344        });
345        let r = result_with(vec![], vec![edge("a"), edge("b")]);
346        let err = v.validate_graph_extraction(&r).unwrap_err();
347        assert_matches!(err, MemoryValidationError::TooManyEdges { .. });
348    }
349
350    #[test]
351    fn entity_name_too_long_rejected() {
352        let v = MemoryWriteValidator::new(MemoryWriteValidationConfig {
353            max_entity_name_bytes: 5,
354            ..MemoryWriteValidationConfig::default()
355        });
356        let r = result_with(vec![entity("TooLongName")], vec![]);
357        let err = v.validate_graph_extraction(&r).unwrap_err();
358        assert_matches!(err, MemoryValidationError::EntityNameTooLong { .. });
359    }
360
361    #[test]
362    fn fact_too_long_rejected() {
363        let v = MemoryWriteValidator::new(MemoryWriteValidationConfig {
364            max_fact_bytes: 10,
365            ..MemoryWriteValidationConfig::default()
366        });
367        let r = result_with(vec![], vec![edge("this fact is longer than ten chars")]);
368        let err = v.validate_graph_extraction(&r).unwrap_err();
369        assert_matches!(err, MemoryValidationError::FactTooLong { .. });
370    }
371
372    #[test]
373    fn email_in_entity_name_rejected() {
374        let r = result_with(vec![entity("user@example.com")], vec![]);
375        let err = validator().validate_graph_extraction(&r).unwrap_err();
376        assert_matches!(err, MemoryValidationError::SuspiciousPiiInEntityName { .. });
377    }
378
379    #[test]
380    fn ssn_in_entity_name_rejected() {
381        let r = result_with(vec![entity("123-45-6789")], vec![]);
382        let err = validator().validate_graph_extraction(&r).unwrap_err();
383        assert_matches!(err, MemoryValidationError::SuspiciousPiiInEntityName { .. });
384    }
385
386    #[test]
387    fn disabled_skips_graph_validation() {
388        let v = validator_disabled();
389        let big_entities: Vec<_> = (0..200).map(|i| entity(&format!("E{i}"))).collect();
390        let r = result_with(big_entities, vec![]);
391        assert!(v.validate_graph_extraction(&r).is_ok());
392    }
393
394    // --- exact boundary: max_content_bytes ---
395
396    #[test]
397    fn content_exactly_at_limit_passes() {
398        let v = MemoryWriteValidator::new(MemoryWriteValidationConfig {
399            max_content_bytes: 10,
400            ..MemoryWriteValidationConfig::default()
401        });
402        // Exactly 10 bytes — must pass.
403        assert!(v.validate_memory_save("1234567890").is_ok());
404    }
405
406    #[test]
407    fn content_one_byte_over_limit_rejected() {
408        let v = MemoryWriteValidator::new(MemoryWriteValidationConfig {
409            max_content_bytes: 10,
410            ..MemoryWriteValidationConfig::default()
411        });
412        // 11 bytes — must fail.
413        let err = v.validate_memory_save("12345678901").unwrap_err();
414        assert_matches!(err, MemoryValidationError::ContentTooLarge { .. });
415    }
416
417    // --- multiple forbidden patterns: first match blocks ---
418
419    #[test]
420    fn multiple_forbidden_patterns_first_match_blocks() {
421        let v = MemoryWriteValidator::new(MemoryWriteValidationConfig {
422            forbidden_content_patterns: vec!["<script".to_owned(), "javascript:".to_owned()],
423            ..MemoryWriteValidationConfig::default()
424        });
425        let err = v.validate_memory_save("javascript:alert(1)").unwrap_err();
426        assert_matches!(err, MemoryValidationError::ForbiddenPattern { .. });
427    }
428
429    #[test]
430    fn content_without_forbidden_pattern_passes() {
431        let v = MemoryWriteValidator::new(MemoryWriteValidationConfig {
432            forbidden_content_patterns: vec!["<script".to_owned()],
433            ..MemoryWriteValidationConfig::default()
434        });
435        assert!(v.validate_memory_save("safe content here").is_ok());
436    }
437
438    // --- is_enabled ---
439
440    #[test]
441    fn is_enabled_true_by_default() {
442        assert!(validator().is_enabled());
443    }
444
445    #[test]
446    fn is_enabled_false_when_disabled() {
447        assert!(!validator_disabled().is_enabled());
448    }
449
450    // --- empty ExtractionResult passes ---
451
452    #[test]
453    fn empty_extraction_passes() {
454        let r = result_with(vec![], vec![]);
455        assert!(validator().validate_graph_extraction(&r).is_ok());
456    }
457
458    // --- exact boundary: entity name ---
459
460    #[test]
461    fn entity_name_exactly_at_limit_passes() {
462        let v = MemoryWriteValidator::new(MemoryWriteValidationConfig {
463            max_entity_name_bytes: 5,
464            ..MemoryWriteValidationConfig::default()
465        });
466        let r = result_with(vec![entity("Alice")], vec![]); // 5 bytes exactly
467        assert!(v.validate_graph_extraction(&r).is_ok());
468    }
469
470    #[test]
471    fn entity_name_one_byte_over_limit_rejected() {
472        let v = MemoryWriteValidator::new(MemoryWriteValidationConfig {
473            max_entity_name_bytes: 5,
474            ..MemoryWriteValidationConfig::default()
475        });
476        let r = result_with(vec![entity("AliceX")], vec![]); // 6 bytes
477        let err = v.validate_graph_extraction(&r).unwrap_err();
478        assert_matches!(err, MemoryValidationError::EntityNameTooLong { .. });
479    }
480
481    // --- min entity name length (FIX-3) ---
482
483    #[test]
484    fn entity_name_below_min_rejected() {
485        let r = result_with(vec![entity("go")], vec![]);
486        let err = validator().validate_graph_extraction(&r).unwrap_err();
487        assert_matches!(err, MemoryValidationError::EntityNameTooShort { .. });
488    }
489
490    #[test]
491    fn entity_name_at_min_passes() {
492        let r = result_with(vec![entity("git")], vec![]);
493        assert!(validator().validate_graph_extraction(&r).is_ok());
494    }
495
496    // --- exact boundary: entities count ---
497
498    #[test]
499    fn entities_exactly_at_limit_passes() {
500        let v = MemoryWriteValidator::new(MemoryWriteValidationConfig {
501            max_entities_per_extraction: 3,
502            ..MemoryWriteValidationConfig::default()
503        });
504        let r = result_with(vec![entity("Abc"), entity("Def"), entity("Ghi")], vec![]);
505        assert!(v.validate_graph_extraction(&r).is_ok());
506    }
507
508    // --- error message content ---
509
510    #[test]
511    fn content_too_large_error_message() {
512        let big = "x".repeat(5000);
513        let err = validator().validate_memory_save(&big).unwrap_err();
514        let msg = err.to_string();
515        assert!(msg.contains("5000"), "error must include actual size");
516        assert!(msg.contains("4096"), "error must include max size");
517    }
518
519    // --- forbidden_content_patterns in validate_graph_extraction ---
520
521    #[test]
522    fn forbidden_pattern_in_entity_name_rejected_by_graph_extraction() {
523        let v = MemoryWriteValidator::new(MemoryWriteValidationConfig {
524            forbidden_content_patterns: vec!["internal-host".to_owned()],
525            ..MemoryWriteValidationConfig::default()
526        });
527        let r = result_with(vec![entity("internal-hostname.corp")], vec![]);
528        assert_matches!(
529            v.validate_graph_extraction(&r),
530            Err(MemoryValidationError::ForbiddenPattern { .. })
531        );
532    }
533
534    #[test]
535    fn forbidden_pattern_in_edge_fact_rejected_by_graph_extraction() {
536        let v = MemoryWriteValidator::new(MemoryWriteValidationConfig {
537            forbidden_content_patterns: vec!["BEGIN PRIVATE KEY".to_owned()],
538            ..MemoryWriteValidationConfig::default()
539        });
540        let r = result_with(
541            vec![entity("Alice")],
542            vec![edge("Alice has BEGIN PRIVATE KEY material")],
543        );
544        assert_matches!(
545            v.validate_graph_extraction(&r),
546            Err(MemoryValidationError::ForbiddenPattern { .. })
547        );
548    }
549
550    #[test]
551    fn no_forbidden_patterns_graph_extraction_passes() {
552        let v = MemoryWriteValidator::new(MemoryWriteValidationConfig {
553            forbidden_content_patterns: vec!["<script".to_owned()],
554            ..MemoryWriteValidationConfig::default()
555        });
556        let r = result_with(vec![entity("CleanEntity")], vec![]);
557        assert!(v.validate_graph_extraction(&r).is_ok());
558    }
559}