use super::{AdversarialError, AdversarialResult};
#[derive(Debug, Clone)]
pub struct InputValidator {
pub max_size: usize,
pub verify_checksum: bool,
pub detect_nan: bool,
pub detect_inf: bool,
}
impl Default for InputValidator {
fn default() -> Self {
Self {
max_size: 1024 * 1024 * 1024, verify_checksum: true,
detect_nan: true,
detect_inf: true,
}
}
}
impl InputValidator {
pub fn new() -> Self {
Self::default()
}
pub fn with_max_size(mut self, max_size: usize) -> Self {
self.max_size = max_size;
self
}
pub fn validate_bytes(&self, data: &[u8]) -> AdversarialResult<()> {
if data.is_empty() {
return Err(AdversarialError::ZeroSizeInput);
}
if data.len() > self.max_size {
return Err(AdversarialError::MaxSizeExceeded {
size: data.len(),
max: self.max_size,
});
}
Ok(())
}
pub fn validate_floats(&self, data: &[f32]) -> AdversarialResult<()> {
if data.is_empty() {
return Err(AdversarialError::ZeroSizeInput);
}
let byte_size = std::mem::size_of_val(data);
if byte_size > self.max_size {
return Err(AdversarialError::MaxSizeExceeded {
size: byte_size,
max: self.max_size,
});
}
if self.detect_nan {
for (i, &v) in data.iter().enumerate() {
if v.is_nan() {
return Err(AdversarialError::NaNDetected { index: i });
}
}
}
if self.detect_inf {
for (i, &v) in data.iter().enumerate() {
if v.is_infinite() {
return Err(AdversarialError::InfinityDetected {
index: i,
positive: v.is_sign_positive(),
});
}
}
}
Ok(())
}
pub fn compute_checksum(data: &[u8]) -> u32 {
let mut a: u32 = 1;
let mut b: u32 = 0;
for &byte in data {
a = (a.wrapping_add(u32::from(byte))) % 65521;
b = (b.wrapping_add(a)) % 65521;
}
(b << 16) | a
}
pub fn verify_checksum(&self, data: &[u8], expected: u32) -> AdversarialResult<()> {
if !self.verify_checksum {
return Ok(());
}
let actual = Self::compute_checksum(data);
if actual != expected {
return Err(AdversarialError::CorruptedInput {
byte_index: 0, expected_checksum: expected,
actual_checksum: actual,
});
}
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct ConfigValidator {
pub mins: std::collections::HashMap<String, f64>,
pub maxs: std::collections::HashMap<String, f64>,
}
impl Default for ConfigValidator {
fn default() -> Self {
Self {
mins: std::collections::HashMap::new(),
maxs: std::collections::HashMap::new(),
}
}
}
impl ConfigValidator {
pub fn new() -> Self {
Self::default()
}
pub fn with_bound(mut self, field: &str, min: f64, max: f64) -> Self {
self.mins.insert(field.to_string(), min);
self.maxs.insert(field.to_string(), max);
self
}
pub fn validate_numeric(&self, field: &str, value: f64) -> AdversarialResult<f64> {
if value.is_nan() {
return Err(AdversarialError::ConfigParseError {
field: field.to_string(),
reason: "value is NaN".to_string(),
});
}
if let Some(&min) = self.mins.get(field) {
if value < min {
return Err(AdversarialError::ConfigOutOfBounds {
field: field.to_string(),
value: value.to_string(),
min: min.to_string(),
max: self
.maxs
.get(field)
.map_or("unbounded".to_string(), |m| m.to_string()),
});
}
}
if let Some(&max) = self.maxs.get(field) {
if value > max {
return Err(AdversarialError::ConfigOutOfBounds {
field: field.to_string(),
value: value.to_string(),
min: self
.mins
.get(field)
.map_or("unbounded".to_string(), |m| m.to_string()),
max: max.to_string(),
});
}
}
Ok(value)
}
pub fn validate_toml_string(&self, input: &str) -> AdversarialResult<()> {
let trimmed = input.trim();
if trimmed.is_empty() {
return Err(AdversarialError::ConfigParseError {
field: "root".to_string(),
reason: "empty config".to_string(),
});
}
let open_brackets = trimmed.matches('[').count();
let close_brackets = trimmed.matches(']').count();
if open_brackets != close_brackets {
return Err(AdversarialError::ConfigParseError {
field: "root".to_string(),
reason: format!(
"mismatched brackets: {} open, {} close",
open_brackets, close_brackets
),
});
}
let quotes = trimmed.matches('"').count();
if !quotes.is_multiple_of(2) {
return Err(AdversarialError::ConfigParseError {
field: "root".to_string(),
reason: "unclosed string literal".to_string(),
});
}
Ok(())
}
}