#![allow(dead_code)]
use crate::Result;
use crate::error::CoreError;
use apache_avro::schema::Schema;
fn actual_schema_from_union(s: &Schema) -> Result<&Schema> {
let Schema::Union(u) = s else {
return Ok(s);
};
let variants = u.variants();
match variants.len() {
2 if matches!(variants[0], Schema::Null) => Ok(&variants[1]),
2 if matches!(variants[1], Schema::Null) => Ok(&variants[0]),
1 => Ok(&variants[0]),
_ => Err(CoreError::Schema(format!("Union is malformed: {s:?}"))),
}
}
fn schemas_equal(a: &Schema, b: &Schema) -> bool {
match (a, b) {
(Schema::Record(ra), Schema::Record(rb)) => {
ra.fields.len() == rb.fields.len()
&& ra
.fields
.iter()
.zip(rb.fields.iter())
.all(|(fa, fb)| fa.name == fb.name && schemas_equal(&fa.schema, &fb.schema))
}
(Schema::Array(x), Schema::Array(y)) => schemas_equal(&x.items, &y.items),
(Schema::Map(x), Schema::Map(y)) => schemas_equal(&x.types, &y.types),
(Schema::Union(x), Schema::Union(y)) => {
x.variants().len() == y.variants().len()
&& x.variants()
.iter()
.zip(y.variants().iter())
.all(|(s, t)| schemas_equal(s, t))
}
_ => a == b,
}
}
fn is_logical_type(s: &Schema) -> bool {
matches!(
s,
Schema::Decimal(_)
| Schema::BigDecimal
| Schema::Uuid
| Schema::Date
| Schema::TimeMillis
| Schema::TimeMicros
| Schema::TimestampMillis
| Schema::TimestampMicros
| Schema::TimestampNanos
| Schema::LocalTimestampMillis
| Schema::LocalTimestampMicros
| Schema::LocalTimestampNanos
| Schema::Duration
)
}
fn physical_type(s: &Schema) -> PhysicalType {
match s {
Schema::Null => PhysicalType::Null,
Schema::Boolean => PhysicalType::Boolean,
Schema::Int | Schema::Date | Schema::TimeMillis => PhysicalType::Int,
Schema::Long
| Schema::TimeMicros
| Schema::TimestampMillis
| Schema::TimestampMicros
| Schema::TimestampNanos
| Schema::LocalTimestampMillis
| Schema::LocalTimestampMicros
| Schema::LocalTimestampNanos => PhysicalType::Long,
Schema::Float => PhysicalType::Float,
Schema::Double => PhysicalType::Double,
Schema::Decimal(d) => physical_type(&d.inner),
Schema::Bytes | Schema::BigDecimal => PhysicalType::Bytes,
Schema::String | Schema::Uuid => PhysicalType::String,
Schema::Fixed(_) | Schema::Duration => PhysicalType::Fixed,
Schema::Record(_) => PhysicalType::Record,
Schema::Enum(_) => PhysicalType::Enum,
Schema::Array(_) => PhysicalType::Array,
Schema::Map(_) => PhysicalType::Map,
Schema::Union(_) => PhysicalType::Union,
Schema::Ref { .. } => PhysicalType::Ref,
}
}
#[derive(PartialEq)]
enum PhysicalType {
Null,
Boolean,
Int,
Long,
Float,
Double,
Bytes,
String,
Fixed,
Record,
Enum,
Array,
Map,
Union,
Ref,
}
fn needs_rewrite_to_string(writer: &Schema, reader_is_enum: bool) -> bool {
if is_logical_type(writer) {
return true;
}
if let Schema::Enum(_) = writer {
return !reader_is_enum;
}
true
}
pub fn record_needs_rewrite_for_extended_promotion(
writer: &Schema,
reader: &Schema,
) -> Result<bool> {
if schemas_equal(writer, reader) {
return Ok(false);
}
if is_logical_type(reader) {
return Ok(!logical_types_equal(writer, reader));
}
match reader {
Schema::Record(rrec) => match writer {
Schema::Record(wrec) => {
if rrec.fields.len() > wrec.fields.len() {
return Ok(true);
}
for rf in &rrec.fields {
match wrec.fields.iter().find(|wf| wf.name == rf.name) {
None => return Ok(true),
Some(wf) => {
if record_needs_rewrite_for_extended_promotion(&wf.schema, &rf.schema)?
{
return Ok(true);
}
}
}
}
Ok(false)
}
_ => Err(CoreError::Schema(format!(
"Not a record: {writer:?} (reader expects a record)"
))),
},
Schema::Array(relem) => match writer {
Schema::Array(welem) => {
record_needs_rewrite_for_extended_promotion(&welem.items, &relem.items)
}
_ => Ok(false),
},
Schema::Map(rval) => match writer {
Schema::Map(wval) => {
record_needs_rewrite_for_extended_promotion(&wval.types, &rval.types)
}
_ => Ok(false),
},
Schema::Union(_) => {
let w = actual_schema_from_union(writer)?;
let r = actual_schema_from_union(reader)?;
record_needs_rewrite_for_extended_promotion(w, r)
}
Schema::Enum(_) => Ok(needs_rewrite_to_string(writer, true)),
Schema::String => Ok(needs_rewrite_to_string(writer, false)),
Schema::Double | Schema::Float | Schema::Long => Ok(!matches!(
physical_type(writer),
PhysicalType::Int | PhysicalType::Long
)),
_ => Ok(physical_type(writer) != physical_type(reader)),
}
}
fn logical_types_equal(writer: &Schema, reader: &Schema) -> bool {
if !is_logical_type(writer) {
return false;
}
schemas_equal(writer, reader)
}
#[cfg(test)]
mod tests {
use super::*;
use apache_avro::Schema as AvroSchema;
fn rec(fields: &str) -> AvroSchema {
AvroSchema::parse_str(&format!(
r#"{{"type":"record","name":"r","fields":[{fields}]}}"#
))
.unwrap()
}
#[test]
fn test_detector_matches_gold_matrix() {
let cases = vec![
(
r#"{"name":"a","type":"int"}"#,
r#"{"name":"a","type":"int"}"#,
false,
),
(
r#"{"name":"a","type":"int"}"#,
r#"{"name":"a","type":"long"}"#,
false,
),
(
r#"{"name":"a","type":"int"}"#,
r#"{"name":"a","type":"float"}"#,
false,
),
(
r#"{"name":"a","type":"int"}"#,
r#"{"name":"a","type":"double"}"#,
false,
),
(
r#"{"name":"a","type":"long"}"#,
r#"{"name":"a","type":"double"}"#,
false,
),
(
r#"{"name":"a","type":"float"}"#,
r#"{"name":"a","type":"double"}"#,
true,
),
(
r#"{"name":"a","type":"int"}"#,
r#"{"name":"a","type":"string"}"#,
true,
),
(
r#"{"name":"a","type":"long"}"#,
r#"{"name":"a","type":"string"}"#,
true,
),
(
r#"{"name":"a","type":"float"}"#,
r#"{"name":"a","type":"string"}"#,
true,
),
(
r#"{"name":"a","type":"bytes"}"#,
r#"{"name":"a","type":"string"}"#,
true,
),
(
r#"{"name":"a","type":"string"}"#,
r#"{"name":"a","type":"bytes"}"#,
true,
),
(
r#"{"name":"a","type":"int"}"#,
r#"{"name":"a","type":"int"},{"name":"b","type":["null","string"],"default":null}"#,
true,
),
(
r#"{"name":"a","type":"int"},{"name":"b","type":"string"}"#,
r#"{"name":"a","type":"int"}"#,
false,
),
(
r#"{"name":"a","type":["null","int"],"default":null}"#,
r#"{"name":"a","type":["null","long"],"default":null}"#,
false,
),
(
r#"{"name":"a","type":{"type":"array","items":"int"}}"#,
r#"{"name":"a","type":{"type":"array","items":"string"}}"#,
true,
),
(
r#"{"name":"a","type":{"type":"map","values":"int"}}"#,
r#"{"name":"a","type":{"type":"map","values":"long"}}"#,
false,
),
(
r#"{"name":"a","type":{"type":"int","logicalType":"time-millis"}}"#,
r#"{"name":"a","type":{"type":"int","logicalType":"date"}}"#,
true,
),
(
r#"{"name":"a","type":{"type":"long","logicalType":"timestamp-millis"}}"#,
r#"{"name":"a","type":{"type":"long","logicalType":"timestamp-micros"}}"#,
true,
),
(
r#"{"name":"a","type":{"type":"long","logicalType":"timestamp-micros"}}"#,
r#"{"name":"a","type":{"type":"long","logicalType":"timestamp-micros"}}"#,
false,
),
(
r#"{"name":"a","type":{"type":"long","logicalType":"timestamp-micros"}}"#,
r#"{"name":"a","type":"long"}"#,
false,
),
(
r#"{"name":"a","type":"long"}"#,
r#"{"name":"a","type":{"type":"long","logicalType":"timestamp-micros"}}"#,
true,
),
(
r#"{"name":"s","type":{"type":"record","name":"s","fields":[{"name":"x","type":"int"}]}}"#,
r#"{"name":"s","type":{"type":"record","name":"s","fields":[{"name":"x","type":"long"}]}}"#,
false,
),
(
r#"{"name":"s","type":{"type":"record","name":"s","fields":[{"name":"x","type":"int"}]}}"#,
r#"{"name":"s","type":{"type":"record","name":"s","fields":[{"name":"x","type":"int"},{"name":"y","type":["null","string"],"default":null}]}}"#,
true,
),
(
r#"{"name":"a","type":{"type":"int","logicalType":"date"}}"#,
r#"{"name":"a","type":"int"}"#,
false,
),
(
r#"{"name":"a","type":{"type":"int","logicalType":"time-millis"}}"#,
r#"{"name":"a","type":"int"}"#,
false,
),
(
r#"{"name":"a","type":{"type":"bytes","logicalType":"decimal","precision":10,"scale":2}}"#,
r#"{"name":"a","type":"bytes"}"#,
false,
),
(
r#"{"name":"a","type":"int"}"#,
r#"{"name":"a","type":"bytes"}"#,
true,
),
(
r#"{"name":"a","type":"int"}"#,
r#"{"name":"b","type":"int"}"#,
true,
),
(
r#"{"name":"s","type":{"type":"record","name":"s","fields":[{"name":"x","type":"int"}]}}"#,
r#"{"name":"s","type":{"type":"record","name":"s","fields":[{"name":"y","type":"int"}]}}"#,
true,
),
];
for (w, r, expect) in cases {
let got = record_needs_rewrite_for_extended_promotion(&rec(w), &rec(r)).unwrap();
assert_eq!(got, expect, "writer=[{w}] reader=[{r}]");
}
}
#[test]
fn test_reader_record_non_record_writer_errors_like_gold() {
let w = rec(
r#"{"name":"s","type":["null",{"type":"record","name":"inner","fields":[{"name":"x","type":"int"}]}],"default":null}"#,
);
let r = rec(
r#"{"name":"s","type":{"type":"record","name":"inner","fields":[{"name":"x","type":"int"}]}}"#,
);
assert!(record_needs_rewrite_for_extended_promotion(&w, &r).is_err());
let w = rec(r#"{"name":"s","type":"string"}"#);
let r = rec(
r#"{"name":"s","type":{"type":"record","name":"inner","fields":[{"name":"x","type":"int"}]}}"#,
);
assert!(record_needs_rewrite_for_extended_promotion(&w, &r).is_err());
}
#[test]
fn test_writer_union_vs_plain_numeric_matches_gold() {
let w = rec(r#"{"name":"a","type":["null","int"],"default":null}"#);
let r = rec(r#"{"name":"a","type":"long"}"#);
assert!(record_needs_rewrite_for_extended_promotion(&w, &r).unwrap());
let w = rec(r#"{"name":"a","type":"long"}"#);
let r = rec(r#"{"name":"a","type":["null","int"],"default":null}"#);
assert!(record_needs_rewrite_for_extended_promotion(&w, &r).unwrap());
}
#[test]
fn test_malformed_union_errors_like_gold() {
let w = rec(r#"{"name":"a","type":["null","int","string"],"default":null}"#);
let r = rec(r#"{"name":"a","type":["null","string","int"]}"#);
assert!(record_needs_rewrite_for_extended_promotion(&w, &r).is_err());
let w = rec(r#"{"name":"a","type":["int","string"]}"#);
let r = rec(r#"{"name":"a","type":["string","int"]}"#);
assert!(record_needs_rewrite_for_extended_promotion(&w, &r).is_err());
let w = rec(
r#"{"name":"ok","type":["null","int"],"default":null},{"name":"bad","type":["null","int","long"],"default":null}"#,
);
let r = rec(
r#"{"name":"ok","type":["null","long"],"default":null},{"name":"bad","type":["null","long","int"]}"#,
);
assert!(record_needs_rewrite_for_extended_promotion(&w, &r).is_err());
}
#[test]
fn test_recursive_schema_terminates_and_matches_gold() {
let w = AvroSchema::parse_str(
r#"{"type":"record","name":"Node","fields":[
{"name":"value","type":"int"},
{"name":"next","type":["null","Node"],"default":null}]}"#,
)
.unwrap();
let r = AvroSchema::parse_str(
r#"{"type":"record","name":"Node","fields":[
{"name":"value","type":"long"},
{"name":"next","type":["null","Node"],"default":null}]}"#,
)
.unwrap();
assert!(!record_needs_rewrite_for_extended_promotion(&w, &r).unwrap());
}
}