use std::{
borrow::Borrow,
collections::{BTreeMap, HashMap},
fmt::{Debug, Formatter},
};
use strum::IntoDiscriminant;
use crate::{
AvroResult, Error,
error::Details,
schema::{
DecimalSchema, FixedSchema, InnerDecimalSchema, Name, NamespaceRef, RecordSchema, Schema,
SchemaKind, UuidSchema,
},
types,
};
#[derive(Clone)]
pub struct UnionSchema {
pub(crate) schemas: Vec<Schema>,
variant_index: BTreeMap<SchemaKind, usize>,
named_index: Vec<usize>,
}
impl Debug for UnionSchema {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
f.debug_struct("UnionSchema")
.field("schemas", &self.schemas)
.finish()
}
}
impl UnionSchema {
pub fn new(schemas: Vec<Schema>) -> AvroResult<Self> {
let mut builder = Self::builder();
for schema in schemas {
builder.variant(schema)?;
}
Ok(builder.build())
}
pub fn builder() -> UnionSchemaBuilder {
UnionSchemaBuilder::new()
}
pub fn variants(&self) -> &[Schema] {
&self.schemas
}
pub fn get_variant(&self, index: usize) -> Result<&Schema, Error> {
self.schemas.get(index).ok_or_else(|| {
Details::GetUnionVariant {
index: index as i64,
num_variants: self.schemas.len(),
}
.into()
})
}
pub(crate) fn index_of_schema_kind(&self, kind: SchemaKind) -> Option<usize> {
self.variant_index.get(&kind).copied()
}
pub(crate) fn find_named_schema<'s>(
&'s self,
name: &str,
names: &'s HashMap<Name, impl Borrow<Schema>>,
) -> Result<Option<(usize, &'s Schema)>, Error> {
for index in self.named_index.iter().copied() {
let schema = &self.schemas[index];
if let Some(schema_name) = schema.name()
&& schema_name.name() == name
{
let schema = if let Schema::Ref { name } = schema {
names
.get(name)
.ok_or_else(|| Details::SchemaResolutionError(name.clone()))?
.borrow()
} else {
schema
};
return Ok(Some((index, schema)));
}
}
Ok(None)
}
pub(crate) fn find_fully_qualified_named_schema<'s>(
&'s self,
full_name: &str,
names: &'s HashMap<Name, impl Borrow<Schema>>,
) -> Result<Option<(usize, &'s Schema)>, Error> {
for index in self.named_index.iter().copied() {
let schema = &self.schemas[index];
if let Some(schema_name) = schema.name()
&& schema_name.as_ref() == full_name
{
let schema = if let Schema::Ref { name } = schema {
names
.get(name)
.ok_or_else(|| Details::SchemaResolutionError(name.clone()))?
.borrow()
} else {
schema
};
return Ok(Some((index, schema)));
}
}
Ok(None)
}
pub(crate) fn find_fixed_of_size_n<'s>(
&'s self,
size: usize,
names: &'s HashMap<Name, impl Borrow<Schema>>,
) -> Result<Option<(usize, &'s FixedSchema)>, Error> {
for index in self.named_index.iter().copied() {
let schema = &self.schemas[index];
let schema = if let Schema::Ref { name } = schema {
names
.get(name)
.ok_or_else(|| Details::SchemaResolutionError(name.clone()))?
.borrow()
} else {
schema
};
match schema {
Schema::Fixed(fixed)
| Schema::Uuid(UuidSchema::Fixed(fixed))
| Schema::Decimal(DecimalSchema {
inner: InnerDecimalSchema::Fixed(fixed),
..
})
| Schema::Duration(fixed)
if fixed.size == size =>
{
return Ok(Some((index, fixed)));
}
_ => {}
}
}
Ok(None)
}
pub(crate) fn find_record_with_n_fields<'s>(
&'s self,
n_fields: usize,
names: &'s HashMap<Name, impl Borrow<Schema>>,
) -> Result<Option<(usize, &'s RecordSchema)>, Error> {
for index in self.named_index.iter().copied() {
let schema = &self.schemas[index];
let schema = if let Schema::Ref { name } = schema {
names
.get(name)
.ok_or_else(|| Details::SchemaResolutionError(name.clone()))?
.borrow()
} else {
schema
};
match schema {
Schema::Record(record) if record.fields.len() == n_fields => {
return Ok(Some((index, record)));
}
_ => {}
}
}
Ok(None)
}
pub fn is_nullable(&self) -> bool {
self.variant_index.contains_key(&SchemaKind::Null)
}
pub fn find_schema_with_known_schemata<S: Borrow<Schema> + Debug>(
&self,
value: &types::Value,
known_schemata: Option<&HashMap<Name, S>>,
enclosing_namespace: NamespaceRef,
) -> Option<(usize, &Schema)> {
let known_schemata_if_none = HashMap::new();
let known_schemata = known_schemata.unwrap_or(&known_schemata_if_none);
let ValueSchemaKind { unnamed, named } = Self::value_to_base_schemakind(value);
let unnamed = unnamed
.and_then(|kind| self.variant_index.get(&kind).copied())
.map(|index| (index, &self.schemas[index]))
.and_then(|(index, schema)| {
let kind = schema.discriminant();
if kind == SchemaKind::Map || kind == SchemaKind::Array {
let namespace = schema.namespace().or(enclosing_namespace);
value
.clone()
.resolve_internal(schema, known_schemata, namespace, None)
.ok()
.map(|_| (index, schema))
} else {
Some((index, schema))
}
});
let named = named.and_then(|kind| {
self.named_index
.iter()
.copied()
.map(|i| (i, &self.schemas[i]))
.filter(|(_i, s)| {
let s_kind = schema_to_base_schemakind(s);
s_kind == kind || s_kind == SchemaKind::Ref
})
.find(|(_i, schema)| {
let namespace = schema.namespace().or(enclosing_namespace);
value
.clone()
.resolve_internal(schema, known_schemata, namespace, None)
.is_ok()
})
});
match (unnamed, named) {
(Some((u_i, _)), Some((n_i, _))) if u_i < n_i => unnamed,
(Some(_), Some(_)) => named,
(Some(_), None) => unnamed,
(None, Some(_)) => named,
(None, None) => {
self.schemas.iter().enumerate().find(|(_i, schema)| {
let namespace = schema.namespace().or(enclosing_namespace);
value
.clone()
.resolve_internal(schema, known_schemata, namespace, None)
.is_ok()
})
}
}
}
fn value_to_base_schemakind(value: &types::Value) -> ValueSchemaKind {
let schemakind = SchemaKind::from(value);
match schemakind {
SchemaKind::Decimal => ValueSchemaKind {
unnamed: Some(SchemaKind::Bytes),
named: Some(SchemaKind::Fixed),
},
SchemaKind::BigDecimal => ValueSchemaKind {
unnamed: Some(SchemaKind::Bytes),
named: None,
},
SchemaKind::Uuid => ValueSchemaKind {
unnamed: Some(SchemaKind::String),
named: Some(SchemaKind::Fixed),
},
SchemaKind::Date | SchemaKind::TimeMillis => ValueSchemaKind {
unnamed: Some(SchemaKind::Int),
named: None,
},
SchemaKind::TimeMicros
| SchemaKind::TimestampMillis
| SchemaKind::TimestampMicros
| SchemaKind::TimestampNanos
| SchemaKind::LocalTimestampMillis
| SchemaKind::LocalTimestampMicros
| SchemaKind::LocalTimestampNanos => ValueSchemaKind {
unnamed: Some(SchemaKind::Long),
named: None,
},
SchemaKind::Duration => ValueSchemaKind {
unnamed: None,
named: Some(SchemaKind::Fixed),
},
SchemaKind::Record | SchemaKind::Enum | SchemaKind::Fixed => ValueSchemaKind {
unnamed: None,
named: Some(schemakind),
},
SchemaKind::Map => ValueSchemaKind {
unnamed: Some(SchemaKind::Map),
named: Some(SchemaKind::Record),
},
_ => ValueSchemaKind {
unnamed: Some(schemakind),
named: None,
},
}
}
}
struct ValueSchemaKind {
unnamed: Option<SchemaKind>,
named: Option<SchemaKind>,
}
impl PartialEq for UnionSchema {
fn eq(&self, other: &UnionSchema) -> bool {
self.schemas.eq(&other.schemas)
}
}
#[derive(Default, Debug)]
pub struct UnionSchemaBuilder {
schemas: Vec<Schema>,
names: HashMap<Name, usize>,
variant_index: BTreeMap<SchemaKind, usize>,
}
impl UnionSchemaBuilder {
pub fn new() -> Self {
Self::default()
}
#[doc(hidden)]
pub fn variant_ignore_duplicates(&mut self, schema: Schema) -> Result<&mut Self, Error> {
if let Some(name) = schema.name() {
if let Some(current) = self.names.get(name).copied() {
if self.schemas[current] != schema {
return Err(Details::GetUnionDuplicateNamedSchemas(name.to_string()).into());
}
} else {
self.names.insert(name.clone(), self.schemas.len());
self.schemas.push(schema);
}
} else if let Schema::Map(_) = &schema {
if let Some(index) = self.variant_index.get(&SchemaKind::Map).copied() {
if self.schemas[index] != schema {
return Err(
Details::GetUnionDuplicateMap(self.schemas[index].clone(), schema).into(),
);
}
} else {
self.variant_index
.insert(SchemaKind::Map, self.schemas.len());
self.schemas.push(schema);
}
} else if let Schema::Array(_) = &schema {
if let Some(index) = self.variant_index.get(&SchemaKind::Array).copied() {
if self.schemas[index] != schema {
return Err(Details::GetUnionDuplicateArray(
self.schemas[index].clone(),
schema,
)
.into());
}
} else {
self.variant_index
.insert(SchemaKind::Array, self.schemas.len());
self.schemas.push(schema);
}
} else {
let discriminant = schema_to_base_schemakind(&schema);
if discriminant == SchemaKind::Union {
return Err(Details::GetNestedUnion.into());
}
if !self.variant_index.contains_key(&discriminant) {
self.variant_index.insert(discriminant, self.schemas.len());
self.schemas.push(schema);
}
}
Ok(self)
}
pub fn variant(&mut self, schema: Schema) -> Result<&mut Self, Error> {
if let Some(name) = schema.name() {
if self.names.contains_key(name) {
return Err(Details::GetUnionDuplicateNamedSchemas(name.to_string()).into());
} else {
self.names.insert(name.clone(), self.schemas.len());
self.schemas.push(schema);
}
} else {
let discriminant = schema_to_base_schemakind(&schema);
if discriminant == SchemaKind::Union {
return Err(Details::GetNestedUnion.into());
}
if self.variant_index.contains_key(&discriminant) {
return Err(Details::GetUnionDuplicate(discriminant).into());
} else {
self.variant_index.insert(discriminant, self.schemas.len());
self.schemas.push(schema);
}
}
Ok(self)
}
pub fn contains(&self, schema: &Schema) -> bool {
if let Some(name) = schema.name() {
if let Some(current) = self.names.get(name).copied() {
&self.schemas[current] == schema
} else {
false
}
} else {
let discriminant = schema_to_base_schemakind(schema);
if let Some(index) = self.variant_index.get(&discriminant).copied() {
&self.schemas[index] == schema
} else {
false
}
}
}
pub fn build(mut self) -> UnionSchema {
self.schemas.shrink_to_fit();
let mut named_index: Vec<_> = self.names.into_values().collect();
named_index.sort();
UnionSchema {
variant_index: self.variant_index,
named_index,
schemas: self.schemas,
}
}
}
fn schema_to_base_schemakind(schema: &Schema) -> SchemaKind {
let kind = schema.discriminant();
match kind {
SchemaKind::Date | SchemaKind::TimeMillis => SchemaKind::Int,
SchemaKind::TimeMicros
| SchemaKind::TimestampMillis
| SchemaKind::TimestampMicros
| SchemaKind::TimestampNanos
| SchemaKind::LocalTimestampMillis
| SchemaKind::LocalTimestampMicros
| SchemaKind::LocalTimestampNanos => SchemaKind::Long,
SchemaKind::Uuid => match schema {
Schema::Uuid(UuidSchema::Bytes) => SchemaKind::Bytes,
Schema::Uuid(UuidSchema::String) => SchemaKind::String,
Schema::Uuid(UuidSchema::Fixed(_)) => SchemaKind::Fixed,
_ => unreachable!(),
},
SchemaKind::Decimal => match schema {
Schema::Decimal(DecimalSchema {
inner: InnerDecimalSchema::Bytes,
..
}) => SchemaKind::Bytes,
Schema::Decimal(DecimalSchema {
inner: InnerDecimalSchema::Fixed(_),
..
}) => SchemaKind::Fixed,
_ => unreachable!(),
},
SchemaKind::Duration => SchemaKind::Fixed,
_ => kind,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::error::{Details, Error};
use crate::schema::RecordSchema;
use crate::types::Value;
use apache_avro_test_helper::TestResult;
#[test]
fn avro_rs_402_new_union_schema() -> TestResult {
let schema1 = Schema::Int;
let schema2 = Schema::String;
let union_schema = UnionSchema::new(vec![schema1.clone(), schema2.clone()])?;
assert_eq!(union_schema.variants(), &[schema1, schema2]);
Ok(())
}
#[test]
fn avro_rs_402_new_union_schema_duplicate_names() -> TestResult {
let res = UnionSchema::new(vec![
Schema::Record(RecordSchema::builder().try_name("Same_name")?.build()),
Schema::Record(RecordSchema::builder().try_name("Same_name")?.build()),
])
.map_err(Error::into_details);
match res {
Err(Details::GetUnionDuplicateNamedSchemas(name)) => {
assert_eq!(name, Name::new("Same_name")?.to_string());
}
err => panic!("Expected GetUnionDuplicateNamedSchemas error, got: {err:?}"),
}
Ok(())
}
#[test]
fn avro_rs_489_union_schema_builder_primitive_type() -> TestResult {
let mut builder = UnionSchema::builder();
builder.variant(Schema::Null)?;
assert!(builder.variant(Schema::Null).is_err());
builder.variant_ignore_duplicates(Schema::Null)?;
builder.variant(Schema::Int)?;
assert!(builder.variant(Schema::Int).is_err());
builder.variant_ignore_duplicates(Schema::Int)?;
builder.variant(Schema::Long)?;
assert!(builder.variant(Schema::Long).is_err());
builder.variant_ignore_duplicates(Schema::Long)?;
let union = builder.build();
assert_eq!(union.schemas, &[Schema::Null, Schema::Int, Schema::Long]);
Ok(())
}
#[test]
fn avro_rs_489_union_schema_builder_complex_types() -> TestResult {
let enum_abc = Schema::parse_str(
r#"{
"type": "enum",
"name": "ABC",
"symbols": ["A", "B", "C"]
}"#,
)?;
let enum_abc_with_extra_symbol = Schema::parse_str(
r#"{
"type": "enum",
"name": "ABC",
"symbols": ["A", "B", "C", "D"]
}"#,
)?;
let enum_def = Schema::parse_str(
r#"{
"type": "enum",
"name": "DEF",
"symbols": ["D", "E", "F"]
}"#,
)?;
let fixed_abc = Schema::parse_str(
r#"{
"type": "fixed",
"name": "ABC",
"size": 1
}"#,
)?;
let fixed_foo = Schema::parse_str(
r#"{
"type": "fixed",
"name": "Foo",
"size": 1
}"#,
)?;
let mut builder = UnionSchema::builder();
builder.variant(enum_abc.clone())?;
assert!(builder.variant(enum_abc.clone()).is_err());
builder.variant_ignore_duplicates(enum_abc.clone())?;
assert!(builder.variant(fixed_abc.clone()).is_err());
assert!(
builder
.variant_ignore_duplicates(fixed_abc.clone())
.is_err()
);
assert!(builder.variant(enum_abc_with_extra_symbol.clone()).is_err());
assert!(
builder
.variant_ignore_duplicates(enum_abc_with_extra_symbol.clone())
.is_err()
);
builder.variant(enum_def.clone())?;
assert!(builder.variant(enum_def.clone()).is_err());
builder.variant_ignore_duplicates(enum_def.clone())?;
builder.variant(fixed_foo.clone())?;
assert!(builder.variant(fixed_foo.clone()).is_err());
builder.variant_ignore_duplicates(fixed_foo.clone())?;
let union = builder.build();
assert_eq!(union.variants(), &[enum_abc, enum_def, fixed_foo]);
Ok(())
}
#[test]
fn avro_rs_489_union_schema_builder_logical_types() -> TestResult {
let fixed_uuid = Schema::parse_str(
r#"{
"type": "fixed",
"name": "Uuid",
"size": 16
}"#,
)?;
let uuid = Schema::parse_str(
r#"{
"type": "fixed",
"logicalType": "uuid",
"name": "Uuid",
"size": 16
}"#,
)?;
let mut builder = UnionSchema::builder();
builder.variant(Schema::Date)?;
assert!(builder.variant(Schema::Date).is_err());
builder.variant_ignore_duplicates(Schema::Date)?;
assert!(builder.variant(Schema::Int).is_err());
builder.variant_ignore_duplicates(Schema::Int)?;
builder.variant(uuid.clone())?;
assert!(builder.variant(uuid.clone()).is_err());
builder.variant_ignore_duplicates(uuid.clone())?;
assert!(builder.variant(fixed_uuid.clone()).is_err());
assert!(
builder
.variant_ignore_duplicates(fixed_uuid.clone())
.is_err()
);
let union = builder.build();
assert_eq!(union.schemas, &[Schema::Date, uuid]);
Ok(())
}
#[test]
fn avro_rs_489_find_schema_with_known_schemata_wrong_map() -> TestResult {
let union = UnionSchema::new(vec![Schema::map(Schema::Int).build(), Schema::Null])?;
let value = Value::Map(
[("key".to_string(), Value::String("value".to_string()))]
.into_iter()
.collect(),
);
assert!(
union
.find_schema_with_known_schemata(&value, None::<&HashMap<Name, Schema>>, None)
.is_none()
);
Ok(())
}
#[test]
fn avro_rs_489_find_schema_with_known_schemata_type_promotion() -> TestResult {
let union = UnionSchema::new(vec![Schema::Long, Schema::Null])?;
let value = Value::Int(42);
assert_eq!(
union.find_schema_with_known_schemata(&value, None::<&HashMap<Name, Schema>>, None),
Some((0, &Schema::Long))
);
Ok(())
}
#[test]
fn avro_rs_489_find_schema_with_known_schemata_uuid_vs_fixed() -> TestResult {
let uuid = Schema::parse_str(
r#"{
"type": "fixed",
"logicalType": "uuid",
"name": "Uuid",
"size": 16
}"#,
)?;
let union = UnionSchema::new(vec![uuid.clone(), Schema::Null])?;
let value = Value::Fixed(16, vec![0; 16]);
assert_eq!(
union.find_schema_with_known_schemata(&value, None::<&HashMap<Name, Schema>>, None),
Some((0, &uuid))
);
Ok(())
}
}