1use thiserror::Error;
20use zeph_memory::graph::extractor::ExtractionResult;
21
22pub use zeph_config::MemoryWriteValidationConfig;
23
24use crate::pii::{EMAIL_RE, SSN_RE};
25
26#[derive(Debug, Error)]
36#[non_exhaustive]
37pub enum MemoryValidationError {
38 #[error("content too large: {size} bytes exceeds max {max}")]
40 ContentTooLarge { size: usize, max: usize },
41
42 #[error("entity name too short: '{name}' is below min {min} bytes")]
44 EntityNameTooShort { name: String, min: usize },
45
46 #[error("entity name too long: '{name}' exceeds max {max} bytes")]
48 EntityNameTooLong { name: String, max: usize },
49
50 #[error("fact text too long: exceeds max {max} bytes")]
52 FactTooLong { max: usize },
53
54 #[error("too many entities: {count} exceeds max {max}")]
56 TooManyEntities { count: usize, max: usize },
57
58 #[error("too many edges: {count} exceeds max {max}")]
60 TooManyEdges { count: usize, max: usize },
61
62 #[error("forbidden pattern detected: {pattern}")]
64 ForbiddenPattern { pattern: String },
65
66 #[error("PII detected in entity name: '{entity}'")]
68 SuspiciousPiiInEntityName { entity: String },
69}
70
71#[derive(Debug, Clone)]
97pub struct MemoryWriteValidator {
98 config: MemoryWriteValidationConfig,
99}
100
101impl MemoryWriteValidator {
102 #[must_use]
114 pub fn new(config: MemoryWriteValidationConfig) -> Self {
115 Self { config }
116 }
117
118 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 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 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 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 #[must_use]
237 pub fn is_enabled(&self) -> bool {
238 self.config.enabled
239 }
240}
241
242#[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 #[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 #[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 #[test]
397 fn content_exactly_at_limit_passes() {
398 let v = MemoryWriteValidator::new(MemoryWriteValidationConfig {
399 max_content_bytes: 10,
400 ..MemoryWriteValidationConfig::default()
401 });
402 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 let err = v.validate_memory_save("12345678901").unwrap_err();
414 assert_matches!(err, MemoryValidationError::ContentTooLarge { .. });
415 }
416
417 #[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 #[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 #[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 #[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![]); 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![]); let err = v.validate_graph_extraction(&r).unwrap_err();
478 assert_matches!(err, MemoryValidationError::EntityNameTooLong { .. });
479 }
480
481 #[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 #[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 #[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 #[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}