use std::{
borrow::Cow,
collections::{HashMap, HashSet},
};
use crate::{
Schema,
schema::{FixedSchema, Name, NamespaceRef, RecordField, RecordSchema, UnionSchema, UuidSchema},
};
pub trait AvroSchema {
fn get_schema() -> Schema;
}
pub trait AvroSchemaComponent {
fn get_schema_in_ctxt(
named_schemas: &mut HashSet<Name>,
enclosing_namespace: NamespaceRef,
) -> Schema;
fn get_record_fields_in_ctxt(
named_schemas: &mut HashSet<Name>,
enclosing_namespace: NamespaceRef,
) -> Option<Vec<RecordField>> {
get_record_fields_in_ctxt(named_schemas, enclosing_namespace, Self::get_schema_in_ctxt)
}
fn field_default() -> Option<serde_json::Value> {
None
}
}
#[doc(hidden)]
pub fn get_record_fields_in_ctxt(
named_schemas: &mut HashSet<Name>,
enclosing_namespace: NamespaceRef,
schema_fn: fn(named_schemas: &mut HashSet<Name>, enclosing_namespace: NamespaceRef) -> Schema,
) -> Option<Vec<RecordField>> {
let mut record = match schema_fn(named_schemas, enclosing_namespace) {
Schema::Record(record) => record,
Schema::Ref { name } => {
assert!(
named_schemas.remove(&name),
"Name '{name}' should exist in `named_schemas` otherwise Ref is invalid: {named_schemas:?}"
);
let schema = schema_fn(named_schemas, enclosing_namespace);
named_schemas.insert(name);
let Schema::Record(record) = schema else {
return None;
};
return Some(record.fields);
}
_ => return None,
};
fn find_first_ref<'a>(schema: &'a mut Schema, target: &Name) -> Option<&'a mut Schema> {
match schema {
Schema::Ref { name } if name == target => Some(schema),
Schema::Array(array) => find_first_ref(&mut array.items, target),
Schema::Map(map) => find_first_ref(&mut map.types, target),
Schema::Union(union) => {
for schema in &mut union.schemas {
if let Some(schema) = find_first_ref(schema, target) {
return Some(schema);
}
}
None
}
Schema::Record(record) => {
assert_ne!(
&record.name, target,
"Only expecting a Ref named {target:?}"
);
for field in &mut record.fields {
if let Some(schema) = find_first_ref(&mut field.schema, target) {
return Some(schema);
}
}
None
}
_ => None,
}
}
let new_fields = record
.fields
.iter()
.map(|field| RecordField {
name: field.name.clone(),
doc: field.doc.clone(),
aliases: field.aliases.clone(),
default: field.default.clone(),
schema: if field.schema.is_named() {
Schema::Ref {
name: field.schema.name().expect("Schema is named").clone(),
}
} else {
field.schema.clone()
},
custom_attributes: field.custom_attributes.clone(),
})
.collect();
named_schemas.remove(&record.name);
for field in &mut record.fields {
if let Some(schema) = find_first_ref(&mut field.schema, &record.name) {
let new_schema = RecordSchema {
name: record.name,
aliases: record.aliases,
doc: record.doc,
fields: new_fields,
lookup: record.lookup,
attributes: record.attributes,
};
let name = match std::mem::replace(schema, Schema::Record(new_schema)) {
Schema::Ref { name } => name,
schema => {
panic!("Only expected `Schema::Ref` from `find_first_ref`, got: {schema:?}")
}
};
named_schemas.insert(name.clone());
break;
}
}
Some(record.fields)
}
impl<T> AvroSchema for T
where
T: AvroSchemaComponent + ?Sized,
{
fn get_schema() -> Schema {
T::get_schema_in_ctxt(&mut HashSet::default(), None)
}
}
macro_rules! impl_schema (
($type:ty, $variant_constructor:expr) => (
impl AvroSchemaComponent for $type {
fn get_schema_in_ctxt(_: &mut HashSet<Name>, _: NamespaceRef) -> Schema {
$variant_constructor
}
fn get_record_fields_in_ctxt(_: &mut HashSet<Name>, _: NamespaceRef) -> Option<Vec<RecordField>> {
None
}
}
);
);
impl_schema!(bool, Schema::Boolean);
impl_schema!(i8, Schema::Int);
impl_schema!(i16, Schema::Int);
impl_schema!(i32, Schema::Int);
impl_schema!(i64, Schema::Long);
impl_schema!(u8, Schema::Int);
impl_schema!(u16, Schema::Int);
impl_schema!(u32, Schema::Long);
impl_schema!(f32, Schema::Float);
impl_schema!(f64, Schema::Double);
impl_schema!(String, Schema::String);
impl_schema!(str, Schema::String);
impl_schema!(char, Schema::String);
impl_schema!((), Schema::Null);
macro_rules! impl_passthrough_schema (
($type:ty where T: AvroSchemaComponent + ?Sized $(+ $bound:tt)*) => (
impl<T: AvroSchemaComponent $(+ $bound)* + ?Sized> AvroSchemaComponent for $type {
fn get_schema_in_ctxt(named_schemas: &mut HashSet<Name>, enclosing_namespace: NamespaceRef) -> Schema {
T::get_schema_in_ctxt(named_schemas, enclosing_namespace)
}
fn get_record_fields_in_ctxt(named_schemas: &mut HashSet<Name>, enclosing_namespace: NamespaceRef) -> Option<Vec<RecordField>> {
T::get_record_fields_in_ctxt(named_schemas, enclosing_namespace)
}
fn field_default() -> Option<serde_json::Value> {
T::field_default()
}
}
);
);
impl_passthrough_schema!(&T where T: AvroSchemaComponent + ?Sized);
impl_passthrough_schema!(&mut T where T: AvroSchemaComponent + ?Sized);
impl_passthrough_schema!(Box<T> where T: AvroSchemaComponent + ?Sized);
impl_passthrough_schema!(Cow<'_, T> where T: AvroSchemaComponent + ?Sized + ToOwned);
impl_passthrough_schema!(std::sync::Mutex<T> where T: AvroSchemaComponent + ?Sized);
macro_rules! impl_array_schema (
($type:ty where T: AvroSchemaComponent) => (
impl<T: AvroSchemaComponent> AvroSchemaComponent for $type {
fn get_schema_in_ctxt(named_schemas: &mut HashSet<Name>, enclosing_namespace: NamespaceRef) -> Schema {
Schema::array(T::get_schema_in_ctxt(named_schemas, enclosing_namespace)).build()
}
fn get_record_fields_in_ctxt(_: &mut HashSet<Name>, _: NamespaceRef) -> Option<Vec<RecordField>> {
None
}
}
);
);
impl_array_schema!([T] where T: AvroSchemaComponent);
impl_array_schema!(Vec<T> where T: AvroSchemaComponent);
impl<T> AvroSchemaComponent for HashMap<String, T>
where
T: AvroSchemaComponent,
{
fn get_schema_in_ctxt(
named_schemas: &mut HashSet<Name>,
enclosing_namespace: NamespaceRef,
) -> Schema {
Schema::map(T::get_schema_in_ctxt(named_schemas, enclosing_namespace)).build()
}
fn get_record_fields_in_ctxt(
_: &mut HashSet<Name>,
_: NamespaceRef,
) -> Option<Vec<RecordField>> {
None
}
}
impl<T> AvroSchemaComponent for Option<T>
where
T: AvroSchemaComponent,
{
fn get_schema_in_ctxt(
named_schemas: &mut HashSet<Name>,
enclosing_namespace: NamespaceRef,
) -> Schema {
let variants = vec![
Schema::Null,
T::get_schema_in_ctxt(named_schemas, enclosing_namespace),
];
Schema::Union(
UnionSchema::new(variants).expect("Option<T> must produce a valid (non-nested) union"),
)
}
fn get_record_fields_in_ctxt(
_: &mut HashSet<Name>,
_: NamespaceRef,
) -> Option<Vec<RecordField>> {
None
}
fn field_default() -> Option<serde_json::Value> {
Some(serde_json::Value::Null)
}
}
impl AvroSchemaComponent for core::time::Duration {
fn get_schema_in_ctxt(
named_schemas: &mut HashSet<Name>,
enclosing_namespace: NamespaceRef,
) -> Schema {
let name = Name::new("org.apache.avro.rust.Duration").expect("Name is valid");
if named_schemas.contains(&name) {
Schema::Ref { name }
} else {
named_schemas.insert(name.clone());
Schema::record(name)
.fields(
Self::get_record_fields_in_ctxt(named_schemas, enclosing_namespace)
.expect("Unreachable!"),
)
.build()
}
}
fn get_record_fields_in_ctxt(
named_schemas: &mut HashSet<Name>,
enclosing_namespace: NamespaceRef,
) -> Option<Vec<RecordField>> {
Some(vec![
RecordField::builder()
.name("secs")
.schema(u64::get_schema_in_ctxt(named_schemas, enclosing_namespace))
.build(),
RecordField::builder()
.name("nanos")
.schema(Schema::Long)
.build(),
])
}
}
impl AvroSchemaComponent for uuid::Uuid {
fn get_schema_in_ctxt(named_schemas: &mut HashSet<Name>, _: NamespaceRef) -> Schema {
let name = Name::new("org.apache.avro.rust.Uuid").expect("Name is valid");
if named_schemas.contains(&name) {
Schema::Ref { name }
} else {
let schema = Schema::Uuid(UuidSchema::Fixed(FixedSchema {
name: name.clone(),
aliases: None,
doc: None,
size: 16,
attributes: Default::default(),
}));
named_schemas.insert(name);
schema
}
}
fn get_record_fields_in_ctxt(
_: &mut HashSet<Name>,
_: NamespaceRef,
) -> Option<Vec<RecordField>> {
None
}
}
impl AvroSchemaComponent for u64 {
fn get_schema_in_ctxt(named_schemas: &mut HashSet<Name>, _: NamespaceRef) -> Schema {
let name = Name::new("org.apache.avro.rust.u64").expect("Name is valid");
if named_schemas.contains(&name) {
Schema::Ref { name }
} else {
let schema = Schema::Fixed(FixedSchema {
name: name.clone(),
aliases: None,
doc: None,
size: 8,
attributes: Default::default(),
});
named_schemas.insert(name);
schema
}
}
fn get_record_fields_in_ctxt(
_: &mut HashSet<Name>,
_: NamespaceRef,
) -> Option<Vec<RecordField>> {
None
}
}
impl AvroSchemaComponent for u128 {
fn get_schema_in_ctxt(named_schemas: &mut HashSet<Name>, _: NamespaceRef) -> Schema {
let name = Name::new("org.apache.avro.rust.u128").expect("Name is valid");
if named_schemas.contains(&name) {
Schema::Ref { name }
} else {
let schema = Schema::Fixed(FixedSchema {
name: name.clone(),
aliases: None,
doc: None,
size: 16,
attributes: Default::default(),
});
named_schemas.insert(name);
schema
}
}
fn get_record_fields_in_ctxt(
_: &mut HashSet<Name>,
_: NamespaceRef,
) -> Option<Vec<RecordField>> {
None
}
}
impl AvroSchemaComponent for i128 {
fn get_schema_in_ctxt(named_schemas: &mut HashSet<Name>, _: NamespaceRef) -> Schema {
let name = Name::new("org.apache.avro.rust.i128").expect("Name is valid");
if named_schemas.contains(&name) {
Schema::Ref { name }
} else {
let schema = Schema::Fixed(FixedSchema {
name: name.clone(),
aliases: None,
doc: None,
size: 16,
attributes: Default::default(),
});
named_schemas.insert(name);
schema
}
}
fn get_record_fields_in_ctxt(
_: &mut HashSet<Name>,
_: NamespaceRef,
) -> Option<Vec<RecordField>> {
None
}
}
impl<const N: usize, T: AvroSchemaComponent> AvroSchemaComponent for [T; N] {
fn get_schema_in_ctxt(
named_schemas: &mut HashSet<Name>,
enclosing_namespace: NamespaceRef,
) -> Schema {
if N == 0 {
Schema::Null
} else if N == 1 {
T::get_schema_in_ctxt(named_schemas, enclosing_namespace)
} else {
let t_schema = T::get_schema_in_ctxt(named_schemas, enclosing_namespace);
let name = Name::new_with_enclosing_namespace(
format!("A{N}_{}", t_schema.unique_normalized_name()),
enclosing_namespace,
)
.expect("Name is valid");
if named_schemas.contains(&name) {
Schema::Ref { name }
} else {
named_schemas.insert(name.clone());
let t_default = T::field_default();
let t_ref = T::get_schema_in_ctxt(named_schemas, enclosing_namespace);
let fields = std::iter::once(
RecordField::builder()
.name("field_0".to_string())
.schema(t_schema)
.maybe_default(t_default.clone())
.build(),
)
.chain((1..N).map(|n| {
RecordField::builder()
.name(format!("field_{n}"))
.schema(t_ref.clone())
.maybe_default(t_default.clone())
.build()
}))
.collect();
Schema::record(name).fields(fields).build()
}
}
}
fn get_record_fields_in_ctxt(
named_schemas: &mut HashSet<Name>,
enclosing_namespace: NamespaceRef,
) -> Option<Vec<RecordField>> {
if N == 0 {
None
} else if N == 1 {
T::get_record_fields_in_ctxt(named_schemas, enclosing_namespace)
} else {
let t_schema = T::get_schema_in_ctxt(named_schemas, enclosing_namespace);
let t_default = T::field_default();
let t_ref = T::get_schema_in_ctxt(named_schemas, enclosing_namespace);
let fields = std::iter::once(
RecordField::builder()
.name("field_0".to_string())
.schema(t_schema)
.maybe_default(t_default.clone())
.build(),
)
.chain((1..N).map(|n| {
RecordField::builder()
.name(format!("field_{n}"))
.schema(t_ref.clone())
.maybe_default(t_default.clone())
.build()
}))
.collect();
Some(fields)
}
}
fn field_default() -> Option<serde_json::Value> {
if N == 1 { T::field_default() } else { None }
}
}
#[cfg_attr(docsrs, doc(fake_variadic))]
impl<T: AvroSchemaComponent> AvroSchemaComponent for (T,) {
fn get_schema_in_ctxt(
named_schemas: &mut HashSet<Name>,
enclosing_namespace: NamespaceRef,
) -> Schema {
T::get_schema_in_ctxt(named_schemas, enclosing_namespace)
}
fn get_record_fields_in_ctxt(
named_schemas: &mut HashSet<Name>,
enclosing_namespace: NamespaceRef,
) -> Option<Vec<RecordField>> {
T::get_record_fields_in_ctxt(named_schemas, enclosing_namespace)
}
fn field_default() -> Option<serde_json::Value> {
T::field_default()
}
}
macro_rules! tuple_impls {
($($len:expr => ($($name:ident)+))+) => {
$(
#[cfg_attr(docsrs, doc(hidden))]
impl<$($name: AvroSchemaComponent),+> AvroSchemaComponent for ($($name),+) {
fn get_schema_in_ctxt(named_schemas: &mut HashSet<Name>, enclosing_namespace: NamespaceRef) -> Schema {
let schemas: [Schema; $len] = [$($name::get_schema_in_ctxt(named_schemas, enclosing_namespace)),+];
let mut name = format!("T{}", $len);
for schema in &schemas {
name.push('_');
name.push_str(&schema.unique_normalized_name());
}
let name = Name::new_with_enclosing_namespace(name, enclosing_namespace).expect("Name is valid");
if named_schemas.contains(&name) {
Schema::Ref { name }
} else {
named_schemas.insert(name.clone());
let defaults: [Option<serde_json::Value>; $len] = [$($name::field_default()),+];
let fields = schemas.into_iter().zip(defaults.into_iter()).enumerate().map(|(n, (schema, default))| {
RecordField::builder()
.name(format!("field_{n}"))
.schema(schema)
.maybe_default(default)
.build()
}).collect();
Schema::record(name).fields(fields).build()
}
}
}
)+
}
}
tuple_impls! {
2 => (T0 T1)
3 => (T0 T1 T2)
4 => (T0 T1 T2 T3)
5 => (T0 T1 T2 T3 T4)
6 => (T0 T1 T2 T3 T4 T5)
7 => (T0 T1 T2 T3 T4 T5 T6)
8 => (T0 T1 T2 T3 T4 T5 T6 T7)
9 => (T0 T1 T2 T3 T4 T5 T6 T7 T8)
10 => (T0 T1 T2 T3 T4 T5 T6 T7 T8 T9)
11 => (T0 T1 T2 T3 T4 T5 T6 T7 T8 T9 T10)
12 => (T0 T1 T2 T3 T4 T5 T6 T7 T8 T9 T10 T11)
13 => (T0 T1 T2 T3 T4 T5 T6 T7 T8 T9 T10 T11 T12)
14 => (T0 T1 T2 T3 T4 T5 T6 T7 T8 T9 T10 T11 T12 T13)
15 => (T0 T1 T2 T3 T4 T5 T6 T7 T8 T9 T10 T11 T12 T13 T14)
16 => (T0 T1 T2 T3 T4 T5 T6 T7 T8 T9 T10 T11 T12 T13 T14 T15)
}
#[cfg(test)]
mod tests {
use apache_avro_test_helper::TestResult;
use crate::{
AvroSchema, Schema,
reader::datum::GenericDatumReader,
schema::{FixedSchema, Name},
writer::datum::GenericDatumWriter,
};
#[test]
fn avro_rs_401_str() -> TestResult {
let schema = str::get_schema();
assert_eq!(schema, Schema::String);
Ok(())
}
#[test]
fn avro_rs_401_references() -> TestResult {
let schema_ref = <&str>::get_schema();
let schema_ref_mut = <&mut str>::get_schema();
assert_eq!(schema_ref, Schema::String);
assert_eq!(schema_ref_mut, Schema::String);
Ok(())
}
#[test]
fn avro_rs_401_slice() -> TestResult {
let schema = <[u8]>::get_schema();
assert_eq!(schema, Schema::array(Schema::Int).build());
Ok(())
}
#[test]
fn avro_rs_401_option_ref_slice_array() -> TestResult {
let schema = <Option<&[u8]>>::get_schema();
assert_eq!(
schema,
Schema::union(vec![Schema::Null, Schema::array(Schema::Int).build()])?
);
Ok(())
}
#[test]
fn avro_rs_414_char() -> TestResult {
let schema = char::get_schema();
assert_eq!(schema, Schema::String);
Ok(())
}
#[test]
fn avro_rs_414_u64() -> TestResult {
let schema = u64::get_schema();
assert_eq!(
schema,
Schema::Fixed(FixedSchema {
name: Name::new("org.apache.avro.rust.u64")?,
aliases: None,
doc: None,
size: 8,
attributes: Default::default(),
})
);
Ok(())
}
#[test]
fn avro_rs_414_i128() -> TestResult {
let schema = i128::get_schema();
assert_eq!(
schema,
Schema::Fixed(FixedSchema {
name: Name::new("org.apache.avro.rust.i128")?,
aliases: None,
doc: None,
size: 16,
attributes: Default::default(),
})
);
Ok(())
}
#[test]
fn avro_rs_414_u128() -> TestResult {
let schema = u128::get_schema();
assert_eq!(
schema,
Schema::Fixed(FixedSchema {
name: Name::new("org.apache.avro.rust.u128")?,
aliases: None,
doc: None,
size: 16,
attributes: Default::default(),
})
);
Ok(())
}
#[test]
fn avro_rs_486_unit() -> TestResult {
let schema = <()>::get_schema();
assert_eq!(schema, Schema::Null);
Ok(())
}
#[test]
#[should_panic(
expected = "Option<T> must produce a valid (non-nested) union: Error { details: Unions cannot contain duplicate types, found at least two Null }"
)]
fn avro_rs_489_some_unit() {
<Option<()>>::get_schema();
}
#[test]
#[should_panic(
expected = "Option<T> must produce a valid (non-nested) union: Error { details: Unions may not directly contain a union }"
)]
fn avro_rs_489_option_option() {
<Option<Option<i32>>>::get_schema();
}
#[test]
fn avro_rs_512_std_time_duration() -> TestResult {
let schema = Schema::parse_str(
r#"{
"type": "record",
"name": "Duration",
"namespace": "org.apache.avro.rust",
"fields": [
{ "name": "secs", "type": {"type": "fixed", "name": "u64", "namespace": "org.apache.avro.rust", "size": 8} },
{ "name": "nanos", "type": "long" }
]
}"#,
)?;
let zero = std::time::Duration::ZERO;
let max = std::time::Duration::MAX;
assert_eq!(schema, std::time::Duration::get_schema());
let writer = GenericDatumWriter::builder(&schema).build()?;
let written_zero = writer.write_ser_to_vec(&zero)?;
let written_max = writer.write_ser_to_vec(&max)?;
let reader = GenericDatumReader::builder(&schema).build()?;
let read_zero = reader.read_deser(&mut &written_zero[..])?;
assert_eq!(zero, read_zero);
let read_max = reader.read_deser(&mut &written_max[..])?;
assert_eq!(max, read_max);
Ok(())
}
#[test]
fn avro_rs_512_0_array() -> TestResult {
assert_eq!(Schema::Null, <[String; 0]>::get_schema());
assert_eq!(Schema::Null, <[(); 0]>::get_schema());
assert_eq!(Schema::Null, <[bool; 0]>::get_schema());
Ok(())
}
#[test]
fn avro_rs_512_1_array() -> TestResult {
assert_eq!(Schema::String, <[String; 1]>::get_schema());
assert_eq!(Schema::Null, <[(); 1]>::get_schema());
assert_eq!(Schema::Boolean, <[bool; 1]>::get_schema());
Ok(())
}
#[test]
fn avro_rs_512_n_array() -> TestResult {
let schema = Schema::parse_str(
r#"{
"type": "record",
"name": "A5_s",
"fields": [
{ "name": "field_0", "type": "string" },
{ "name": "field_1", "type": "string" },
{ "name": "field_2", "type": "string" },
{ "name": "field_3", "type": "string" },
{ "name": "field_4", "type": "string" }
]
}"#,
)?;
assert_eq!(schema, <[String; 5]>::get_schema());
Ok(())
}
#[test]
fn avro_rs_512_n_array_complex_type() -> TestResult {
let schema = Schema::parse_str(
r#"{
"type": "record",
"name": "A2_u2_n_r25_org_apache_avro_rust_Uuid",
"fields": [
{ "name": "field_0", "type": ["null", {"type": "fixed", "logicalType": "uuid", "size": 16, "name": "Uuid", "namespace": "org.apache.avro.rust"}], "default": null },
{ "name": "field_1", "type": ["null", "org.apache.avro.rust.Uuid"], "default": null }
]
}"#,
)?;
assert_eq!(schema, <[Option<uuid::Uuid>; 2]>::get_schema());
Ok(())
}
#[test]
fn avro_rs_512_1_tuple() -> TestResult {
assert_eq!(Schema::String, <(String,)>::get_schema());
assert_eq!(Schema::Null, <((),)>::get_schema());
assert_eq!(Schema::Boolean, <(bool,)>::get_schema());
Ok(())
}
#[test]
fn avro_rs_512_n_tuple() -> TestResult {
let schema = Schema::parse_str(
r#"{
"type": "record",
"name": "T5_s_i_l_B_n",
"fields": [
{ "name": "field_0", "type": "string" },
{ "name": "field_1", "type": "int" },
{ "name": "field_2", "type": "long" },
{ "name": "field_3", "type": "boolean" },
{ "name": "field_4", "type": "null" }
]
}"#,
)?;
assert_eq!(schema, <(String, i32, i64, bool, ())>::get_schema());
Ok(())
}
#[test]
fn avro_rs_512_n_tuple_complex_type() -> TestResult {
let schema = Schema::parse_str(
r#"{
"type": "record",
"name": "T3_u2_n_r25_org_apache_avro_rust_Uuid_r25_org_apache_avro_rust_Uuid_s",
"fields": [
{ "name": "field_0", "type": ["null", {"type": "fixed", "logicalType": "uuid", "size": 16, "name": "Uuid", "namespace": "org.apache.avro.rust"}], "default": null },
{ "name": "field_1", "type": "org.apache.avro.rust.Uuid" },
{ "name": "field_2", "type": "string" }
]
}"#,
)?;
assert_eq!(
schema,
<(Option<uuid::Uuid>, uuid::Uuid, String)>::get_schema()
);
Ok(())
}
}