use crate::tree::ast::identifier::SimpleIdentifier;
use crate::types::Type;
use std::fmt::{Debug, Display, Formatter};
use std::sync::Arc;
use super::{INTERVAL, UNKNOWN};
pub trait Matcher: Display + Debug {
fn matches(&self, other: &Type) -> bool;
}
#[derive(Debug)]
pub struct ExactMatcher {
pub ty: Type,
}
impl ExactMatcher {
pub fn of(ty: Type) -> Self {
Self { ty }
}
}
impl Display for ExactMatcher {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.ty)
}
}
impl Matcher for ExactMatcher {
fn matches(&self, other: &Type) -> bool {
self.ty == *other || UNKNOWN == *other
}
}
#[derive(Debug, Default)]
pub struct AnyMatcher;
impl Display for AnyMatcher {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(f, "any")
}
}
impl Matcher for AnyMatcher {
fn matches(&self, _other: &Type) -> bool {
true
}
}
#[derive(Debug, Default)]
pub struct OrMatcher {
pub matchers: Vec<Arc<dyn Matcher + Send + Sync>>,
}
impl Display for OrMatcher {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(
f,
"({})",
self.matchers
.iter()
.map(|m| m.to_string())
.collect::<Vec<_>>()
.join(" | ")
)
}
}
impl OrMatcher {
pub fn with<M: Matcher + Send + Sync + 'static>(mut self, matcher: M) -> Self {
self.matchers.push(Arc::new(matcher));
self
}
}
impl Matcher for OrMatcher {
fn matches(&self, other: &Type) -> bool {
self.matchers.iter().any(|m| m.matches(other))
}
}
#[derive(Debug, Default)]
pub struct NumericMatcher;
impl Display for NumericMatcher {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(f, "numeric")
}
}
pub fn numeric_or_interval_matcher() -> OrMatcher {
OrMatcher::default()
.with(NumericMatcher::default())
.with(ExactMatcher::of(INTERVAL))
}
impl Matcher for NumericMatcher {
fn matches(&self, other: &Type) -> bool {
match other {
Type::Binary => false,
Type::Boolean => false,
Type::Interval => false,
Type::CalendarInterval => false,
Type::Int => true,
Type::Double => true,
Type::Rows => true,
Type::String => false,
Type::Timestamp => false,
Type::Unknown => true,
Type::Decimal(_) => true,
Type::Array(_) => false,
Type::Function(_) => false,
Type::Map(_) => false,
Type::Tuple(_) => false,
Type::Variant => false,
Type::Range(_) | Type::RangeInclusive(_) => false,
Type::Struct(_) => false,
}
}
}
#[derive(Debug, Default)]
pub struct IntervalMatcher;
impl Display for IntervalMatcher {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(f, "interval")
}
}
impl Matcher for IntervalMatcher {
fn matches(&self, other: &Type) -> bool {
match other {
Type::Interval => true,
Type::CalendarInterval => true,
Type::Unknown => true,
_ => false,
}
}
}
#[derive(Debug, Default)]
pub struct BaseMatcher;
impl Display for BaseMatcher {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(f, "any base")
}
}
impl Matcher for BaseMatcher {
fn matches(&self, other: &Type) -> bool {
match other {
Type::Binary => true,
Type::Boolean => true,
Type::Interval => true,
Type::CalendarInterval => true,
Type::Int => true,
Type::Double => true,
Type::Rows => false,
Type::String => true,
Type::Timestamp => true,
Type::Unknown => true,
Type::Decimal(_) => true,
Type::Array(_) => false,
Type::Function(_) => false,
Type::Map(_) => false,
Type::Tuple(_) => false,
Type::Variant => false,
Type::Range(_) | Type::RangeInclusive(_) => false,
Type::Struct(_) => false,
}
}
}
#[derive(Debug, Default)]
pub struct MapKeyMatcher;
impl Display for MapKeyMatcher {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(f, "any base (excluding float)")
}
}
impl Matcher for MapKeyMatcher {
fn matches(&self, other: &Type) -> bool {
match other {
Type::Double => false,
_ => BaseMatcher.matches(other),
}
}
}
#[derive(Debug)]
pub struct ArrayMatcher {
pub matcher: Arc<dyn Matcher + Send + Sync>,
}
impl ArrayMatcher {
pub fn of<M: Matcher + Send + Sync + 'static>(matcher: M) -> Self {
Self {
matcher: Arc::new(matcher),
}
}
}
impl Default for ArrayMatcher {
fn default() -> Self {
Self {
matcher: Arc::new(AnyMatcher),
}
}
}
impl Display for ArrayMatcher {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(f, "array({})", self.matcher)
}
}
impl Matcher for ArrayMatcher {
fn matches(&self, other: &Type) -> bool {
matches!(other, Type::Array(a) if self.matcher.matches(a.element_type.as_ref()))
}
}
#[derive(Debug)]
pub struct MapMatcher {
pub key: Arc<dyn Matcher + Send + Sync>,
pub value: Arc<dyn Matcher + Send + Sync>,
}
impl MapMatcher {
pub fn of<KM: Matcher + Send + Sync + 'static, VM: Matcher + Send + Sync + 'static>(
key: KM,
value: VM,
) -> Self {
Self {
key: Arc::new(key),
value: Arc::new(value),
}
}
}
impl Default for MapMatcher {
fn default() -> Self {
Self {
key: Arc::new(AnyMatcher),
value: Arc::new(AnyMatcher),
}
}
}
impl Display for MapMatcher {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(f, "map({}, {})", self.key, self.value)
}
}
impl Matcher for MapMatcher {
fn matches(&self, other: &Type) -> bool {
matches!(
other,
Type::Map(m)
if self.key.matches(m.key_type.as_ref()) && self.value.matches(m.value_type.as_ref())
)
}
}
#[derive(Default, Debug)]
pub struct StructMatcher {
pub matchers: Vec<(SimpleIdentifier, Arc<dyn Matcher + Send + Sync>)>,
}
impl StructMatcher {
pub fn with<M: Matcher + Send + Sync + 'static>(
mut self,
name: SimpleIdentifier,
matcher: M,
) -> Self {
self.matchers.push((name, Arc::new(matcher)));
self
}
}
impl Display for StructMatcher {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(f, "{{")?;
let matchers_length = self.matchers.len();
for (i, (name, matcher)) in self.matchers.iter().enumerate() {
write!(f, "{}", name)?;
write!(f, ": ")?;
write!(f, "{}", matcher)?;
if i != matchers_length - 1 {
write!(f, ", ")?;
}
}
write!(f, "}}")
}
}
impl Matcher for StructMatcher {
fn matches(&self, other: &Type) -> bool {
matches!(other, Type::Struct(s) if {
self.matchers.iter().all(|(name, matcher)| {
s.lookup(name).map_or(false, |ty| matcher.matches(ty))
})
})
}
}
#[derive(Debug)]
pub struct TupleMatcher {
pub matchers: Vec<Arc<dyn Matcher + Send + Sync>>,
}
impl TupleMatcher {
pub fn of<L, R>(left: L, right: R) -> Self
where
L: Matcher + Send + Sync + 'static,
R: Matcher + Send + Sync + 'static,
{
Self {
matchers: vec![Arc::new(left), Arc::new(right)],
}
}
pub fn with<M: Matcher + Send + Sync + 'static>(mut self, matcher: M) -> Self {
self.matchers.push(Arc::new(matcher));
self
}
}
impl Display for TupleMatcher {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
if self.matchers.len() == 2 {
write!(f, "{}: {}", self.matchers[0], self.matchers[1])
} else {
write!(f, "(")?;
let matchers_length = self.matchers.len();
for (i, matcher) in self.matchers.iter().enumerate() {
write!(f, "{}", matcher)?;
if i != matchers_length - 1 {
write!(f, ", ")?;
}
}
write!(f, ")")
}
}
}
impl Matcher for TupleMatcher {
fn matches(&self, other: &Type) -> bool {
let inner = match other {
Type::Tuple(t) => t,
_ => return false,
};
inner.elements.len() == self.matchers.len()
&& self
.matchers
.iter()
.zip(inner.elements.iter())
.all(|(matcher, ty)| matcher.matches(ty))
}
}
#[derive(Debug)]
pub struct RangeMatcher {
pub matcher: Arc<dyn Matcher + Send + Sync>,
}
impl RangeMatcher {
pub fn of<M: Matcher + Send + Sync + 'static>(matcher: M) -> Self {
Self {
matcher: Arc::new(matcher),
}
}
}
impl Default for RangeMatcher {
fn default() -> Self {
Self {
matcher: Arc::new(AnyMatcher),
}
}
}
impl Display for RangeMatcher {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(f, "range({})", self.matcher)
}
}
impl Matcher for RangeMatcher {
fn matches(&self, other: &Type) -> bool {
matches!(other, Type::Range(r) | Type::RangeInclusive(r) if self.matcher.matches(r.of.as_ref()))
}
}
#[derive(Debug)]
pub struct UnboundedTupleMatcher {
pub matcher: Arc<dyn Matcher + Send + Sync>,
}
impl UnboundedTupleMatcher {
pub fn of<M: Matcher + Send + Sync + 'static>(matcher: M) -> Self {
Self {
matcher: Arc::new(matcher),
}
}
}
impl Default for UnboundedTupleMatcher {
fn default() -> Self {
Self {
matcher: Arc::new(AnyMatcher),
}
}
}
impl Display for UnboundedTupleMatcher {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(f, "({}..)", self.matcher)
}
}
impl Matcher for UnboundedTupleMatcher {
fn matches(&self, other: &Type) -> bool {
matches!(other, Type::Tuple(r) if r.elements.iter().all(|t| self.matcher.matches(t)))
}
}