use crate::antlr::hamelinparser::{
HamelintypeContextAll, ParameterizedTypeContextAttrs, StructTypeContextAttrs,
TupleTypeContextAttrs, TypeWithArgumentsContextAttrs,
};
use crate::catalog::HamelinType;
use crate::err::{NonMergeableTypes, TranslationError, TranslationErrors};
use crate::sql::expression::identifier::SimpleIdentifier as SqlSimpleIdentifier;
use crate::sql::types::{
SQLAnonRowType, SQLArrayType, SQLBaseType, SQLDecimalType, SQLMapType, SQLRowType,
SQLTimestampTzType, SQLType,
};
use crate::tree::ast::context::{FromCst, ParseContext, TryFromCst};
use crate::tree::ast::identifier::{ParsedSimpleIdentifier, SimpleIdentifier};
use crate::types::array::Array;
use crate::types::decimal_type::Decimal;
use crate::types::function::Function;
use crate::types::map::Map;
use crate::types::range::Range;
use crate::types::struct_type::Struct;
use crate::types::tuple::Tuple;
use antlr_rust::tree::ParseTree;
use anyhow::bail;
use derive_more::{From, TryUnwrap};
use ordermap::OrderMap;
use std::fmt::{Display, Formatter, Write};
use std::sync::Arc;
use vecmap::VecMap;
pub mod array;
pub mod decimal_type;
pub mod function;
pub mod map;
pub mod matcher;
pub mod range;
pub mod struct_type;
pub mod tuple;
#[derive(Debug, Clone, From, TryUnwrap, PartialEq, Eq, Hash)]
pub enum Type {
Binary,
Boolean,
Interval,
CalendarInterval,
Int,
Double,
Rows,
String,
Timestamp,
Unknown,
Decimal(Decimal),
Array(Array),
Function(Function),
Map(Map),
Tuple(Tuple),
Variant,
Range(Range),
Struct(Struct),
}
pub const INT: Type = Type::Int;
pub const DOUBLE: Type = Type::Double;
pub const ROWS: Type = Type::Rows;
pub const BINARY: Type = Type::Binary;
pub const BOOLEAN: Type = Type::Boolean;
pub const INTERVAL: Type = Type::Interval;
pub const CALENDAR_INTERVAL: Type = Type::CalendarInterval;
pub const STRING: Type = Type::String;
pub const TIMESTAMP: Type = Type::Timestamp;
pub const UNKNOWN: Type = Type::Unknown;
pub const VARIANT: Type = Type::Variant;
impl Type {
pub fn from_sql(sql_type: SQLType) -> anyhow::Result<Self> {
match sql_type {
SQLType::SQLBaseType(sbt) => match sbt {
SQLBaseType::TinyInt
| SQLBaseType::SmallInt
| SQLBaseType::Integer
| SQLBaseType::BigInt => Ok(INT),
SQLBaseType::Real | SQLBaseType::Double => Ok(DOUBLE),
SQLBaseType::Char | SQLBaseType::VarChar | SQLBaseType::Json => Ok(STRING),
SQLBaseType::Boolean => Ok(BOOLEAN),
SQLBaseType::VarBinary => Ok(BINARY),
SQLBaseType::Date | SQLBaseType::TimeStamp | SQLBaseType::TimeStampTz => {
Ok(TIMESTAMP)
}
SQLBaseType::Time | SQLBaseType::TimeTz => {
bail!("Time type not supported in Hamelin")
}
SQLBaseType::IntervalDayToSecond => Ok(INTERVAL),
SQLBaseType::IntervalYearToMonth => Ok(CALENDAR_INTERVAL),
SQLBaseType::Unknown => Ok(UNKNOWN),
SQLBaseType::IpAddress => bail!("IPADDRESS type not supported in hamelin"),
},
SQLType::SQLArrayType(SQLArrayType { element_type }) => {
Ok(Array::new(Type::from_sql(*element_type)?).into())
}
SQLType::SQLCharType(_) => Ok(STRING),
SQLType::SQLDecimalType(SQLDecimalType { precision, scale }) => {
Ok(Decimal::new(precision as i32, scale as i32)?.into())
}
SQLType::SQLMapType(SQLMapType {
key_type,
value_type,
}) => Ok(Map::new(Type::from_sql(*key_type)?, Type::from_sql(*value_type)?).into()),
SQLType::SQLRowType(SQLRowType { bindings }) => {
let mut ret = VecMap::new();
for (k, v) in bindings.into_iter() {
ret.insert(SimpleIdentifier::new(&k.name), Type::from_sql(v)?);
}
Ok(Struct::new(ret).into())
}
SQLType::SQLAnonRowType(SQLAnonRowType { elements }) => Ok(Tuple::new(
elements
.into_iter()
.map(|t| Type::from_sql(t))
.collect::<anyhow::Result<Vec<_>>>()?,
)
.into()),
SQLType::SQLTimeType(_) => {
bail!("Time type not supported in Hamelin")
}
SQLType::SQLTimestampType(_) | SQLType::SQLTimestampTzType(_) => Ok(TIMESTAMP),
SQLType::SQLVarBinaryType(_) => Ok(BINARY),
SQLType::SQLVarcharType(_) => Ok(STRING),
}
}
pub fn from_parse_tree(
hamelin_type: &HamelintypeContextAll<'static>,
) -> Result<Self, TranslationErrors> {
let mut ctx = ParseContext::new();
Type::try_from_cst_with_context(hamelin_type, &mut ctx).map_err(|_| ctx.take_errors())
}
pub fn to_sql(self) -> anyhow::Result<SQLType> {
match self {
Type::Int => Ok(SQLBaseType::BigInt.into()),
Type::Double => Ok(SQLBaseType::Double.into()),
Type::Decimal(Decimal { precision, scale }) => {
Ok(SQLDecimalType::new(precision as u64, scale as u64).into())
}
Type::Rows => Ok(SQLBaseType::BigInt.into()),
Type::Binary => Ok(SQLBaseType::VarBinary.into()),
Type::Boolean => Ok(SQLBaseType::Boolean.into()),
Type::Interval => Ok(SQLBaseType::IntervalDayToSecond.into()),
Type::CalendarInterval => Ok(SQLBaseType::IntervalYearToMonth.into()),
Type::String => Ok(SQLBaseType::VarChar.into()),
Type::Timestamp => Ok(SQLTimestampTzType::new(6).into()),
Type::Unknown => bail!("Unknown type."),
Type::Array(Array { element_type }) => {
Ok(SQLArrayType::new(element_type.as_ref().clone().to_sql()?).into())
}
Type::Map(Map {
key_type,
value_type,
}) => Ok(SQLMapType::new(
key_type.as_ref().clone().to_sql()?,
value_type.as_ref().clone().to_sql()?,
)
.into()),
Type::Tuple(Tuple { elements }) => Ok(SQLAnonRowType::new(
elements
.into_iter()
.map(|t| t.as_ref().clone().to_sql())
.collect::<anyhow::Result<Vec<SQLType>>>()?,
)
.into()),
Type::Variant => Ok(SQLBaseType::Json.into()),
Type::Range(Range { of }) => {
let mut ret = OrderMap::new();
let of_sql = of.as_ref().clone().to_sql()?;
ret.insert(SqlSimpleIdentifier::new("begin"), of_sql.clone());
ret.insert(SqlSimpleIdentifier::new("end"), of_sql);
Ok(SQLRowType::new(ret).into())
}
Type::Struct(s) => {
let mut ret = OrderMap::new();
for (k, v) in s.iter() {
ret.insert(k.clone().into(), v.clone().to_sql()?);
}
Ok(SQLRowType::new(ret).into())
}
Type::Function(_) => {
bail!("Function types cannot be converted to SQL types")
}
}
}
pub fn merge(self, other: Type) -> Result<Type, NonMergeableTypes> {
match (self, other) {
(left, right) if left == right => Ok(right),
(left, UNKNOWN) => Ok(left),
(UNKNOWN, right) => Ok(right),
(Type::Array(a1), Type::Array(a2)) => a1
.element_type
.as_ref()
.clone()
.merge(a2.element_type.as_ref().clone())
.map(|t| Array::new(t).into()),
(Type::Map(m1), Type::Map(m2)) => {
let merged_key = m1
.key_type
.as_ref()
.clone()
.merge(m2.key_type.as_ref().clone())?;
let merged_val = m1
.value_type
.as_ref()
.clone()
.merge(m2.value_type.as_ref().clone())?;
Ok(Map::new(merged_key, merged_val).into())
}
(Type::Tuple(t1), Type::Tuple(t2)) => {
if t1.elements.len() != t2.elements.len() {
return Err(NonMergeableTypes::new(Type::Tuple(t1), Type::Tuple(t2)));
}
let mut merged = Vec::with_capacity(t1.elements.len());
for (l, r) in t1.elements.iter().zip(t2.elements.iter()) {
merged.push(l.as_ref().clone().merge(r.as_ref().clone())?);
}
Ok(Tuple::new(merged).into())
}
(Type::Struct(left), Type::Struct(right)) => left.merge(&right).map(|t| t.into()),
(left, right) => Err(NonMergeableTypes::new(left, right)),
}
}
pub fn fmt_indented(
&self,
f: &mut impl Write,
indentation: usize,
max_fields: usize,
) -> std::fmt::Result {
match self {
Type::Binary => write!(f, "binary"),
Type::Boolean => write!(f, "boolean"),
Type::Interval => write!(f, "interval"),
Type::CalendarInterval => write!(f, "calendar_interval"),
Type::Int => write!(f, "int"),
Type::Double => write!(f, "double"),
Type::Decimal(decimal) => {
write!(f, "decimal({}, {})", decimal.precision, decimal.scale)
}
Type::Rows => write!(f, "rows"),
Type::String => write!(f, "string"),
Type::Timestamp => write!(f, "timestamp"),
Type::Unknown => write!(f, "unknown"),
Type::Array(a) => a.fmt_indented(f, indentation + 4, max_fields),
Type::Function(func) => func.fmt_indented(f, indentation + 4, max_fields),
Type::Map(m) => m.fmt_indented(f, indentation + 4, max_fields),
Type::Tuple(t) => t.fmt_indented(f, indentation + 4, max_fields),
Type::Variant => write!(f, "variant"),
Type::Range(r) => r.fmt_indented(f, indentation + 4, max_fields),
Type::Struct(s) => s.fmt_indented(f, indentation, max_fields),
}
}
pub fn short_name(&self) -> String {
match self {
Type::Struct(_) => "struct".to_string(),
Type::Tuple(_) => "tuple".to_string(),
Type::Array(_) => "array".to_string(),
Type::Function(_) => "fn".to_string(),
Type::Map(_) => "map".to_string(),
_ => format!("{}", self),
}
}
pub fn subfields(&self) -> usize {
match self {
Type::Array(a) => a.subfields(),
Type::Function(f) => f.subfields(),
Type::Map(m) => m.subfields(),
Type::Tuple(t) => t.subfields(),
Type::Struct(s) => s.subfields(),
_ => 0,
}
}
}
impl TryFromCst<&HamelintypeContextAll<'static>> for Type {
fn try_from_cst_with_context(
cst: &HamelintypeContextAll<'static>,
ctx: &mut ParseContext,
) -> Result<Self, Arc<TranslationError>> {
match cst {
HamelintypeContextAll::SimpleTypeContext(node) => {
match node.get_text().to_lowercase().as_str() {
"int" => Ok(INT),
"string" => Ok(STRING),
"double" => Ok(DOUBLE),
"decimal" => Decimal::new(38, 17).map(Into::into).map_err(|e| {
ctx.error("invalid decimal type")
.at(node)
.with_source_boxed(e.into())
.emit()
}),
"timestamp" => Ok(TIMESTAMP),
"interval" => Ok(INTERVAL),
"calendar_interval" => Ok(CALENDAR_INTERVAL),
"boolean" => Ok(BOOLEAN),
"variant" => Ok(VARIANT),
t => Err(ctx
.error(&format!("Unknown hamelin type: {}", t))
.at(node)
.emit()),
}
}
HamelintypeContextAll::ParameterizedTypeContext(node) => {
let zero_type = node
.hamelintype(0)
.ok_or_else(|| ctx.error("expected type argument").at(node).emit())?;
let type_name = node
.simpleIdentifier()
.ok_or_else(|| ctx.error("expected type name").at(node).emit())?;
match type_name.get_text().to_lowercase().as_str() {
"array" => {
let inner = Type::try_from_cst_with_context(zero_type.as_ref(), ctx)?;
Ok(Array::new(inner).into())
}
"map" => {
let key_type = Type::try_from_cst_with_context(zero_type.as_ref(), ctx)?;
let value_type_cst = node.hamelintype(1).ok_or_else(|| {
ctx.error("Map type must have two arguments")
.at(node)
.emit()
})?;
let value_type =
Type::try_from_cst_with_context(value_type_cst.as_ref(), ctx)?;
Ok(Map::new(key_type, value_type).into())
}
"range" => {
let inner_type = Type::try_from_cst_with_context(zero_type.as_ref(), ctx)?;
Ok(Range::new(inner_type).into())
}
t => Err(ctx
.error(&format!("Unknown hamelin type: {}", t))
.at(node)
.emit()),
}
}
HamelintypeContextAll::TypeWithArgumentsContext(node) => {
let type_name = node
.simpleIdentifier()
.ok_or_else(|| ctx.error("expected type name").at(node).emit())?;
match type_name.get_text().to_lowercase().as_str() {
"decimal" => {
let precision_token = node.INTEGER_VALUE(0).ok_or_else(|| {
ctx.error("Decimal type must have precision argument")
.at(node)
.emit()
})?;
let precision = precision_token.get_text().parse().map_err(|e| {
ctx.error(&format!("invalid precision: {}", e))
.at(node)
.emit()
})?;
let scale_token = node.INTEGER_VALUE(1).ok_or_else(|| {
ctx.error("Decimal type must have two arguments")
.at(node)
.emit()
})?;
let scale = scale_token.get_text().parse().map_err(|e| {
ctx.error(&format!("invalid scale: {}", e)).at(node).emit()
})?;
Decimal::new(precision, scale).map(Into::into).map_err(|e| {
ctx.error("invalid decimal type")
.at(node)
.with_source_boxed(e.into())
.emit()
})
}
t => Err(ctx
.error(&format!("Unknown hamelin type: {}", t))
.at(node)
.emit()),
}
}
HamelintypeContextAll::StructTypeContext(node) => {
let idents = node.simpleIdentifier_all();
let types = node.hamelintype_all();
if idents.len() != types.len() {
return Err(ctx
.error(&format!(
"struct has {} field names but {} types",
idents.len(),
types.len()
))
.at(node)
.emit());
}
let mut ret = VecMap::new();
for (ident_cst, type_cst) in idents.into_iter().zip(types.into_iter()) {
let ident =
ParsedSimpleIdentifier::from_cst_with_context(ident_cst.as_ref(), ctx)
.valid()?;
let r#type = Type::try_from_cst_with_context(type_cst.as_ref(), ctx)?;
ret.insert(ident, r#type);
}
Ok(Struct::new(ret).into())
}
HamelintypeContextAll::TupleTypeContext(node) => {
let all = node
.hamelintype_all()
.into_iter()
.map(|t| Type::try_from_cst_with_context(t.as_ref(), ctx))
.collect::<Result<Vec<Type>, _>>()?;
Ok(Tuple::new(all).into())
}
_ => Err(ctx.error("invalid type expression").at(cst).emit()),
}
}
}
impl Display for Type {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
self.fmt_indented(f, 0, usize::MAX)
}
}
impl TryFrom<HamelinType> for Type {
type Error = anyhow::Error;
fn try_from(value: HamelinType) -> Result<Self, Self::Error> {
let res = match value {
HamelinType::Int => Type::Int,
HamelinType::Double => Type::Double,
HamelinType::Decimal { precision, scale } => Decimal::new(precision, scale)?.into(),
HamelinType::String => Type::String,
HamelinType::Timestamp => Type::Timestamp,
HamelinType::Interval => Type::Interval,
HamelinType::CalendarInterval => Type::CalendarInterval,
HamelinType::Boolean => Type::Boolean,
HamelinType::Variant => Type::Variant,
HamelinType::Binary => Type::Binary,
HamelinType::Rows => Type::Rows,
HamelinType::Unknown => Type::Unknown,
HamelinType::Array { element_type } => {
Array::new(Self::try_from(*element_type)?).into()
}
HamelinType::Map {
key_type,
value_type,
} => Map::new(Self::try_from(*key_type)?, Self::try_from(*value_type)?).into(),
HamelinType::Tuple { elements } => {
let mut new_elements = vec![];
for t in elements {
new_elements.push(Self::try_from(t)?);
}
Tuple::new(new_elements).into()
}
HamelinType::Range { of } => Range::new(Self::try_from(*of)?).into(),
HamelinType::Struct(fields) => {
let mut new_fields = VecMap::new();
for f in fields {
new_fields.insert(f.name, Self::try_from(f.typ)?);
}
Struct::new(new_fields).into()
}
};
Ok(res)
}
}