use crate::error::{Error, Result};
use crate::prelude::*;
use crate::value::{Number, Value};
pub fn validate_against_schema(value: &Value, schema: &Value) -> Result<()> {
CompiledSchema::compile(schema)?.validate(value)
}
pub const DEFAULT_MAX_SCHEMA_ERRORS: usize = 100;
pub const DEFAULT_MAX_SCHEMA_ERROR_BYTES: usize = 64 * 1024;
pub const MAX_SCHEMA_MESSAGE_BYTES: usize = 1024;
pub const DEFAULT_MAX_SCHEMA_NODES: usize = 100_000;
pub const DEFAULT_MAX_SCHEMA_DEPTH: usize = 128;
pub struct CompiledSchema {
validator: jsonschema::Validator,
limits: ErrorLimits,
}
#[derive(Debug, Clone, Copy)]
struct ErrorLimits {
max_errors: usize,
max_bytes: usize,
}
impl Default for ErrorLimits {
fn default() -> Self {
Self {
max_errors: DEFAULT_MAX_SCHEMA_ERRORS,
max_bytes: DEFAULT_MAX_SCHEMA_ERROR_BYTES,
}
}
}
impl fmt::Debug for CompiledSchema {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("CompiledSchema").finish_non_exhaustive()
}
}
impl CompiledSchema {
pub fn compile(schema: &Value) -> Result<Self> {
Self::builder(schema).build()
}
#[must_use]
pub fn builder(schema: &Value) -> CompiledSchemaBuilder {
CompiledSchemaBuilder {
schema_json: value_to_json(schema),
validate_formats: None,
formats: Vec::new(),
limits: ErrorLimits::default(),
max_schema_nodes: DEFAULT_MAX_SCHEMA_NODES,
max_schema_depth: DEFAULT_MAX_SCHEMA_DEPTH,
backtrack_limit: None,
}
}
pub fn validate(&self, value: &Value) -> Result<()> {
let instance_json = value_to_json(value)
.map_err(|e| Error::Parse(format!("validate_against_schema: value -> JSON: {e}")))?;
let report = self.collect(&instance_json);
if report.found == 0 {
return Ok(());
}
Err(Error::Custom(report.summary()))
}
pub fn iter_errors(&self, value: &Value) -> Result<Vec<SchemaViolation>> {
let instance_json = value_to_json(value)
.map_err(|e| Error::Parse(format!("iter_errors: value -> JSON: {e}")))?;
Ok(self.collect(&instance_json).violations)
}
fn collect(&self, instance: &serde_json::Value) -> Report {
let mut report = Report::default();
let mut bytes_left = self.limits.max_bytes;
for err in self.validator.iter_errors(instance) {
report.found += 1;
if report.violations.len() >= self.limits.max_errors || bytes_left == 0 {
continue;
}
let violation = SchemaViolation {
instance_path: bounded_display(&err.instance_path(), bytes_left),
keyword: keyword_of(&err),
message: bounded_display(&err, bytes_left.min(MAX_SCHEMA_MESSAGE_BYTES)),
};
let cost = violation.message.len() + violation.instance_path.len();
bytes_left = bytes_left.saturating_sub(cost);
report.violations.push(violation);
}
report
}
}
fn keyword_of(err: &jsonschema::ValidationError<'_>) -> String {
err.schema_path()
.to_string()
.rsplit('/')
.next()
.unwrap_or_default()
.to_string()
}
#[derive(Default)]
struct Report {
violations: Vec<SchemaViolation>,
found: usize,
}
impl Report {
fn summary(&self) -> String {
let lines: Vec<String> = self
.violations
.iter()
.map(|v| format!("{} (at `{}`)", v.message, v.instance_path))
.collect();
if self.found == 1 && lines.len() == 1 {
return format!("schema violation: {}", lines[0]);
}
let mut out = format!("schema violations ({} total):", self.found);
for line in &lines {
out.push_str("\n - ");
out.push_str(line);
}
let omitted = self.found - lines.len();
if omitted > 0 {
out.push_str(&format!("\n - ... and {omitted} more not shown"));
}
out
}
}
fn bounded_display(item: &dyn fmt::Display, limit: usize) -> String {
use fmt::Write as _;
let mut out = BoundedWriter {
buf: String::new(),
limit,
cut: false,
};
let _ = write!(out, "{item}");
if out.cut {
out.buf.push_str("...");
}
out.buf
}
struct BoundedWriter {
buf: String,
limit: usize,
cut: bool,
}
impl fmt::Write for BoundedWriter {
fn write_str(&mut self, s: &str) -> fmt::Result {
let room = self.limit - self.buf.len();
if s.len() <= room {
self.buf.push_str(s);
return Ok(());
}
let mut end = room;
while !s.is_char_boundary(end) {
end -= 1;
}
self.buf.push_str(&s[..end]);
self.cut = true;
Err(fmt::Error)
}
}
pub struct CompiledSchemaBuilder {
schema_json: core::result::Result<serde_json::Value, String>,
validate_formats: Option<bool>,
#[allow(clippy::type_complexity)]
formats: Vec<(String, Arc<dyn Fn(&str) -> bool + Send + Sync>)>,
limits: ErrorLimits,
max_schema_nodes: usize,
max_schema_depth: usize,
backtrack_limit: Option<usize>,
}
impl fmt::Debug for CompiledSchemaBuilder {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("CompiledSchemaBuilder")
.field("validate_formats", &self.validate_formats)
.field("limits", &self.limits)
.field("max_schema_nodes", &self.max_schema_nodes)
.field("max_schema_depth", &self.max_schema_depth)
.field("backtrack_limit", &self.backtrack_limit)
.field(
"formats",
&self.formats.iter().map(|(n, _)| n).collect::<Vec<_>>(),
)
.finish_non_exhaustive()
}
}
impl CompiledSchemaBuilder {
#[must_use]
pub fn validate_formats(mut self, yes: bool) -> Self {
self.validate_formats = Some(yes);
self
}
#[must_use]
pub fn with_format<N, F>(mut self, name: N, check: F) -> Self
where
N: Into<String>,
F: Fn(&str) -> bool + Send + Sync + 'static,
{
self.formats.push((name.into(), Arc::new(check)));
self
}
#[must_use]
pub fn max_errors(mut self, n: usize) -> Self {
self.limits.max_errors = n;
self
}
#[must_use]
pub fn max_error_bytes(mut self, n: usize) -> Self {
self.limits.max_bytes = n;
self
}
#[must_use]
pub fn max_schema_nodes(mut self, n: usize) -> Self {
self.max_schema_nodes = n;
self
}
#[must_use]
pub fn max_schema_depth(mut self, n: usize) -> Self {
self.max_schema_depth = n;
self
}
#[must_use]
pub fn backtracking_patterns(mut self, limit: usize) -> Self {
self.backtrack_limit = Some(limit);
self
}
pub fn build(self) -> Result<CompiledSchema> {
let schema_json = self
.schema_json
.map_err(|e| Error::Custom(format!("validate_against_schema: schema -> JSON: {e}")))?;
check_schema_shape(&schema_json, self.max_schema_nodes, self.max_schema_depth)?;
let mut options = hardened_options();
if let Some(limit) = self.backtrack_limit {
options = options.with_pattern_options(
jsonschema::PatternOptions::fancy_regex().backtrack_limit(limit),
);
}
if let Some(yes) = self.validate_formats {
options = options.should_validate_formats(yes);
}
for (name, check) in self.formats {
options = options.with_format(name, move |s: &str| check(s));
}
let validator = options.build(&schema_json).map_err(|e| {
Error::Custom(format!(
"validate_against_schema: schema is not a valid JSON Schema: {e}"
))
})?;
Ok(CompiledSchema {
validator,
limits: self.limits,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SchemaViolation {
pub instance_path: String,
pub keyword: String,
pub message: String,
}
pub fn validate_against_schema_str(yaml: &str, schema_yaml: &str) -> Result<()> {
let value: Value = crate::from_str(yaml)?;
let schema: Value = crate::from_str(schema_yaml)?;
validate_against_schema(&value, &schema)
}
fn hardened_options() -> jsonschema::ValidationOptions<'static> {
jsonschema::options()
.with_retriever(RefuseExternalRefs)
.with_pattern_options(jsonschema::PatternOptions::regex())
}
pub(crate) fn compile_for_coercion(schema: &Value, context: &str) -> Result<jsonschema::Validator> {
let schema_json = value_to_json(schema)
.map_err(|e| Error::Custom(format!("{context}: schema -> JSON: {e}")))?;
check_schema_shape(
&schema_json,
DEFAULT_MAX_SCHEMA_NODES,
DEFAULT_MAX_SCHEMA_DEPTH,
)?;
hardened_options()
.build(&schema_json)
.map_err(|e| Error::Custom(format!("{context}: schema is not a valid JSON Schema: {e}")))
}
fn check_schema_shape(
schema: &serde_json::Value,
max_nodes: usize,
max_depth: usize,
) -> Result<()> {
let mut pending: Vec<(&serde_json::Value, usize)> = vec![(schema, 0)];
let mut nodes = 0usize;
while let Some((node, depth)) = pending.pop() {
nodes += 1;
if nodes > max_nodes {
return Err(Error::Custom(format!(
"validate_against_schema: schema exceeds {max_nodes} nodes"
)));
}
let children: Vec<&serde_json::Value> = match node {
serde_json::Value::Array(items) => items.iter().collect(),
serde_json::Value::Object(map) => map.values().collect(),
_ => continue,
};
if depth + 1 > max_depth {
return Err(Error::Custom(format!(
"validate_against_schema: schema nesting exceeds the recursion depth limit of {max_depth}"
)));
}
pending.extend(children.into_iter().map(|c| (c, depth + 1)));
}
Ok(())
}
struct RefuseExternalRefs;
impl jsonschema::Retrieve for RefuseExternalRefs {
fn retrieve(
&self,
uri: &jsonschema::Uri<String>,
) -> core::result::Result<serde_json::Value, Box<dyn core::error::Error + Send + Sync>> {
Err(Box::new(ExternalRefRefused(uri.as_str().to_string())))
}
}
#[derive(Debug)]
struct ExternalRefRefused(String);
impl fmt::Display for ExternalRefRefused {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"external $ref `{}` refused: only references inside the schema document are resolved",
self.0
)
}
}
impl core::error::Error for ExternalRefRefused {}
pub(crate) fn value_to_json(v: &Value) -> core::result::Result<serde_json::Value, String> {
serde_json::to_value(v).map_err(|e| e.to_string())
}
pub fn coerce_to_schema(value: &mut Value, schema: &Value) -> Result<usize> {
use jsonschema::JsonType;
use jsonschema::error::{TypeKind, ValidationErrorKind};
let validator = compile_for_coercion(schema, "coerce_to_schema")?;
let mut applied: usize = 0;
let max_iterations = 1024;
for _ in 0..max_iterations {
let instance_json = value_to_json(value)
.map_err(|e| Error::Parse(format!("coerce_to_schema: value -> JSON: {e}")))?;
let mut applied_this_pass = false;
let mut targets: Vec<(String, JsonType)> = Vec::new();
for err in validator.iter_errors(&instance_json) {
if let ValidationErrorKind::Type {
kind: TypeKind::Single(target),
} = err.kind()
{
targets.push((err.instance_path().to_string(), *target));
}
}
for (path, target) in targets {
let segments = parse_json_pointer(&path);
if let Some(node) = navigate_mut(value, &segments) {
if try_coerce(node, target) {
applied += 1;
applied_this_pass = true;
}
}
}
if !applied_this_pass {
break;
}
}
Ok(applied)
}
fn parse_json_pointer(s: &str) -> Vec<String> {
if s.is_empty() || s == "/" {
return Vec::new();
}
s.trim_start_matches('/')
.split('/')
.map(|seg| seg.replace("~1", "/").replace("~0", "~"))
.collect()
}
fn navigate_mut<'a>(value: &'a mut Value, path: &[String]) -> Option<&'a mut Value> {
let mut cursor = value;
for seg in path {
cursor = match cursor {
Value::Mapping(m) => m.get_mut(seg.as_str())?,
Value::Sequence(s) => {
let idx: usize = seg.parse().ok()?;
s.get_mut(idx)?
}
_ => return None,
};
}
Some(cursor)
}
fn try_coerce(node: &mut Value, target: jsonschema::JsonType) -> bool {
use jsonschema::JsonType;
let s = match node {
Value::String(s) => s.clone(),
_ => return false,
};
let coerced = match target {
JsonType::Integer => s
.parse::<i64>()
.ok()
.map(|n| Value::Number(Number::Integer(n))),
JsonType::Number => s
.parse::<f64>()
.ok()
.map(|f| Value::Number(Number::Float(f))),
JsonType::Boolean => match s.as_str() {
"true" => Some(Value::Bool(true)),
"false" => Some(Value::Bool(false)),
_ => None,
},
_ => None,
};
match coerced {
Some(new_v) => {
*node = new_v;
true
}
None => false,
}
}
#[cfg(test)]
mod tests {
use super::*;
fn parse(s: &str) -> Value {
crate::from_str(s).unwrap()
}
#[test]
fn valid_value_returns_ok() {
let schema =
parse("type: object\nrequired: [port]\nproperties:\n port:\n type: integer\n");
let value = parse("port: 8080\n");
assert!(validate_against_schema(&value, &schema).is_ok());
}
#[test]
fn type_mismatch_returns_err() {
let schema = parse("type: object\nproperties:\n port:\n type: integer\n");
let value = parse("port: hello\n");
let err = validate_against_schema(&value, &schema).unwrap_err();
let msg = err.to_string();
assert!(msg.contains("schema violation"), "got: {msg}");
assert!(msg.contains("/port"), "path missing: {msg}");
}
#[test]
fn missing_required_field_returns_err() {
let schema = parse("type: object\nrequired: [port]\n");
let value = parse("host: localhost\n");
let err = validate_against_schema(&value, &schema).unwrap_err();
assert!(err.to_string().contains("port"));
}
#[test]
fn multiple_violations_aggregated() {
let schema = parse(
"type: object
required: [port, host]
properties:
port:
type: integer
host:
type: string
",
);
let value = parse("port: not-int\n");
let err = validate_against_schema(&value, &schema).unwrap_err();
let msg = err.to_string();
assert!(msg.contains("schema violations"), "got: {msg}");
assert!(msg.contains("port"));
assert!(msg.contains("host"));
}
#[test]
fn invalid_schema_distinguished_from_invalid_data() {
let schema = parse("type: 42\n");
let value = parse("anything: 1\n");
let err = validate_against_schema(&value, &schema).unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("not a valid JSON Schema"),
"expected schema-side error, got: {msg}"
);
}
#[test]
fn enum_constraint_enforced() {
let schema = parse(
"type: object
properties:
level:
enum: [trace, debug, info, warn, error]
",
);
assert!(validate_against_schema(&parse("level: warn\n"), &schema).is_ok());
assert!(validate_against_schema(&parse("level: ULTRA\n"), &schema).is_err());
}
#[test]
fn integer_bounds_enforced() {
let schema = parse(
"type: object
properties:
port:
type: integer
minimum: 0
maximum: 65535
",
);
assert!(validate_against_schema(&parse("port: 8080\n"), &schema).is_ok());
assert!(validate_against_schema(&parse("port: 70000\n"), &schema).is_err());
assert!(validate_against_schema(&parse("port: -1\n"), &schema).is_err());
}
#[test]
fn nested_object_validated() {
let schema = parse(
"type: object
properties:
db:
type: object
required: [host]
properties:
host:
type: string
",
);
let good = parse("db:\n host: localhost\n");
let bad = parse("db: {}\n");
assert!(validate_against_schema(&good, &schema).is_ok());
assert!(validate_against_schema(&bad, &schema).is_err());
}
#[test]
fn validate_against_schema_str_parses_both_inputs() {
let schema = "type: object\nrequired: [port]\n";
let yaml = "port: 8080\n";
assert!(validate_against_schema_str(yaml, schema).is_ok());
}
#[test]
fn schema_for_codegen_round_trip_validates_self() {
#[derive(serde::Serialize, serde::Deserialize, crate::JsonSchema)]
#[allow(dead_code)]
struct Cfg {
port: u16,
#[serde(default)]
host: String,
}
let schema = crate::schema_for::<Cfg>().unwrap();
let good = parse("port: 8080\nhost: localhost\n");
assert!(validate_against_schema(&good, &schema).is_ok());
let bad = parse("host: localhost\n"); let err = validate_against_schema(&bad, &schema).unwrap_err();
assert!(err.to_string().contains("port"));
}
}