use anyhow::{anyhow, Result};
use regex::Regex;
pub const MAX_USER_ID_LENGTH: usize = 128;
pub const MAX_CONTENT_LENGTH: usize = 50_000; pub const MAX_PATTERN_LENGTH: usize = 256; pub const MAX_ENTITY_LENGTH: usize = 256; #[allow(unused)] pub const MAX_METADATA_SIZE: usize = 10_000; #[allow(unused)] pub const MAX_ENTITIES_PER_MEMORY: usize = 50;
pub fn validate_user_id(user_id: &str) -> Result<()> {
if user_id.is_empty() {
return Err(anyhow!("user_id cannot be empty"));
}
if user_id.len() > MAX_USER_ID_LENGTH {
return Err(anyhow!(
"user_id too long: {} chars (max: {})",
user_id.len(),
MAX_USER_ID_LENGTH
));
}
if !user_id
.chars()
.all(|c| c.is_alphanumeric() || c == '-' || c == '_' || c == '@' || c == '.')
{
return Err(anyhow!(
"user_id contains invalid characters (allowed: alphanumeric, -, _, @, .)"
));
}
if user_id.contains("..") {
return Err(anyhow!(
"user_id contains invalid path traversal sequence (..)"
));
}
if user_id.starts_with('.') || user_id.ends_with('.') {
return Err(anyhow!("user_id cannot start or end with a dot"));
}
if std::path::Path::new(user_id).is_absolute() {
return Err(anyhow!("user_id cannot be an absolute path"));
}
{
let upper = user_id.to_uppercase();
let stem = upper.split('.').next().unwrap_or(&upper);
const DEVICE_NAMES: &[&str] = &[
"CON", "PRN", "AUX", "NUL", "COM1", "COM2", "COM3", "COM4", "COM5", "COM6", "COM7",
"COM8", "COM9", "LPT1", "LPT2", "LPT3", "LPT4", "LPT5", "LPT6", "LPT7", "LPT8", "LPT9",
];
if DEVICE_NAMES.contains(&stem) {
return Err(anyhow!(
"user_id cannot be a Windows reserved device name: {}",
user_id
));
}
}
Ok(())
}
pub fn validate_memory_id(memory_id: &str) -> Result<uuid::Uuid> {
uuid::Uuid::parse_str(memory_id).map_err(|e| anyhow!("Invalid memory_id UUID format: {e}"))
}
pub fn validate_memory_id_or_prefix(memory_id: &str) -> Result<Option<uuid::Uuid>> {
if let Ok(uuid) = uuid::Uuid::parse_str(memory_id) {
return Ok(Some(uuid));
}
let trimmed = memory_id.trim();
if trimmed.len() < 8 {
return Err(anyhow!(
"Memory ID must be a full UUID or at least 8 hex characters, got {} chars",
trimmed.len()
));
}
if !trimmed.chars().all(|c| c.is_ascii_hexdigit()) {
return Err(anyhow!(
"Memory ID prefix contains invalid characters (only hex digits 0-9, a-f allowed)"
));
}
Ok(None)
}
pub const MIN_MEANINGFUL_CONTENT_LENGTH: usize = 10;
pub fn validate_content(content: &str, allow_empty: bool) -> Result<()> {
let trimmed = content.trim();
if !allow_empty && trimmed.is_empty() {
return Err(anyhow!("content cannot be empty"));
}
if !allow_empty && !trimmed.is_empty() && trimmed.len() < MIN_MEANINGFUL_CONTENT_LENGTH {
return Err(anyhow!(
"content too short: {} chars (min: {})",
trimmed.len(),
MIN_MEANINGFUL_CONTENT_LENGTH
));
}
if content.len() > MAX_CONTENT_LENGTH {
return Err(anyhow!(
"content too long: {} bytes (max: {})",
content.len(),
MAX_CONTENT_LENGTH
));
}
Ok(())
}
pub fn validate_embeddings(embeddings: &[f32]) -> Result<()> {
if embeddings.is_empty() {
return Err(anyhow!("embeddings cannot be empty"));
}
let valid_dims = [128, 256, 384, 512, 768, 1024, 1536, 2048];
if !valid_dims.contains(&embeddings.len()) {
return Err(anyhow!(
"Unusual embedding dimension: {}. Common dimensions: {:?}",
embeddings.len(),
valid_dims
));
}
if embeddings.iter().any(|&v| !v.is_finite()) {
return Err(anyhow!("embeddings contain NaN or Inf values"));
}
Ok(())
}
pub fn validate_importance_threshold(threshold: f32) -> Result<()> {
if !(0.0..=1.0).contains(&threshold) {
return Err(anyhow!(
"importance_threshold must be between 0.0 and 1.0, got: {threshold}"
));
}
Ok(())
}
pub fn validate_max_results(max_results: usize) -> Result<()> {
if max_results == 0 {
return Err(anyhow!("max_results must be greater than 0"));
}
if max_results > 10_000 {
return Err(anyhow!(
"max_results too large: {max_results} (max: 10,000)"
));
}
Ok(())
}
pub fn validate_and_compile_pattern(pattern: &str) -> Result<Regex> {
if pattern.is_empty() {
return Err(anyhow!("Pattern cannot be empty"));
}
if pattern.len() > MAX_PATTERN_LENGTH {
return Err(anyhow!(
"Pattern too long: {} chars (max: {})",
pattern.len(),
MAX_PATTERN_LENGTH
));
}
Regex::new(pattern).map_err(|e| anyhow!("Invalid regex pattern: {e}"))
}
pub fn validate_entity(entity: &str) -> Result<()> {
if entity.is_empty() {
return Err(anyhow!("Entity name cannot be empty"));
}
if entity.len() > MAX_ENTITY_LENGTH {
return Err(anyhow!(
"Entity name too long: {} chars (max: {})",
entity.len(),
MAX_ENTITY_LENGTH
));
}
if entity.chars().any(|c| c.is_control()) {
return Err(anyhow!("Entity name contains invalid control characters"));
}
if entity.contains("..") || entity.contains('/') || entity.contains('\\') {
return Err(anyhow!("Entity name contains invalid path characters"));
}
Ok(())
}
#[allow(unused)] pub fn validate_entities(entities: &[String]) -> Result<()> {
if entities.len() > MAX_ENTITIES_PER_MEMORY {
return Err(anyhow!(
"Too many entities: {} (max: {})",
entities.len(),
MAX_ENTITIES_PER_MEMORY
));
}
for entity in entities {
validate_entity(entity)?;
}
Ok(())
}
#[allow(unused)] pub fn validate_metadata(metadata: &serde_json::Value) -> Result<()> {
let size = metadata.to_string().len();
if size > MAX_METADATA_SIZE {
return Err(anyhow!(
"Metadata too large: {size} bytes (max: {MAX_METADATA_SIZE})"
));
}
Ok(())
}
pub fn validate_relationship_strength(strength: f32) -> Result<()> {
if !(0.0..=1.0).contains(&strength) {
return Err(anyhow!(
"Relationship strength must be between 0.0 and 1.0, got: {strength}"
));
}
Ok(())
}
pub fn validate_weight(name: &str, value: f32) -> Result<()> {
if !value.is_finite() || !(0.0..=1.0).contains(&value) {
return Err(anyhow!("{name} must be between 0.0 and 1.0, got: {value}"));
}
Ok(())
}
pub fn validate_geo_location(geo: &[f64; 3]) -> Result<()> {
if !geo[0].is_finite() || !(-90.0..=90.0).contains(&geo[0]) {
return Err(anyhow!(
"latitude must be between -90.0 and 90.0, got: {}",
geo[0]
));
}
if !geo[1].is_finite() || !(-180.0..=180.0).contains(&geo[1]) {
return Err(anyhow!(
"longitude must be between -180.0 and 180.0, got: {}",
geo[1]
));
}
if !geo[2].is_finite() {
return Err(anyhow!("altitude must be a finite number, got: {}", geo[2]));
}
Ok(())
}
pub fn validate_geo_filter(lat: f64, lon: f64, radius_meters: f64) -> Result<()> {
if !lat.is_finite() || !(-90.0..=90.0).contains(&lat) {
return Err(anyhow!(
"geo_filter latitude must be between -90.0 and 90.0, got: {lat}"
));
}
if !lon.is_finite() || !(-180.0..=180.0).contains(&lon) {
return Err(anyhow!(
"geo_filter longitude must be between -180.0 and 180.0, got: {lon}"
));
}
if !radius_meters.is_finite() || radius_meters <= 0.0 {
return Err(anyhow!(
"geo_filter radius_meters must be > 0, got: {radius_meters}"
));
}
if radius_meters > 40_075_000.0 {
return Err(anyhow!(
"geo_filter radius_meters exceeds Earth's circumference: {radius_meters}"
));
}
Ok(())
}
pub fn validate_reminder_timestamp(at: &chrono::DateTime<chrono::Utc>) -> Result<()> {
let now = chrono::Utc::now();
let max_future = now + chrono::Duration::days(365 * 5); let max_past = now - chrono::Duration::hours(1);
if *at < max_past {
return Err(anyhow!("Reminder timestamp is in the past: {at}"));
}
if *at > max_future {
return Err(anyhow!(
"Reminder timestamp is too far in the future (max 5 years): {at}"
));
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_valid_user_id() {
assert!(validate_user_id("alice").is_ok());
assert!(validate_user_id("user-123").is_ok());
assert!(validate_user_id("test_user").is_ok());
assert!(validate_user_id("user@example.com").is_ok());
}
#[test]
fn test_invalid_user_id() {
assert!(validate_user_id("").is_err()); assert!(validate_user_id("user/123").is_err()); assert!(validate_user_id(&"a".repeat(200)).is_err()); }
#[test]
fn test_path_traversal_prevention() {
assert!(validate_user_id("user..admin").is_err()); assert!(validate_user_id("..").is_err()); assert!(validate_user_id("a..b..c").is_err()); assert!(validate_user_id(".hidden").is_err()); assert!(validate_user_id("user.").is_err()); assert!(validate_user_id("user.name@example.com").is_ok());
assert!(validate_user_id("first.last").is_ok());
}
#[test]
fn test_valid_content() {
assert!(validate_content("Hello world", false).is_ok());
assert!(validate_content("", true).is_ok()); }
#[test]
fn test_invalid_content() {
assert!(validate_content("", false).is_err()); assert!(validate_content(&"x".repeat(100_000), false).is_err()); }
#[test]
fn test_content_min_length_gate() {
assert!(validate_content("TT", false).is_err());
assert!(validate_content("OK", false).is_err());
assert!(validate_content("WTesti", false).is_err());
assert!(validate_content("A", false).is_err());
assert!(validate_content(" TT ", false).is_err());
assert!(validate_content("short note", false).is_ok()); assert!(validate_content("real memory content here", false).is_ok());
assert!(validate_content("", true).is_ok());
assert!(validate_content("TT", true).is_ok());
}
#[test]
fn test_valid_embeddings() {
let emb_384 = vec![0.5_f32; 384];
assert!(validate_embeddings(&emb_384).is_ok());
let emb_768 = vec![0.5_f32; 768];
assert!(validate_embeddings(&emb_768).is_ok());
}
#[test]
fn test_invalid_embeddings() {
assert!(validate_embeddings(&[]).is_err()); assert!(validate_embeddings(&[f32::NAN, 0.5]).is_err()); assert!(validate_embeddings(&vec![0.5; 999]).is_err()); }
#[test]
fn test_importance_threshold() {
assert!(validate_importance_threshold(0.0).is_ok());
assert!(validate_importance_threshold(0.5).is_ok());
assert!(validate_importance_threshold(1.0).is_ok());
assert!(validate_importance_threshold(-0.1).is_err());
assert!(validate_importance_threshold(1.5).is_err());
}
#[test]
fn test_max_results() {
assert!(validate_max_results(1).is_ok());
assert!(validate_max_results(100).is_ok());
assert!(validate_max_results(10_000).is_ok());
assert!(validate_max_results(0).is_err());
assert!(validate_max_results(20_000).is_err());
}
#[test]
fn test_valid_patterns() {
assert!(validate_and_compile_pattern("hello").is_ok());
assert!(validate_and_compile_pattern("user.*").is_ok());
assert!(validate_and_compile_pattern("[a-z]+").is_ok());
assert!(validate_and_compile_pattern("^start").is_ok());
assert!(validate_and_compile_pattern("end$").is_ok());
}
#[test]
fn test_regex_edge_cases() {
assert!(validate_and_compile_pattern("(a+)+").is_ok());
assert!(validate_and_compile_pattern("(.*)*").is_ok());
assert!(validate_and_compile_pattern("(.+)+").is_ok());
assert!(validate_and_compile_pattern(&"a".repeat(300)).is_err());
assert!(validate_and_compile_pattern("").is_err());
}
#[test]
fn test_valid_entity() {
assert!(validate_entity("user").is_ok());
assert!(validate_entity("John Doe").is_ok());
assert!(validate_entity("entity-123").is_ok());
}
#[test]
fn test_invalid_entity() {
assert!(validate_entity("").is_err()); assert!(validate_entity(&"a".repeat(300)).is_err()); assert!(validate_entity("../etc/passwd").is_err()); assert!(validate_entity("entity\x00null").is_err()); }
#[test]
fn test_entities_list() {
let valid: Vec<String> = vec!["a".to_string(), "b".to_string()];
assert!(validate_entities(&valid).is_ok());
let too_many: Vec<String> = (0..100).map(|i| format!("entity{i}")).collect();
assert!(validate_entities(&too_many).is_err());
}
#[test]
fn test_memory_id_or_prefix_full_uuid() {
let result = validate_memory_id_or_prefix("c77bb954-1234-5678-abcd-ef0123456789");
assert!(result.is_ok());
assert!(result.unwrap().is_some());
}
#[test]
fn test_memory_id_or_prefix_valid_prefix() {
let result = validate_memory_id_or_prefix("c77bb954");
assert!(result.is_ok());
assert!(result.unwrap().is_none());
}
#[test]
fn test_memory_id_or_prefix_long_prefix() {
let result = validate_memory_id_or_prefix("c77bb9541234abcd");
assert!(result.is_ok());
assert!(result.unwrap().is_none());
}
#[test]
fn test_memory_id_or_prefix_too_short() {
assert!(validate_memory_id_or_prefix("c77bb").is_err());
}
#[test]
fn test_memory_id_or_prefix_invalid_chars() {
assert!(validate_memory_id_or_prefix("c77bb95z").is_err());
}
#[test]
fn test_memory_id_or_prefix_empty() {
assert!(validate_memory_id_or_prefix("").is_err());
}
#[test]
fn test_relationship_strength() {
assert!(validate_relationship_strength(0.0).is_ok());
assert!(validate_relationship_strength(0.5).is_ok());
assert!(validate_relationship_strength(1.0).is_ok());
assert!(validate_relationship_strength(-0.1).is_err());
assert!(validate_relationship_strength(1.1).is_err());
}
#[test]
fn test_validate_weight() {
assert!(validate_weight("test", 0.0).is_ok());
assert!(validate_weight("test", 0.5).is_ok());
assert!(validate_weight("test", 1.0).is_ok());
assert!(validate_weight("test", -0.1).is_err());
assert!(validate_weight("test", 1.1).is_err());
assert!(validate_weight("test", f32::NAN).is_err());
assert!(validate_weight("test", f32::INFINITY).is_err());
}
#[test]
fn test_validate_geo_location() {
assert!(validate_geo_location(&[37.7749, -122.4194, 10.0]).is_ok());
assert!(validate_geo_location(&[0.0, 0.0, 0.0]).is_ok());
assert!(validate_geo_location(&[-90.0, -180.0, -100.0]).is_ok());
assert!(validate_geo_location(&[90.0, 180.0, 8848.0]).is_ok());
assert!(validate_geo_location(&[91.0, 0.0, 0.0]).is_err());
assert!(validate_geo_location(&[-91.0, 0.0, 0.0]).is_err());
assert!(validate_geo_location(&[0.0, 181.0, 0.0]).is_err());
assert!(validate_geo_location(&[0.0, -181.0, 0.0]).is_err());
assert!(validate_geo_location(&[f64::NAN, 0.0, 0.0]).is_err());
assert!(validate_geo_location(&[0.0, 0.0, f64::INFINITY]).is_err());
}
#[test]
fn test_validate_geo_filter() {
assert!(validate_geo_filter(37.7749, -122.4194, 1000.0).is_ok());
assert!(validate_geo_filter(0.0, 0.0, 1.0).is_ok());
assert!(validate_geo_filter(91.0, 0.0, 100.0).is_err());
assert!(validate_geo_filter(0.0, 181.0, 100.0).is_err());
assert!(validate_geo_filter(0.0, 0.0, 0.0).is_err());
assert!(validate_geo_filter(0.0, 0.0, -1.0).is_err());
assert!(validate_geo_filter(0.0, 0.0, 50_000_000.0).is_err()); }
#[test]
fn test_validate_reminder_timestamp() {
let now = chrono::Utc::now();
assert!(validate_reminder_timestamp(&(now + chrono::Duration::hours(1))).is_ok());
assert!(validate_reminder_timestamp(&(now - chrono::Duration::minutes(30))).is_ok());
assert!(validate_reminder_timestamp(&(now - chrono::Duration::hours(2))).is_err());
assert!(validate_reminder_timestamp(&(now + chrono::Duration::days(365 * 10))).is_err());
}
}