use alloc::boxed::Box;
use alloc::string::{String, ToString};
use alloc::vec::Vec;
use regex::Regex;
use crate::nodes::node::{Node, Numeric};
use crate::validation::error::ValidationError;
use crate::validation::messages;
use crate::validation::schema::SchemaType;
pub type ValidationResult = Result<(), ValidationError>;
pub trait Validator {
fn validate(&self, node: &Node) -> ValidationResult;
fn description(&self) -> String;
}
#[derive(Debug, Clone)]
pub struct TypeValidator {
expected_type: SchemaType,
}
impl TypeValidator {
pub fn new(expected_type: SchemaType) -> Self {
Self { expected_type }
}
}
impl Validator for TypeValidator {
fn validate(&self, node: &Node) -> ValidationResult {
let matches = match (&self.expected_type, node) {
(SchemaType::String, Node::Str(_, _, _)) => true,
(SchemaType::Number, Node::Number(_)) => true,
(SchemaType::Integer, Node::Number(Numeric::Integer(_))) => true,
(SchemaType::Integer, Node::Number(Numeric::Int32(_))) => true,
(SchemaType::Integer, Node::Number(Numeric::Int16(_))) => true,
(SchemaType::Integer, Node::Number(Numeric::Int8(_))) => true,
(SchemaType::Float, Node::Number(Numeric::Float(_))) => true,
(SchemaType::Boolean, Node::Boolean(_)) => true,
(SchemaType::Null, Node::None) => true,
(SchemaType::Array, Node::Array(_)) => true,
(SchemaType::Array, Node::Set(_)) => true,
(SchemaType::Object, Node::Mapping(_)) => true,
(SchemaType::Any, _) => true,
_ => false,
};
if matches {
Ok(())
} else {
Err(
crate::validation::engine::ValidationContextCore::fail_type_mismatch(
&self.expected_type,
node,
),
)
}
}
fn description(&self) -> String {
messages::type_must_be(&self.expected_type)
}
}
#[derive(Debug, Clone)]
pub struct RangeValidator {
min: Option<f64>,
max: Option<f64>,
}
impl RangeValidator {
pub fn new(min: Option<f64>, max: Option<f64>) -> Self {
Self { min, max }
}
}
impl Validator for RangeValidator {
fn validate(&self, node: &Node) -> ValidationResult {
let value = match node {
Node::Number(Numeric::Integer(i)) => *i as f64,
Node::Number(Numeric::Float(f)) => *f,
Node::Number(Numeric::UInteger(u)) => *u as f64,
Node::Number(Numeric::Int32(i)) => *i as f64,
Node::Number(Numeric::UInt32(u)) => *u as f64,
Node::Number(Numeric::Int16(i)) => *i as f64,
Node::Number(Numeric::UInt16(u)) => *u as f64,
Node::Number(Numeric::Int8(i)) => *i as f64,
Node::Number(Numeric::Byte(b)) => *b as f64,
_ => {
return Err(ValidationError::InvalidNodeType {
validator: "RangeValidator".to_string(),
found: node_type_name(node).to_string(),
});
}
};
if let Some(min) = self.min {
if value < min {
return Err(
crate::validation::engine::ValidationContextCore::fail_range(
value, self.min, self.max,
),
);
}
}
if let Some(max) = self.max {
if value > max {
return Err(
crate::validation::engine::ValidationContextCore::fail_range(
value, self.min, self.max,
),
);
}
}
Ok(())
}
fn description(&self) -> String {
match (self.min, self.max) {
(Some(min), Some(max)) => messages::value_must_be_between(min, max),
(Some(min), None) => messages::value_must_be_at_least(min),
(None, Some(max)) => messages::value_must_be_at_most(max),
(None, None) => messages::no_range_restriction(),
}
}
}
#[derive(Debug, Clone)]
pub struct LengthValidator {
min: Option<usize>,
max: Option<usize>,
}
impl LengthValidator {
pub fn new(min: Option<usize>, max: Option<usize>) -> Self {
Self { min, max }
}
}
impl Validator for LengthValidator {
fn validate(&self, node: &Node) -> ValidationResult {
let length = match node {
n if n.as_str().is_some() => n.as_str().unwrap().len(),
Node::Array(arr) => arr.len(),
Node::Set(set) => set.len(),
_ => {
return Err(ValidationError::InvalidNodeType {
validator: "LengthValidator".to_string(),
found: node_type_name(node).to_string(),
});
}
};
if let Some(min) = self.min {
if length < min {
return Err(ValidationError::LengthError {
length,
min: self.min,
max: self.max,
});
}
}
if let Some(max) = self.max {
if length > max {
return Err(ValidationError::LengthError {
length,
min: self.min,
max: self.max,
});
}
}
Ok(())
}
fn description(&self) -> String {
match (self.min, self.max) {
(Some(min), Some(max)) => messages::length_must_be_between(min, max),
(Some(min), None) => messages::length_must_be_at_least(min),
(None, Some(max)) => messages::length_must_be_at_most(max),
(None, None) => messages::no_length_restriction(),
}
}
}
#[derive(Debug, Clone)]
pub struct PatternValidator {
regex: Regex,
pattern: String,
}
impl PatternValidator {
pub fn new(pattern: impl Into<String>) -> Self {
let pattern_str = pattern.into();
let regex = Regex::new(&pattern_str).expect("Invalid regex pattern");
Self {
regex,
pattern: pattern_str,
}
}
fn matches(&self, s: &str) -> bool {
self.regex.is_match(s)
}
}
impl Validator for PatternValidator {
fn validate(&self, node: &Node) -> ValidationResult {
if let Some(s) = node.as_str() {
if self.matches(s) {
Ok(())
} else {
Err(ValidationError::PatternMismatch {
pattern: self.pattern.clone(),
value: s.to_string(),
})
}
} else {
Err(ValidationError::InvalidNodeType {
validator: "PatternValidator".to_string(),
found: node_type_name(node).to_string(),
})
}
}
fn description(&self) -> String {
format!("Must match regex pattern: {}", self.pattern)
}
}
#[derive(Debug, Clone)]
pub struct EnumValidator {
allowed: Vec<String>,
}
impl EnumValidator {
pub fn new(allowed: Vec<String>) -> Self {
Self { allowed }
}
fn node_scalar_value(node: &Node) -> Option<String> {
match node {
Node::Str(s, _, _) => Some(s.clone()),
Node::Number(n) => Some(match n {
Numeric::Integer(i) => i.to_string(),
Numeric::Float(f) => f.to_string(),
Numeric::UInteger(u) => u.to_string(),
Numeric::Byte(b) => b.to_string(),
Numeric::Int32(i) => i.to_string(),
Numeric::UInt32(u) => u.to_string(),
Numeric::Int16(i) => i.to_string(),
Numeric::UInt16(u) => u.to_string(),
Numeric::Int8(i) => i.to_string(),
Numeric::UInt8(u) => u.to_string(),
}),
Node::Boolean(b) => Some(b.to_string()),
Node::None => Some("null".to_string()),
_ => None,
}
}
}
impl Validator for EnumValidator {
fn validate(&self, node: &Node) -> ValidationResult {
match Self::node_scalar_value(node) {
Some(value) => {
if self.allowed.contains(&value) {
Ok(())
} else {
Err(ValidationError::EnumMismatch {
allowed: self.allowed.clone(),
value,
})
}
}
None => Err(ValidationError::InvalidNodeType {
validator: "EnumValidator".to_string(),
found: node_type_name(node).to_string(),
}),
}
}
fn description(&self) -> String {
messages::must_be_one_of(&self.allowed)
}
}
#[derive(Debug, Clone)]
pub struct RequiredValidator {
field_name: String,
}
impl RequiredValidator {
pub fn new(field_name: impl Into<String>) -> Self {
Self {
field_name: field_name.into(),
}
}
}
impl Validator for RequiredValidator {
fn validate(&self, node: &Node) -> ValidationResult {
match node {
Node::Mapping(pairs) => {
let found = pairs.iter().any(|(k, _)| match k {
Node::Str(s, _, _) => s == &self.field_name,
Node::Number(n) => {
let key_str = match n {
Numeric::Integer(i) => i.to_string(),
Numeric::Float(f) => f.to_string(),
Numeric::UInteger(u) => u.to_string(),
Numeric::Byte(b) => b.to_string(),
Numeric::Int32(i) => i.to_string(),
Numeric::UInt32(u) => u.to_string(),
Numeric::Int16(i) => i.to_string(),
Numeric::UInt16(u) => u.to_string(),
Numeric::Int8(i) => i.to_string(),
Numeric::UInt8(u) => u.to_string(),
};
key_str == self.field_name
}
Node::Boolean(b) => b.to_string() == self.field_name,
Node::None => self.field_name == "null",
_ => false,
});
if found {
Ok(())
} else {
Err(
crate::validation::engine::ValidationContextCore::fail_required(
&self.field_name,
),
)
}
}
_ => Err(ValidationError::InvalidNodeType {
validator: "RequiredValidator".to_string(),
found: node_type_name(node).to_string(),
}),
}
}
fn description(&self) -> String {
format!("Field '{}' is required", self.field_name)
}
}
pub struct CustomValidator {
validate_fn: Box<dyn Fn(&Node) -> ValidationResult>,
description: String,
}
impl CustomValidator {
pub fn new<F>(description: impl Into<String>, validate_fn: F) -> Self
where
F: Fn(&Node) -> ValidationResult + 'static,
{
Self {
validate_fn: Box::new(validate_fn),
description: description.into(),
}
}
}
impl Validator for CustomValidator {
fn validate(&self, node: &Node) -> ValidationResult {
(self.validate_fn)(node)
}
fn description(&self) -> String {
self.description.clone()
}
}
pub fn node_type_name(node: &Node) -> &'static str {
match node {
Node::Boolean(_) => "Boolean",
Node::Number(_) => "Number",
Node::Str(_, _, _) => "String",
Node::Array(_) => "Array",
Node::Set(_) => "Set",
Node::Mapping(_) => "Mapping",
Node::Comment(_) => "Comment",
Node::Document(_) => "Document",
Node::Anchored(_, _) => "Anchored",
Node::Tagged(_, _) => "Tagged",
Node::Alias(_) => "Alias",
Node::Documents(_) => "Documents",
Node::None => "Null",
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_type_validator() {
let validator = TypeValidator::new(SchemaType::String);
assert!(validator.validate(&Node::from("hello")).is_ok());
assert!(validator.validate(&Node::from(42)).is_err());
}
#[test]
fn test_range_validator() {
let validator = RangeValidator::new(Some(0.0), Some(100.0));
assert!(
validator
.validate(&Node::Number(Numeric::Integer(50)))
.is_ok()
);
assert!(
validator
.validate(&Node::Number(Numeric::Integer(150)))
.is_err()
);
assert!(
validator
.validate(&Node::Number(Numeric::Integer(-10)))
.is_err()
);
}
#[test]
fn test_length_validator() {
let validator = LengthValidator::new(Some(3), Some(10));
assert!(validator.validate(&Node::from("hello")).is_ok());
assert!(validator.validate(&Node::from("hi")).is_err());
assert!(validator.validate(&Node::from("verylongstring")).is_err());
}
#[test]
fn test_pattern_validator() {
let validator = PatternValidator::new("@");
assert!(validator.validate(&Node::from("user@example.com")).is_ok());
assert!(validator.validate(&Node::from("invalid")).is_err());
let validator = PatternValidator::new(r"^[\w.-]+@[\w.-]+\.[a-zA-Z]{2,}$");
assert!(validator.validate(&Node::from("user@example.com")).is_ok());
assert!(validator.validate(&Node::from("user@domain")).is_err());
assert!(validator.validate(&Node::from("@domain.com")).is_err());
assert!(
validator
.validate(&Node::Number(Numeric::Integer(42)))
.is_err()
);
}
#[test]
fn test_enum_validator() {
let validator = EnumValidator::new(vec![
"red".to_string(),
"green".to_string(),
"blue".to_string(),
]);
assert!(validator.validate(&Node::from("red")).is_ok());
assert!(validator.validate(&Node::from("yellow")).is_err());
let validator = EnumValidator::new(vec!["1".to_string(), "2".to_string()]);
assert!(
validator
.validate(&Node::Number(Numeric::Integer(1)))
.is_ok()
);
assert!(
validator
.validate(&Node::Number(Numeric::Integer(3)))
.is_err()
);
let validator = EnumValidator::new(vec!["true".to_string(), "false".to_string()]);
assert!(validator.validate(&Node::Boolean(true)).is_ok());
assert!(validator.validate(&Node::Boolean(false)).is_ok());
assert!(validator.validate(&Node::from("true")).is_ok());
assert!(validator.validate(&Node::from("maybe")).is_err());
let validator = EnumValidator::new(vec!["null".to_string()]);
assert!(validator.validate(&Node::None).is_ok());
assert!(validator.validate(&Node::from("null")).is_ok());
assert!(validator.validate(&Node::from("notnull")).is_err());
assert!(validator.validate(&Node::Array(vec![])).is_err());
}
#[test]
fn test_required_validator() {
let validator = RequiredValidator::new("name");
let mapping = Node::Mapping(vec![
(Node::from("name"), Node::from("Alice")),
(Node::from("age"), Node::from(30)),
]);
assert!(validator.validate(&mapping).is_ok());
let mapping2 = Node::Mapping(vec![(Node::from("age"), Node::from(30))]);
assert!(validator.validate(&mapping2).is_err());
let validator = RequiredValidator::new("42");
let mapping = Node::Mapping(vec![(
Node::Number(Numeric::Integer(42)),
Node::from("answer"),
)]);
assert!(validator.validate(&mapping).is_ok());
let mapping2 = Node::Mapping(vec![(
Node::Number(Numeric::Integer(43)),
Node::from("not answer"),
)]);
assert!(validator.validate(&mapping2).is_err());
let validator = RequiredValidator::new("true");
let mapping = Node::Mapping(vec![(Node::Boolean(true), Node::from("yes"))]);
assert!(validator.validate(&mapping).is_ok());
let mapping2 = Node::Mapping(vec![(Node::Boolean(false), Node::from("no"))]);
assert!(validator.validate(&mapping2).is_err());
let validator = RequiredValidator::new("null");
let mapping = Node::Mapping(vec![(Node::None, Node::from("missing"))]);
assert!(validator.validate(&mapping).is_ok());
let mapping2 = Node::Mapping(vec![(Node::from("notnull"), Node::from("not missing"))]);
assert!(validator.validate(&mapping2).is_err());
assert!(validator.validate(&Node::Array(vec![])).is_err());
}
#[test]
fn test_custom_validator() {
let validator = CustomValidator::new("Must be positive", |node| match node {
Node::Number(Numeric::Integer(i)) if *i > 0 => Ok(()),
Node::Number(Numeric::Integer(_)) => {
Err(ValidationError::Custom("Number must be positive".to_string()).into())
}
_ => Err(ValidationError::Custom("Not a number".to_string()).into()),
});
assert!(
validator
.validate(&Node::Number(Numeric::Integer(10)))
.is_ok()
);
assert!(
validator
.validate(&Node::Number(Numeric::Integer(-5)))
.is_err()
);
}
}
#[cfg(test)]
mod additional_validators_tests {
use super::*;
#[test]
fn test_type_validator_any() {
let validator = TypeValidator::new(SchemaType::Any);
assert!(validator.validate(&Node::from("string")).is_ok());
assert!(validator.validate(&Node::Number(Numeric::Integer(1))).is_ok());
assert!(validator.validate(&Node::Boolean(true)).is_ok());
}
#[test]
fn test_range_validator_invalid_node() {
let validator = RangeValidator::new(Some(0.0), Some(10.0));
let result = validator.validate(&Node::from("not a number"));
assert!(result.is_err());
if let Err(ValidationError::InvalidNodeType { validator: v, .. }) = result {
assert_eq!(v, "RangeValidator");
} else {
panic!("Expected InvalidNodeType error");
}
}
#[test]
fn test_length_validator_string_and_array() {
let validator = LengthValidator::new(Some(2), Some(4));
assert!(validator.validate(&Node::from("ab")).is_ok());
assert!(validator.validate(&Node::from("abcd")).is_ok());
assert!(validator.validate(&Node::from("a")).is_err());
assert!(validator.validate(&Node::from("abcde")).is_err());
let arr = Node::Array(vec![Node::from(1), Node::from(2)]);
assert!(validator.validate(&arr).is_ok());
let arr = Node::Array(vec![Node::from(1)]);
assert!(validator.validate(&arr).is_err());
}
#[test]
fn test_pattern_validator() {
let validator = PatternValidator::new("^abc[0-9]+$".to_string());
assert!(validator.validate(&Node::from("abc123")).is_ok());
assert!(validator.validate(&Node::from("ab123")).is_err());
}
#[test]
fn test_enum_validator() {
let validator = EnumValidator::new(vec!["A".to_string(), "B".to_string()]);
assert!(validator.validate(&Node::from("A")).is_ok());
assert!(validator.validate(&Node::from("C")).is_err());
}
#[test]
fn test_required_validator_non_mapping() {
let validator = RequiredValidator::new("foo");
let node = Node::Array(vec![]);
assert!(validator.validate(&node).is_err());
}
}