diff --git a/crates/protovalidate/src/error.rs b/crates/protovalidate/src/error.rs index 2ec217f1..d4b36aa1 100644 --- a/crates/protovalidate/src/error.rs +++ b/crates/protovalidate/src/error.rs @@ -12,7 +12,12 @@ // See the License for the specific language governing permissions and // limitations under the License. -use std::fmt; +use std::fmt::{self, Write as _}; + +use buffa::Message as _; + +use crate::validate::__buffa::oneof::field_path_element::Subscript; +use crate::validate::{FieldPath, Violation as ViolationPb, Violations}; /// A descriptor that could not be registered. #[derive(Debug)] @@ -80,21 +85,25 @@ impl std::error::Error for Error { } } -/// One or more rule violations, carried by [`Error::Validation`] as a -/// serialized `buf.validate.Violations`. +/// One or more rule violations, carried by [`Error::Validation`]. pub struct ValidationError { - violations: Vec, + /// Never empty. + violations: Vec, } impl ValidationError { - pub(crate) fn new(violations: Vec) -> Self { + pub(crate) fn new(violations: Vec) -> Self { Self { violations } } - /// The violations, as a serialized `buf.validate.Violations`. + /// Encodes the violations as a `buf.validate.Violations`. #[must_use] - pub fn violations(&self) -> &[u8] { - &self.violations + pub fn encode_violations(&self) -> Vec { + Violations { + violations: self.violations.clone(), + ..Default::default() + } + .encode_to_vec() } } @@ -108,8 +117,151 @@ impl fmt::Debug for ValidationError { } } +/// The first violation, followed by the number of others. +/// For example, `user.email: must be a valid email address, and 2 more violations`. impl fmt::Display for ValidationError { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - f.write_str("validation failed") + write_violation(f, &self.violations[0])?; + match self.violations.len() - 1 { + 0 => Ok(()), + 1 => f.write_str(", and 1 more violation"), + more => write!(f, ", and {more} more violations"), + } + } +} + +/// Writes `violation` as `: `, or `[]` in place +/// of an empty message, omitting an empty field path. +fn write_violation(f: &mut fmt::Formatter<'_>, violation: &ViolationPb) -> fmt::Result { + if !violation.field.elements.is_empty() { + write_field_path(f, &violation.field)?; + f.write_str(": ")?; + } + match ( + violation.message.as_deref().unwrap_or_default(), + violation.rule_id.as_deref().unwrap_or_default(), + ) { + ("", "") => f.write_str("[unknown]"), + ("", rule_id) => write!(f, "[{rule_id}]"), + (message, _) => f.write_str(message), + } +} + +fn write_field_path(f: &mut fmt::Formatter<'_>, path: &FieldPath) -> fmt::Result { + for (i, element) in path.elements.iter().enumerate() { + let name = element.field_name.as_deref().unwrap_or_default(); + // Extension names are already bracketed, as in `[pkg.ext]`. + if i > 0 && !name.starts_with('[') { + f.write_str(".")?; + } + f.write_str(name)?; + match &element.subscript { + None => {} + Some(Subscript::Index(index) | Subscript::UintKey(index)) => write!(f, "[{index}]")?, + Some(Subscript::IntKey(key)) => write!(f, "[{key}]")?, + Some(Subscript::BoolKey(key)) => write!(f, "[{key}]")?, + Some(Subscript::StringKey(key)) => { + f.write_str("[\"")?; + for c in key.chars() { + match c { + '\\' => f.write_str("\\\\")?, + '"' => f.write_str("\\\"")?, + '\r' => f.write_str("\\r")?, + '\n' => f.write_str("\\n")?, + c => f.write_char(c)?, + } + } + f.write_str("\"]")?; + } + } + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use buffa::MessageField; + + use super::{FieldPath, Subscript, ValidationError, ViolationPb}; + use crate::validate::FieldPathElement; + + fn element(name: &str, subscript: Option) -> FieldPathElement { + FieldPathElement { + field_name: Some(name.to_owned()), + subscript, + ..Default::default() + } + } + + fn violation(field: Vec, message: &str, rule_id: &str) -> ViolationPb { + ViolationPb { + field: if field.is_empty() { + MessageField::none() + } else { + MessageField::some(FieldPath { + elements: field, + ..Default::default() + }) + }, + message: Some(message.to_owned()), + rule_id: Some(rule_id.to_owned()), + ..Default::default() + } + } + + fn display(violations: Vec) -> String { + ValidationError::new(violations).to_string() + } + + #[test] + fn displays_field_path() { + let field = vec![ + element("a", None), + element("b", Some(Subscript::Index(0))), + element("[pkg.ext]", None), + element("c", Some(Subscript::BoolKey(true))), + element("d", Some(Subscript::IntKey(-1))), + element("e", Some(Subscript::UintKey(2))), + element("f", Some(Subscript::StringKey("x\\\"\r\n".to_owned()))), + ]; + assert_eq!( + display(vec![violation(field, "must be set", "required")]), + r#"a.b[0][pkg.ext].c[true].d[-1].e[2].f["x\\\"\r\n"]: must be set"#, + ); + } + + #[test] + fn falls_back_to_rule_id() { + assert_eq!( + display(vec![violation(vec![], "must be set", "required")]), + "must be set", + ); + assert_eq!( + display(vec![violation(vec![element("a", None)], "", "custom")]), + "a: [custom]", + ); + assert_eq!(display(vec![violation(vec![], "", "custom")]), "[custom]"); + assert_eq!( + display(vec![violation(vec![element("a", None)], "", "")]), + "a: [unknown]", + ); + assert_eq!(display(vec![violation(vec![], "", "")]), "[unknown]"); + } + + #[test] + fn counts_other_violations() { + let first = || violation(vec![element("a", None)], "bad", "rule"); + assert_eq!( + display(vec![first(), violation(vec![], "", "other")]), + "a: bad, and 1 more violation", + ); + assert_eq!( + display(vec![ + first(), + violation(vec![], "", "other"), + violation(vec![], "", "other"), + ]), + "a: bad, and 2 more violations", + ); } } diff --git a/crates/protovalidate/src/lib.rs b/crates/protovalidate/src/lib.rs index 688bccd5..8964219f 100644 --- a/crates/protovalidate/src/lib.rs +++ b/crates/protovalidate/src/lib.rs @@ -43,8 +43,8 @@ //! Ok(()) => { /* valid */ } //! Err(Error::Validation(e)) => { //! // The message failed to evaluate against rules in the schema. -//! // e.violations() is a serialized buf.validate.Violations describing -//! // each failure. +//! // e.encode_violations() is a serialized buf.validate.Violations +//! // describing each failure. //! } //! // Validation itself failed: unknown type, rules that do not compile //! // or evaluate. diff --git a/crates/protovalidate/src/validator.rs b/crates/protovalidate/src/validator.rs index c2d048b4..1de3c9a1 100644 --- a/crates/protovalidate/src/validator.rs +++ b/crates/protovalidate/src/validator.rs @@ -25,18 +25,9 @@ use crate::protobuf::{Reader, Runtime}; use crate::rules::ValidatorCache; use crate::rules::build::{self, Builder}; use crate::rules::eval::Walk; -use crate::validate::{Violation as ViolationPb, Violations}; +use crate::validate::Violation as ViolationPb; use crate::{DescriptorError, Error, ValidationError}; -/// The violations as a serialized `buf.validate.Violations`. -fn encode_violations(violations: Vec) -> Vec { - Violations { - violations, - ..Default::default() - } - .encode_to_vec() -} - /// Appended to a panic that only a broken build can reach. const BUILD_BUG: &str = "this is a bug in protovalidate, please report it"; @@ -152,9 +143,7 @@ impl Validator { if violations.is_empty() { return Ok(()); } - Err(Error::Validation(ValidationError::new(encode_violations( - violations, - )))) + Err(Error::Validation(ValidationError::new(violations))) } fn run( diff --git a/src/lib.rs b/src/lib.rs index 37b52a19..b01ebf1a 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -27,7 +27,7 @@ use pyo3::exceptions::{PyException, PyValueError}; use pyo3::import_exception; use pyo3::prelude::*; use pyo3::sync::{PyOnceLock, RwLockExt}; -use pyo3::types::{PyBytes, PyList, PyString}; +use pyo3::types::{PyList, PyString}; use protovalidate::{DescriptorError, Error}; @@ -145,17 +145,13 @@ impl Validator { message: PbMessage<'_, '_>, fail_fast: bool, ) -> PyResult<()> { - let (adapter, violations) = self.collect(py, message, fail_fast)?; - if violations.0.is_empty() { + let Some(invalid) = self.check(py, message, fail_fast)? else { return Ok(()); - } - let name = adapter - .descriptor(py) - .getattr(&self.constants.name)? - .cast_into::()?; + }; + let violations = self.violations(py, message, &invalid)?; Err(ValidationError::new_err(( - format!("invalid {}", name.to_str()?), - violations.0.unbind(), + invalid.error.to_string(), + violations.unbind(), ))) } @@ -185,45 +181,60 @@ impl Validator { message: PbMessage<'_, 'py>, fail_fast: bool, ) -> PyResult> { - Ok(self.collect(py, message, fail_fast)?.1) + Ok(ViolationList(match self.check(py, message, fail_fast)? { + None => PyList::empty(py), + Some(invalid) => self.violations(py, message, &invalid)?, + })) } } +/// A message with validation errors. +struct Invalid { + adapter: ProtoAdapter, + error: protovalidate::ValidationError, +} + impl Validator { - /// Resolves a message and collects its violations. - fn collect<'py>( + /// Resolves a message and validates it, returning `None` when it is + /// valid. + fn check( &self, - py: Python<'py>, - message: PbMessage<'_, 'py>, + py: Python<'_>, + message: PbMessage<'_, '_>, fail_fast: bool, - ) -> PyResult<(ProtoAdapter, ViolationList<'py>)> { + ) -> PyResult> { let adapter = ProtoAdapter::resolve(&message.0, self.constants)?; let engine = self.engine(py, adapter.runtime)?; let file = adapter.descriptor(py).getattr(&self.constants.file)?; engine.register(py, adapter.runtime, &file, self.constants)?; let type_name = adapter.type_name(py, self.constants)?; - let Some(serialized) = self.evaluate( + let error = self.evaluate( py, engine, &adapter, type_name.to_str()?, &message.0, fail_fast, - )? - else { - return Ok((adapter, ViolationList(PyList::empty(py)))); - }; - let violations = violation::build_violations( + )?; + Ok(error.map(|error| Invalid { adapter, error })) + } + + /// Converts the violation errors to Python violations. + fn violations<'py>( + &self, + py: Python<'py>, + message: PbMessage<'_, 'py>, + invalid: &Invalid, + ) -> PyResult> { + violation::build_violations( py, - &serialized, + &invalid.error.encode_violations(), &message.0, - &adapter, + &invalid.adapter, self.constants, &self.imports, ) - .map(ViolationList)?; - Ok((adapter, violations)) } /// Returns the engine for `runtime`. @@ -243,8 +254,8 @@ impl Validator { }) } - /// Validates the message in place, returning serialized violations, or - /// `None` when the message is valid. + /// Validates the message in place, returning its violations, or `None` + /// when the message is valid. /// /// Reading the message calls into Python, so the interpreter stays /// attached throughout; the core lock is taken with the @@ -258,7 +269,7 @@ impl Validator { type_name: &str, message: &Bound<'py, PyAny>, fail_fast: bool, - ) -> PyResult>> { + ) -> PyResult> { let core = engine .core .read_py_attached(py) @@ -271,7 +282,7 @@ impl Validator { let ctx = Ctx::new(py, adapter.runtime, self.constants, message, source); match core.validate_message(type_name, &ctx, fail_fast) { Ok(()) => Ok(None), - Err(Error::Validation(error)) => Ok(Some(PyBytes::new(py, error.violations()))), + Err(Error::Validation(error)) => Ok(Some(error)), Err(error) => Err(to_py_err(error)), } } diff --git a/src/violation.rs b/src/violation.rs index 375a35b6..69d09213 100644 --- a/src/violation.rs +++ b/src/violation.rs @@ -21,7 +21,7 @@ use std::sync::Arc; use pyo3::prelude::*; use pyo3::sync::PyOnceLock; -use pyo3::types::{PyBytes, PyDict, PyList, PyString}; +use pyo3::types::{PyDict, PyList, PyString}; use crate::constants::{Constants, Imports}; use crate::hints::{PbFieldPath, ViolationProto}; @@ -359,7 +359,7 @@ fn rules_of<'py>( /// `ValidationError` or reading `rule_id` never pays for path walking. pub fn build_violations<'py>( py: Python<'py>, - serialized: &Bound<'py, PyBytes>, + serialized: &[u8], message: &Bound<'py, PyAny>, adapter: &ProtoAdapter, constants: &Constants, diff --git a/test/test_validate.py b/test/test_validate.py index 9cec26fd..0ef8efe5 100644 --- a/test/test_validate.py +++ b/test/test_validate.py @@ -51,7 +51,7 @@ def test_ninf(validator: ValidatorProtocol) -> None: rule_value=True, ) - check_invalid(validator, msg, [expected_violation]) + check_invalid(validator, msg, [expected_violation], "val: must be finite") @pytest.mark.parametrize("validator", validators) @@ -67,7 +67,7 @@ def test_map_key(validator: ValidatorProtocol) -> None: rule_value=0, ) - check_invalid(validator, msg, [expected_violation]) + check_invalid(validator, msg, [expected_violation], "val[1]: must be less than 0") @pytest.mark.parametrize("validator", validators) @@ -102,7 +102,7 @@ def test_protovalidate_oneof_violation(validator: ValidatorProtocol) -> None: message="only one of a, b can be set", rule_id="message.oneof" ) - check_invalid(validator, msg, [expected_violation]) + check_invalid(validator, msg, [expected_violation], "only one of a, b can be set") @pytest.mark.parametrize("validator", validators) @@ -113,7 +113,7 @@ def test_protovalidate_oneof_required_violation(validator: ValidatorProtocol) -> message="one of a, b must be set", rule_id="message.oneof" ) - check_invalid(validator, msg, [expected_violation]) + check_invalid(validator, msg, [expected_violation], "one of a, b must be set") @pytest.mark.parametrize("validator", validators) @@ -158,7 +158,9 @@ def test_maps(validator: ValidatorProtocol) -> None: rule_value=2, ) - check_invalid(validator, msg, [expected_violation]) + check_invalid( + validator, msg, [expected_violation], "val: map must be at least 2 entries" + ) @pytest.mark.parametrize("validator", validators) @@ -189,7 +191,12 @@ def test_multiple_validations(validator: ValidatorProtocol) -> None: rule_value=5, ) - check_invalid(validator, msg, [expected_violation1, expected_violation2]) + check_invalid( + validator, + msg, + [expected_violation1, expected_violation2], + "title: does not have prefix `foo`, and 1 more violation", + ) @pytest.mark.parametrize("validator", validators) @@ -217,7 +224,7 @@ def test_fail_fast(validator: ValidatorProtocol) -> None: with pytest.raises(protovalidate.ValidationError) as cm: validator.validate(msg, fail_fast=True) e = cm.value - assert str(e) == f"invalid {type(msg).desc().name}" + assert str(e) == "title: does not have prefix `foo`" compare_violations(e.violations, [expected_violation]) # ty: ignore # Test collect_violations @@ -226,16 +233,16 @@ def test_fail_fast(validator: ValidatorProtocol) -> None: def check_invalid( - validator: ValidatorProtocol, msg: protobuf.Message, expected: list[Violation] + validator: ValidatorProtocol, + msg: protobuf.Message, + expected: list[Violation], + expected_message: str, ) -> None: # Test validate with pytest.raises(protovalidate.ValidationError) as exc_info: validator.validate(msg) e = exc_info.value - if isinstance(msg, protobuf.Message): - assert str(e) == f"invalid {type(msg).desc().name}" - else: - assert str(e) == f"invalid {msg.DESCRIPTOR.name}" + assert str(e) == expected_message compare_violations(e.violations, expected) # ty: ignore # Test collect_violations diff --git a/test/test_validate_legacy.py b/test/test_validate_legacy.py index d91c9c97..a8fc240d 100644 --- a/test/test_validate_legacy.py +++ b/test/test_validate_legacy.py @@ -18,9 +18,6 @@ pytest.importorskip("google.protobuf", reason="optional dependency not installed") - -from typing import TYPE_CHECKING - import protovalidate from protovalidate import Violation @@ -28,9 +25,6 @@ from .conftest import make_validator from .gen.tests.example.v1 import validations_pb2 -if TYPE_CHECKING: - from google.protobuf import message as google_message - validators: list[ValidatorProtocol] = [ protovalidate, # global module singleton make_validator(), @@ -61,7 +55,7 @@ def test_legacy_message_invalid(validator: ValidatorProtocol) -> None: with pytest.raises(protovalidate.ValidationError) as exc_info: validator.validate(msg) e = exc_info.value - assert str(e) == f"invalid {msg.DESCRIPTOR.name}" + assert str(e) == "val: must be finite" compare_violations(e.violations, [expected_violation]) # ty: ignore violations = validator.collect_violations(msg) @@ -83,18 +77,3 @@ def test_legacy_message_map_key(validator: ValidatorProtocol) -> None: violations = validator.collect_violations(msg) compare_violations(violations, [expected_violation]) - - -def check_invalid( - validator: ValidatorProtocol, msg: google_message.Message, expected: list[Violation] -) -> None: - # Test validate - with pytest.raises(protovalidate.ValidationError) as exc_info: - validator.validate(msg) - e = exc_info.value - assert str(e) == f"invalid {msg.DESCRIPTOR.name}" - compare_violations(e.violations, expected) # ty: ignore - - # Test collect_violations - violations = validator.collect_violations(msg) - compare_violations(violations, expected)