This is an automated email from the ASF dual-hosted git repository. Kriskras99 pushed a commit to branch fix/precision_must_be_non_zero in repository https://gitbox.apache.org/repos/asf/avro-rs.git
commit 24818400dc15084bd2d8dd5257dc13be3d75799b Author: Kriskras99 <[email protected]> AuthorDate: Wed Sep 9 14:18:05 2026 +0200 fix: Prevent users from setting a non-zero precision --- avro/src/decode.rs | 5 ++- avro/src/error.rs | 2 +- avro/src/schema/mod.rs | 18 ++++---- avro/src/schema/parser.rs | 91 +++++++++++++++++----------------------- avro/src/schema_compatibility.rs | 12 +++--- avro/src/schema_equality.rs | 5 ++- avro/src/types.rs | 16 +++---- avro/src/writer/datum.rs | 5 ++- 8 files changed, 72 insertions(+), 82 deletions(-) diff --git a/avro/src/decode.rs b/avro/src/decode.rs index 00b8b67..9c06eb9 100644 --- a/avro/src/decode.rs +++ b/avro/src/decode.rs @@ -471,6 +471,7 @@ mod tests { use apache_avro_test_helper::TestResult; use pretty_assertions::assert_eq; use std::collections::HashMap; + use std::num::NonZero; use uuid::Uuid; #[test] @@ -692,7 +693,7 @@ mod tests { .size(2) .build(), ), - precision: 4, + precision: NonZero::new(4).unwrap(), scale: 2, }); let bigint = (-423).to_bigint().unwrap(); @@ -720,7 +721,7 @@ mod tests { doc: None, attributes: Default::default(), }), - precision: 4, + precision: NonZero::new(4).unwrap(), scale: 2, }); let value = Value::Decimal(Decimal::from( diff --git a/avro/src/error.rs b/avro/src/error.rs index 58f9145..8168694 100644 --- a/avro/src/error.rs +++ b/avro/src/error.rs @@ -405,7 +405,7 @@ pub enum Details { #[error("The decimal precision ({precision}) must be bigger or equal to the scale ({scale})")] DecimalPrecisionLessThanScale { precision: usize, scale: usize }, - #[error("The decimal precision ({precision}) must be a positive number")] + #[error("The decimal precision ({precision}) must be a non-zero positive number")] DecimalPrecisionMuBePositive { precision: usize }, #[deprecated(since = "0.20.0", note = "This error variant is not generated anymore")] diff --git a/avro/src/schema/mod.rs b/avro/src/schema/mod.rs index b503f12..9c0de17 100644 --- a/avro/src/schema/mod.rs +++ b/avro/src/schema/mod.rs @@ -47,6 +47,7 @@ use serde::{ use serde_json::{Map, Value as JsonValue}; use std::borrow::Cow; use std::fmt::Formatter; +use std::num::NonZero; use std::{ collections::{BTreeMap, HashMap, HashSet}, fmt, @@ -454,9 +455,8 @@ pub enum UuidSchema { Fixed(FixedSchema), } -type DecimalMetadata = usize; -pub(crate) type Precision = DecimalMetadata; -pub(crate) type Scale = DecimalMetadata; +pub(crate) type Precision = NonZero<usize>; +pub(crate) type Scale = usize; impl Schema { /// Converts `self` into its [Parsing Canonical Form]. @@ -4347,14 +4347,14 @@ mod tests { "scale": 2 }); let parse_result = Schema::parse(schema)?; - assert!(matches!( + assert_eq!( parse_result, Schema::Decimal(DecimalSchema { - precision: 9, + precision: NonZero::new(9).unwrap(), scale: 2, - .. + inner: InnerDecimalSchema::Bytes }) - )); + ); // long decimal, represents as native complex type. let schema = json!( @@ -4579,7 +4579,7 @@ mod tests { #[test] fn test_avro_3925_serialize_decimal_inner_fixed() -> TestResult { let schema = Schema::Decimal(DecimalSchema { - precision: 36, + precision: NonZero::new(36).unwrap(), scale: 10, inner: InnerDecimalSchema::Fixed(FixedSchema { name: Name::new("decimal_36_10").unwrap(), @@ -4609,7 +4609,7 @@ mod tests { #[test] fn test_avro_3925_serialize_decimal_inner_bytes() -> TestResult { let schema = Schema::Decimal(DecimalSchema { - precision: 36, + precision: NonZero::new(36).unwrap(), scale: 10, inner: InnerDecimalSchema::Bytes, }); diff --git a/avro/src/schema/parser.rs b/avro/src/schema/parser.rs index 37ef726..4bf78cc 100644 --- a/avro/src/schema/parser.rs +++ b/avro/src/schema/parser.rs @@ -17,9 +17,9 @@ use crate::error::Details; use crate::schema::{ - Alias, Aliases, ArraySchema, DecimalMetadata, DecimalSchema, EnumSchema, FixedSchema, - MapSchema, Name, Names, NamespaceRef, Precision, RecordField, RecordSchema, Scale, Schema, - SchemaKind, UnionSchema, UuidSchema, + Alias, Aliases, ArraySchema, DecimalSchema, EnumSchema, FixedSchema, MapSchema, Name, Names, + NamespaceRef, Precision, RecordField, RecordSchema, Scale, Schema, SchemaKind, UnionSchema, + UuidSchema, }; use crate::util::{JsonValueDescriber, MapHelper}; use crate::validator::validate_enum_symbol_name; @@ -27,6 +27,7 @@ use crate::{AvroResult, Error}; use log::{debug, error, warn}; use serde_json::{Map, Value}; use std::collections::{BTreeMap, HashMap, HashSet}; +use std::num::NonZero; #[derive(Default)] pub(crate) struct Parser { @@ -190,41 +191,46 @@ impl Parser { Ok(Schema::Ref { name: full_name }) } - fn get_decimal_integer( + fn parse_precision_and_scale( &self, complex: &Map<String, Value>, - key: &'static str, - ) -> AvroResult<DecimalMetadata> { - match complex.get(key) { - Some(Value::Number(value)) => self.parse_json_integer_for_decimal(value), - None => { - if key == "scale" { - Ok(0) - } else { - Err(Details::GetDecimalMetadataFromJson(key).into()) + ) -> AvroResult<(Precision, Scale)> { + let precision = match complex.get("precision") { + Some(Value::Number(value)) if let Some(value) = value.as_u64() => { + let value = + usize::try_from(value).map_err(|e| Details::ConvertU64ToUsize(e, value))?; + NonZero::new(value) + .ok_or(Details::DecimalPrecisionMuBePositive { precision: value })? + } + Some(value) => { + return Err(Details::GetDecimalMetadataValueFromJson { + key: "precision".into(), + value: value.clone(), } + .into()); } - Some(value) => Err(Details::GetDecimalMetadataValueFromJson { - key: key.into(), - value: value.clone(), + None => return Err(Details::GetDecimalMetadataFromJson("precision").into()), + }; + let scale = match complex.get("scale") { + Some(Value::Number(value)) if let Some(value) = value.as_u64() => { + usize::try_from(value).map_err(|e| Details::ConvertU64ToUsize(e, value))? } - .into()), - } - } - - fn parse_precision_and_scale( - &self, - complex: &Map<String, Value>, - ) -> AvroResult<(Precision, Scale)> { - let precision = self.get_decimal_integer(complex, "precision")?; - let scale = self.get_decimal_integer(complex, "scale")?; - - if precision < 1 { - return Err(Details::DecimalPrecisionMuBePositive { precision }.into()); - } + Some(value) => { + return Err(Details::GetDecimalMetadataValueFromJson { + key: "scale".into(), + value: value.clone(), + } + .into()); + } + None => 0, + }; - if precision < scale { - Err(Details::DecimalPrecisionLessThanScale { precision, scale }.into()) + if precision.get() < scale { + Err(Details::DecimalPrecisionLessThanScale { + precision: precision.get(), + scale, + } + .into()) } else { Ok((precision, scale)) } @@ -798,25 +804,4 @@ impl Parser { _ => Ok(name), } } - - fn parse_json_integer_for_decimal( - &self, - value: &serde_json::Number, - ) -> AvroResult<DecimalMetadata> { - Ok(if value.is_u64() { - let num = value - .as_u64() - .ok_or_else(|| Details::GetU64FromJson(value.clone()))?; - num.try_into() - .map_err(|e| Details::ConvertU64ToUsize(e, num))? - } else if value.is_i64() { - let num = value - .as_i64() - .ok_or_else(|| Details::GetI64FromJson(value.clone()))?; - num.try_into() - .map_err(|e| Details::ConvertI64ToUsize(e, num))? - } else { - return Err(Details::GetPrecisionOrScaleFromJson(value.clone()).into()); - }) - } } diff --git a/avro/src/schema_compatibility.rs b/avro/src/schema_compatibility.rs index dcbc1dd..ab86921 100644 --- a/avro/src/schema_compatibility.rs +++ b/avro/src/schema_compatibility.rs @@ -320,9 +320,9 @@ impl Checker { Ok(Compatibility::Full) } else { Err(CompatibilityError::DecimalMismatch { - r_precision: *r_precision, + r_precision: r_precision.get(), r_scale: *r_scale, - w_precision: *w_precision, + w_precision: w_precision.get(), w_scale: *w_scale }) } @@ -444,8 +444,6 @@ impl Checker { #[cfg(test)] mod tests { - use std::collections::BTreeMap; - use super::*; use crate::{ Codec, Decimal, Reader, Writer, @@ -454,6 +452,8 @@ mod tests { }; use apache_avro_test_helper::TestResult; use rstest::*; + use std::collections::BTreeMap; + use std::num::NonZero; fn int_array_schema() -> Schema { Schema::parse_str(r#"{"type":"array", "items":"int"}"#).unwrap() @@ -1691,12 +1691,12 @@ mod tests { #[test] fn avro_rs_342_decimal_fixed_and_bytes() -> TestResult { let bytes = Schema::Decimal(DecimalSchema { - precision: 20, + precision: NonZero::new(20).unwrap(), scale: 0, inner: InnerDecimalSchema::Bytes, }); let fixed = Schema::Decimal(DecimalSchema { - precision: 20, + precision: NonZero::new(20).unwrap(), scale: 0, inner: InnerDecimalSchema::Fixed(FixedSchema { name: Name::new("DecimalFixed")?, diff --git a/avro/src/schema_equality.rs b/avro/src/schema_equality.rs index 3e31b96..6e39acc 100644 --- a/avro/src/schema_equality.rs +++ b/avro/src/schema_equality.rs @@ -267,6 +267,7 @@ mod tests { use apache_avro_test_helper::TestResult; use serde_json::Value; use std::collections::BTreeMap; + use std::num::NonZero; const SPECIFICATION_EQ: SpecificationEq = SpecificationEq; const STRUCT_FIELD_EQ: StructFieldEq = StructFieldEq { @@ -505,7 +506,7 @@ mod tests { #[test] fn test_avro_3939_compare_decimal_schemata() { let schema_one = Schema::Decimal(DecimalSchema { - precision: 10, + precision: NonZero::new(10).unwrap(), scale: 2, inner: InnerDecimalSchema::Bytes, }); @@ -513,7 +514,7 @@ mod tests { assert!(!STRUCT_FIELD_EQ.compare(&schema_one, &Schema::Boolean)); let schema_two = Schema::Decimal(DecimalSchema { - precision: 10, + precision: NonZero::new(10).unwrap(), scale: 2, inner: InnerDecimalSchema::Bytes, }); diff --git a/avro/src/types.rs b/avro/src/types.rs index 70759a0..ac977e1 100644 --- a/avro/src/types.rs +++ b/avro/src/types.rs @@ -901,6 +901,7 @@ impl Value { scale: Scale, inner: &InnerDecimalSchema, ) -> Result<Self, Error> { + let precision = precision.get(); if scale > precision { return Err(Details::GetScaleAndPrecision { scale, precision }.into()); } @@ -1376,6 +1377,7 @@ mod tests { use num_bigint::BigInt; use pretty_assertions::assert_eq; use serde_json::json; + use std::num::NonZero; #[test] fn avro_3809_validate_nested_records_with_implicit_namespace() -> TestResult { @@ -1895,7 +1897,7 @@ Field with name '"b"' is not a member of the map items"#, fn resolve_decimal_bytes() -> TestResult { let value = Value::Decimal(Decimal::from(vec![1, 2, 3, 4, 5])); value.clone().resolve(&Schema::Decimal(DecimalSchema { - precision: 10, + precision: NonZero::new(10).unwrap(), scale: 4, inner: InnerDecimalSchema::Bytes, }))?; @@ -1908,7 +1910,7 @@ Field with name '"b"' is not a member of the map items"#, fn avro_rs_580_resolve_decimal_from_string_default() -> TestResult { let value = Value::String("\u{0000}".to_string()); let resolved = value.resolve(&Schema::Decimal(DecimalSchema { - precision: 10, + precision: NonZero::new(10).unwrap(), scale: 4, inner: InnerDecimalSchema::Bytes, }))?; @@ -1919,7 +1921,7 @@ Field with name '"b"' is not a member of the map items"#, all_bytes_str.push(char::from_u32(b as u32).unwrap()); } let resolved = Value::String(all_bytes_str).resolve(&Schema::Decimal(DecimalSchema { - precision: 10, + precision: NonZero::new(10).unwrap(), scale: 0, inner: InnerDecimalSchema::Bytes, }))?; @@ -1932,7 +1934,7 @@ Field with name '"b"' is not a member of the map items"#, assert!( value .resolve(&Schema::Decimal(DecimalSchema { - precision: 10, + precision: NonZero::new(10).unwrap(), scale: 4, inner: InnerDecimalSchema::Bytes, })) @@ -1973,7 +1975,7 @@ Field with name '"b"' is not a member of the map items"#, assert!( value .resolve(&Schema::Decimal(DecimalSchema { - precision: 2, + precision: NonZero::new(2).unwrap(), scale: 3, inner: InnerDecimalSchema::Bytes, })) @@ -1987,7 +1989,7 @@ Field with name '"b"' is not a member of the map items"#, assert!( value .resolve(&Schema::Decimal(DecimalSchema { - precision: 1, + precision: NonZero::new(1).unwrap(), scale: 0, inner: InnerDecimalSchema::Bytes, })) @@ -2002,7 +2004,7 @@ Field with name '"b"' is not a member of the map items"#, value .clone() .resolve(&Schema::Decimal(DecimalSchema { - precision: 10, + precision: NonZero::new(10).unwrap(), scale: 1, inner: InnerDecimalSchema::Fixed(FixedSchema { name: Name::new("decimal").unwrap(), diff --git a/avro/src/writer/datum.rs b/avro/src/writer/datum.rs index 30804ff..b021c80 100644 --- a/avro/src/writer/datum.rs +++ b/avro/src/writer/datum.rs @@ -239,6 +239,7 @@ pub fn to_avro_datum_schemata<T: Into<Value>>( #[cfg(test)] mod tests { use apache_avro_test_helper::TestResult; + use std::num::NonZero; use super::*; use crate::reader::datum::GenericDatumReader; @@ -456,7 +457,7 @@ mod tests { logical_type_test( r#"{"type": {"type": "fixed", "size": 30, "name": "decimal"}, "logicalType": "decimal", "precision": 20, "scale": 5}"#, &Schema::Decimal(DecimalSchema { - precision: 20, + precision: NonZero::new(20).unwrap(), scale: 5, inner, }), @@ -472,7 +473,7 @@ mod tests { logical_type_test( r#"{"type": "bytes", "logicalType": "decimal", "precision": 4, "scale": 3}"#, &Schema::Decimal(DecimalSchema { - precision: 4, + precision: NonZero::new(4).unwrap(), scale: 3, inner: InnerDecimalSchema::Bytes, }),
