use std::{ops::RangeInclusive, rc::Rc, sync::Arc};
use thiserror::Error;
use std::fmt::{self, Display, Formatter};
use crate::{
antlr::{
hamelinparser::{
ArrayLiteralContext, ArrayLiteralContextAttrs, BinaryLiteralContext,
BinaryOperatorContext, BooleanLiteralContext, BooleanLiteralContextAttrs, CastContext,
DecimalLiteralContext, DoubleLiteralContext, ExpressionContextAll,
ExpressionEOFContextAttrs, FieldLookupContext, FieldReferenceAltContext,
FieldReferenceAltContextAttrs, FieldReferenceContextAttrs, FunctionCallContext,
FunctionCallContextAttrs, IndexAccessContext, IntegerLiteralContext,
IntervalLiteralContext, IntervalLiteralContextAttrs, LambdaExpressionContext,
LambdaExpressionContextAttrs, LambdaParamsContextAttrs, NamedArgumentContextAttrs,
NullLiteralContext, NumberContextAll, NumericLiteralContextAttrs, PairLiteralContext,
ParenthesizedExpressionContextAttrs, PositionalArgumentContextAttrs,
RowsLiteralContext, ScientificLiteralContext, StringContextAll, StringLiteralContext,
StringLiteralContextAttrs, StructLiteralContext, StructLiteralContextAttrs,
TemplateParameterExprContextAttrs, TemplatedTsTruncContext,
TemplatedTsTruncContextAttrs, TsTruncContext, TsTruncContextAttrs, TupleLiteralContext,
TupleLiteralContextAttrs, UnaryPostfixOperatorContext,
UnaryPostfixOperatorContextAttrs, UnaryPrefixOperatorContext,
UnaryPrefixOperatorContextAttrs, UnboundRangeLiteralContext,
},
interval, resilient_parse_expression, strip_numeric_separators,
},
err::{TranslationError, TranslationErrors},
tree::{
ast::{
context::{FromCst, ParseContext, TryFromCst},
display::write_comma_list,
identifier::{Identifier, ParsedSimpleIdentifier, SimpleIdentifier},
ops::{BinaryOp, UnaryPostfixOp, UnaryPrefixOp},
ParseWithErrors,
},
builder::IntoExpressionBuilder,
},
types::Type,
write_utils::pad,
};
use antlr_rust::tree::ParseTree;
use derive_more::derive::{From, TryUnwrap};
use serde::Serialize;
use tsify::Tsify;
use vecmap::VecMap;
use super::node::{Span, Spannable};
#[derive(Debug, Clone)]
pub struct Expression {
pub span: Span,
pub kind: ExpressionKind,
}
impl Expression {
pub fn from_kind<T>(kind: T) -> Self
where
T: Into<ExpressionKind>,
{
Self {
span: Span::NONE,
kind: kind.into(),
}
}
}
impl ParseWithErrors for Expression {
fn parse_with_errors(input: impl Into<String>) -> (Self, TranslationErrors) {
let mut ctx = ParseContext::new();
match resilient_parse_expression(input.into()) {
Ok((csteof, parse_errors)) => match csteof.expression() {
Some(cst) => {
let ast = Self::from_cst_with_context(cst, &mut ctx);
let mut errors = ctx.take_errors();
errors.extend(parse_errors);
(ast, errors)
}
None => {
let err = ctx.error("parse tree did not contain an expression").emit();
let mut errors = ctx.take_errors();
errors.extend(parse_errors);
(Self::from_kind(ErrorExpression { error: err }), errors)
}
},
Err(e) => {
let err = ctx
.error("fatal error initializing parser")
.with_source_boxed(e.into())
.emit();
let errors = ctx.take_errors();
(Self::from_kind(ErrorExpression { error: err }), errors)
}
}
}
}
impl Spannable for Expression {
fn span(&self) -> Option<RangeInclusive<usize>> {
self.span.to_range()
}
}
impl From<ExpressionKind> for Expression {
fn from(kind: ExpressionKind) -> Self {
Self {
span: Span::NONE,
kind,
}
}
}
impl From<Identifier> for Expression {
fn from(ident: Identifier) -> Self {
ident.into_expression_builder().build()
}
}
impl PartialEq for Expression {
fn eq(&self, other: &Self) -> bool {
self.kind == other.kind
}
}
#[derive(Debug, Clone, PartialEq, From, TryUnwrap)]
pub enum ExpressionKind {
IntLiteral(IntLiteral),
DecimalLiteral(DecimalLiteral),
ScientificLiteral(ScientificLiteral),
DoubleLiteral(DoubleLiteral),
BooleanLiteral(BooleanLiteral),
StringLiteral(StringLiteral),
BinaryLiteral(BinaryLiteral),
NullLiteral(NullLiteral),
ArrayLiteral(ArrayLiteral),
TupleLiteral(TupleLiteral),
PairLiteral(PairLiteral),
StructLiteral(StructLiteral),
FieldReference(FieldReference),
TemplateParameter(TemplateParameter),
UnaryPrefixOperator(UnaryPrefixOperator),
UnaryPostfixOperator(UnaryPostfixOperator),
BinaryOperator(BinaryOperator),
FunctionCall(FunctionCall),
IndexAccess(IndexAccess),
FieldLookup(FieldLookup),
Cast(Cast),
IntervalLiteral(IntervalLiteral),
TsTrunc(TsTrunc),
RowsLiteral(RowsLiteral),
UnboundRangeLiteral(UnboundRangeLiteral),
Lambda(Lambda),
Error(ErrorExpression),
}
impl From<Rc<ExpressionContextAll<'static>>> for Expression {
fn from(cst: Rc<ExpressionContextAll<'static>>) -> Self {
let mut ctx = ParseContext::new();
Self::from_cst_with_context(cst, &mut ctx)
}
}
impl FromCst<Rc<ExpressionContextAll<'static>>> for Expression {
fn from_cst_with_context(
cst: Rc<ExpressionContextAll<'static>>,
ctx: &mut ParseContext,
) -> Self {
fn helper_cst<T, C>(parse_ctx: C, ctx: &mut ParseContext) -> ExpressionKind
where
T: FromCst<C> + Into<ExpressionKind>,
{
T::from_cst_with_context(parse_ctx, ctx).into()
}
fn helper_try<T, C>(parse_ctx: C, ctx: &mut ParseContext) -> ExpressionKind
where
T: TryFromCst<C> + Into<ExpressionKind>,
{
match T::try_from_cst_with_context(parse_ctx, ctx) {
Ok(r) => r.into(),
Err(error) => ErrorExpression { error }.into(),
}
}
let kind = match cst.as_ref() {
ExpressionContextAll::NumericLiteralContext(num_ctx) => {
if let Some(number_ctx) = num_ctx.number() {
match number_ctx.as_ref() {
NumberContextAll::IntegerLiteralContext(int_ctx) => {
helper_try::<IntLiteral, _>(int_ctx, ctx)
}
NumberContextAll::DecimalLiteralContext(dec_ctx) => {
helper_try::<DecimalLiteral, _>(dec_ctx, ctx)
}
NumberContextAll::ScientificLiteralContext(sci_ctx) => {
helper_try::<ScientificLiteral, _>(sci_ctx, ctx)
}
NumberContextAll::DoubleLiteralContext(dbl_ctx) => {
helper_try::<DoubleLiteral, _>(dbl_ctx, ctx)
}
NumberContextAll::Error(_) => {
let err = ctx.error("numeric parse error").at(cst.as_ref()).emit();
ErrorExpression { error: err }.into()
}
}
} else {
let err = ctx.error("missing number context").at(cst.as_ref()).emit();
ErrorExpression { error: err }.into()
}
}
ExpressionContextAll::BooleanLiteralContext(bool_ctx) => {
helper_try::<BooleanLiteral, _>(bool_ctx, ctx)
}
ExpressionContextAll::StringLiteralContext(str_ctx) => {
helper_try::<StringLiteral, _>(str_ctx, ctx)
}
ExpressionContextAll::BinaryLiteralContext(bin_ctx) => {
helper_try::<BinaryLiteral, _>(bin_ctx, ctx)
}
ExpressionContextAll::NullLiteralContext(null_ctx) => {
helper_try::<NullLiteral, _>(null_ctx, ctx)
}
ExpressionContextAll::PairLiteralContext(pair_ctx) => {
helper_cst::<PairLiteral, _>(pair_ctx, ctx)
}
ExpressionContextAll::RowsLiteralContext(rows_ctx) => {
helper_try::<RowsLiteral, _>(rows_ctx, ctx)
}
ExpressionContextAll::UnboundRangeLiteralContext(range_ctx) => {
helper_try::<UnboundRangeLiteral, _>(range_ctx, ctx)
}
ExpressionContextAll::StructLiteralContext(struct_ctx) => {
helper_cst::<StructLiteral, _>(struct_ctx, ctx)
}
ExpressionContextAll::FieldReferenceAltContext(col_ctx) => {
helper_try::<FieldReference, _>(col_ctx, ctx)
}
ExpressionContextAll::TemplateParameterExprContext(tmpl_ctx) => {
let name = tmpl_ctx
.simpleIdentifier()
.map(|si| ParsedSimpleIdentifier::from_cst_with_context(si.as_ref(), ctx))
.unwrap_or_else(|| {
ctx.error("missing template parameter name")
.at(tmpl_ctx)
.emit()
.into()
});
TemplateParameter { name }.into()
}
ExpressionContextAll::FieldLookupContext(field_lookup_ctx) => {
helper_try::<FieldLookup, _>(field_lookup_ctx, ctx)
}
ExpressionContextAll::ArrayLiteralContext(arr_ctx) => {
helper_cst::<ArrayLiteral, _>(arr_ctx, ctx)
}
ExpressionContextAll::TemplatedTsTruncContext(tmpl_ts_ctx) => {
helper_cst::<TsTrunc, _>(tmpl_ts_ctx, ctx)
}
ExpressionContextAll::TsTruncContext(ts_ctx) => helper_cst::<TsTrunc, _>(ts_ctx, ctx),
ExpressionContextAll::TsTruncTimestampLiteralContext(ts_lit_ctx) => {
let err = ctx
.error("ts_trunc timestamp literal not yet implemented")
.at(ts_lit_ctx)
.emit();
ErrorExpression { error: err }.into()
}
ExpressionContextAll::UnaryPrefixOperatorContext(unary_ctx) => {
helper_cst::<UnaryPrefixOperator, _>(unary_ctx, ctx)
}
ExpressionContextAll::IndexAccessContext(idx_ctx) => {
helper_cst::<IndexAccess, _>(idx_ctx, ctx)
}
ExpressionContextAll::CastContext(cast_ctx) => helper_cst::<Cast, _>(cast_ctx, ctx),
ExpressionContextAll::UnaryPostfixOperatorContext(postfix_ctx) => {
helper_cst::<UnaryPostfixOperator, _>(postfix_ctx, ctx)
}
ExpressionContextAll::TupleLiteralContext(tuple_ctx) => {
helper_cst::<TupleLiteral, _>(tuple_ctx, ctx)
}
ExpressionContextAll::ParenthesizedExpressionContext(paren_ctx) => {
if let Some(inner_ctx) = paren_ctx.expression() {
return Self::from_cst_with_context(inner_ctx, ctx);
} else {
let err = ctx
.error("missing inner expression in parenthesized expression")
.at(paren_ctx)
.emit();
return Self::from_kind(ErrorExpression { error: err });
}
}
ExpressionContextAll::BinaryOperatorContext(bin_op_ctx) => {
helper_cst::<BinaryOperator, _>(bin_op_ctx, ctx)
}
ExpressionContextAll::FunctionCallContext(fn_ctx) => {
helper_try::<FunctionCall, _>(fn_ctx, ctx)
}
ExpressionContextAll::IntervalLiteralContext(interval_ctx) => {
helper_try::<IntervalLiteral, _>(interval_ctx, ctx)
}
ExpressionContextAll::LambdaExpressionContext(lambda_ctx) => {
Lambda::from_cst_with_context(lambda_ctx, ctx).into()
}
ExpressionContextAll::Error(_err_ctx) => {
let err = ctx.error("parse error").at(cst.as_ref()).emit();
ErrorExpression { error: err }.into()
}
};
Expression {
span: interval(cst.as_ref()).into(),
kind,
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct IntLiteral {
pub int: i64,
}
impl TryFromCst<&IntegerLiteralContext<'static>> for IntLiteral {
fn try_from_cst_with_context(
value: &IntegerLiteralContext<'static>,
ctx: &mut ParseContext,
) -> Result<Self, Arc<TranslationError>> {
let raw = value.get_text();
let s = strip_numeric_separators(&raw);
let maybe_int = if s.starts_with("0x") || s.starts_with("0X") {
i64::from_str_radix(&s[2..], 16)
} else if s.starts_with("0b") || s.starts_with("0B") {
i64::from_str_radix(&s[2..], 2)
} else if s.starts_with("0o") || s.starts_with("0O") {
i64::from_str_radix(&s[2..], 8)
} else {
s.parse::<i64>()
};
let int = maybe_int.map_err(|e| {
ctx.error("invalid integer literal")
.at(value)
.with_source_boxed(Box::new(e))
.emit()
})?;
Ok(Self { int })
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct DecimalLiteral {
pub unscaled_value: i128,
pub precision: u32,
pub scale: u32,
}
impl TryFromCst<&DecimalLiteralContext<'static>> for DecimalLiteral {
fn try_from_cst_with_context(
value: &DecimalLiteralContext<'static>,
ctx: &mut ParseContext,
) -> Result<Self, Arc<TranslationError>> {
let raw = value.get_text();
let without_suffix = raw.strip_suffix('m').unwrap_or(&raw);
let s = strip_numeric_separators(without_suffix);
let decimal_pos = s.find('.').ok_or_else(|| {
ctx.error("decimal literal missing decimal point")
.at(value)
.emit()
})?;
let integer_part = &s[..decimal_pos];
let fractional_part = &s[decimal_pos + 1..];
let scale = fractional_part.len() as u32;
let precision = s.replace('.', "").len() as u32;
let combined = format!("{}{}", integer_part, fractional_part);
let unscaled_value = combined.parse::<i128>().map_err(|e| {
ctx.error("invalid decimal literal")
.at(value)
.with_source_boxed(Box::new(e))
.emit()
})?;
Ok(Self {
unscaled_value,
precision,
scale,
})
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct ScientificLiteral {
pub value: f64,
}
impl TryFromCst<&ScientificLiteralContext<'static>> for ScientificLiteral {
fn try_from_cst_with_context(
value: &ScientificLiteralContext<'static>,
ctx: &mut ParseContext,
) -> Result<Self, Arc<TranslationError>> {
let s = value.get_text();
let float_value = strip_numeric_separators(&s).parse::<f64>().map_err(|e| {
ctx.error("invalid scientific literal")
.at(value)
.with_source_boxed(Box::new(e))
.emit()
})?;
Ok(Self { value: float_value })
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct DoubleLiteral {
pub value: f64,
}
impl TryFromCst<&DoubleLiteralContext<'static>> for DoubleLiteral {
fn try_from_cst_with_context(
value: &DoubleLiteralContext<'static>,
ctx: &mut ParseContext,
) -> Result<Self, Arc<TranslationError>> {
let raw = value.get_text();
let s = strip_numeric_separators(&raw);
let float_value = s.parse::<f64>().map_err(|e| {
ctx.error("invalid double literal")
.at(value)
.with_source_boxed(Box::new(e))
.emit()
})?;
Ok(Self { value: float_value })
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct BooleanLiteral {
pub value: bool,
}
impl TryFromCst<&BooleanLiteralContext<'static>> for BooleanLiteral {
fn try_from_cst_with_context(
value: &BooleanLiteralContext<'static>,
_ctx: &mut ParseContext,
) -> Result<Self, Arc<TranslationError>> {
Ok(Self {
value: value.TRUE().is_some(),
})
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct StringLiteral {
pub value: String,
}
impl TryFromCst<&StringLiteralContext<'static>> for StringLiteral {
fn try_from_cst_with_context(
value: &StringLiteralContext<'static>,
ctx: &mut ParseContext,
) -> Result<Self, Arc<TranslationError>> {
let string_ctx = value
.string()
.ok_or_else(|| ctx.error("missing string context").at(value).emit())?;
StringLiteral::try_from_cst_with_context(&string_ctx, ctx)
}
}
impl TryFromCst<&Rc<StringContextAll<'static>>> for StringLiteral {
fn try_from_cst_with_context(
value: &Rc<StringContextAll<'static>>,
ctx: &mut ParseContext,
) -> Result<Self, Arc<TranslationError>> {
let text = value.get_text();
fn literal_text(text: &str, quote_char: char) -> String {
let inner = &text[1..text.len() - 1];
let mut result = String::with_capacity(inner.len());
let mut chars = inner.chars();
while let Some(c) = chars.next() {
if c == '\\' {
match chars.next() {
Some('\\') => result.push('\\'),
Some(c) if c == quote_char => result.push(c),
Some(other) => {
result.push('\\');
result.push(other);
}
None => result.push('\\'),
}
} else {
result.push(c);
}
}
result
}
let processed_value = if text.starts_with("U&'") && text.ends_with("'") {
literal_text(&text[2..], '\'')
} else if text.starts_with("U&\"") && text.ends_with("\"") {
literal_text(&text[2..], '"')
} else if text.len() >= 2 && text.starts_with("'") && text.ends_with("'") {
literal_text(&text, '\'')
} else if text.len() >= 2 && text.starts_with("\"") && text.ends_with("\"") {
literal_text(&text, '"')
} else {
return Err(ctx
.error("malformed string literal")
.at(value.as_ref())
.emit());
};
Ok(Self {
value: processed_value,
})
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct BinaryLiteral {
pub value: Vec<u8>,
}
impl TryFromCst<&BinaryLiteralContext<'static>> for BinaryLiteral {
fn try_from_cst_with_context(
value: &BinaryLiteralContext<'static>,
ctx: &mut ParseContext,
) -> Result<Self, Arc<TranslationError>> {
let text = value.get_text();
if !text.starts_with("x'") || !text.ends_with("'") {
return Err(ctx
.error("binary literal must have format x'...'")
.at(value)
.emit());
}
let hex_str = &text[2..text.len() - 1]; let mut bytes = Vec::new();
for chunk in hex_str.as_bytes().chunks(2) {
let hex_pair = std::str::from_utf8(chunk).map_err(|e| {
ctx.error("invalid UTF-8 in binary literal")
.at(value)
.with_source_boxed(Box::new(e))
.emit()
})?;
let byte_val = u8::from_str_radix(hex_pair, 16).map_err(|e| {
ctx.error("invalid hex digit in binary literal")
.at(value)
.with_source_boxed(Box::new(e))
.emit()
})?;
bytes.push(byte_val);
}
Ok(Self { value: bytes })
}
}
#[derive(Debug, Clone, Default, PartialEq)]
pub struct NullLiteral;
impl TryFromCst<&NullLiteralContext<'static>> for NullLiteral {
fn try_from_cst_with_context(
_value: &NullLiteralContext<'static>,
_ctx: &mut ParseContext,
) -> Result<Self, Arc<TranslationError>> {
Ok(Self)
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct ArrayLiteral {
pub elements: Vec<Arc<Expression>>,
}
impl FromCst<&ArrayLiteralContext<'static>> for ArrayLiteral {
fn from_cst_with_context(value: &ArrayLiteralContext<'static>, ctx: &mut ParseContext) -> Self {
let expressions = value.expression_all();
let mut elements = Vec::new();
for expression in expressions {
let element = Expression::from_cst_with_context(expression, ctx);
elements.push(Arc::new(element));
}
Self { elements }
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct TupleLiteral {
pub elements: Vec<Arc<Expression>>,
}
impl FromCst<&TupleLiteralContext<'static>> for TupleLiteral {
fn from_cst_with_context(value: &TupleLiteralContext<'static>, ctx: &mut ParseContext) -> Self {
let expressions = value.expression_all();
let mut elements = Vec::new();
for expression in expressions {
let element = Expression::from_cst_with_context(expression, ctx);
elements.push(Arc::new(element));
}
Self { elements }
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct PairLiteral {
pub left: Arc<Expression>,
pub right: Arc<Expression>,
}
impl FromCst<&PairLiteralContext<'static>> for PairLiteral {
fn from_cst_with_context(value: &PairLiteralContext<'static>, ctx: &mut ParseContext) -> Self {
let left = if let Some(left_ctx) = value.left.as_ref() {
Arc::new(Expression::from_cst_with_context(left_ctx.clone(), ctx))
} else {
let err = ctx
.error("missing left expression in pair literal")
.at(value)
.emit();
Arc::new(Expression::from_kind(ErrorExpression { error: err }))
};
let right = if let Some(right_ctx) = value.right.as_ref() {
Arc::new(Expression::from_cst_with_context(right_ctx.clone(), ctx))
} else {
let err = ctx
.error("missing right expression in pair literal")
.at(value)
.emit();
Arc::new(Expression::from_kind(ErrorExpression { error: err }))
};
Self { left, right }
}
}
#[derive(Debug, Clone, Default, PartialEq)]
pub struct StructLiteral {
pub fields: Vec<(ParsedSimpleIdentifier, Arc<Expression>)>,
}
impl StructLiteral {
pub fn find_field_mut(
&mut self,
name: &str,
) -> Option<&mut (ParsedSimpleIdentifier, Arc<Expression>)> {
self.fields
.iter_mut()
.find(|(id, _)| id.valid_ref().is_ok_and(|n| n.as_str() == name))
}
pub fn has_field(&self, name: &str) -> bool {
self.fields
.iter()
.any(|(id, _)| id.valid_ref().is_ok_and(|n| n.as_str() == name))
}
pub fn set_field(&mut self, name: &str, value: Arc<Expression>) {
if let Some(existing) = self.find_field_mut(name) {
existing.1 = value;
} else {
self.fields
.push((SimpleIdentifier::new(name).into(), value));
}
}
pub fn remove_field(&mut self, name: &str) {
self.fields
.retain(|(id, _)| !id.valid_ref().is_ok_and(|n| n.as_str() == name));
}
pub fn subelements(&self) -> usize {
self.fields
.iter()
.map(|(_, expr)| expr.subelements() + 1)
.sum::<usize>()
}
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: usize) -> fmt::Result {
let flat = self.subelements() <= 4;
write!(f, "{{")?;
if flat {
for (i, (name, expr)) in self.fields.iter().enumerate() {
if i > 0 {
write!(f, ", ")?;
}
write!(f, "{}: ", name)?;
expr.fmt_indented(f, indentation)?;
}
} else {
let inner_indent = indentation + 4;
for (i, (name, expr)) in self.fields.iter().enumerate() {
if i > 0 {
write!(f, ",")?;
}
writeln!(f)?;
pad(f, inner_indent)?;
write!(f, "{}: ", name)?;
expr.fmt_indented(f, inner_indent + name.to_string().len() + 2)?;
}
writeln!(f)?;
pad(f, indentation)?;
}
write!(f, "}}")
}
}
impl FromCst<&StructLiteralContext<'static>> for StructLiteral {
fn from_cst_with_context(
value: &StructLiteralContext<'static>,
ctx: &mut ParseContext,
) -> Self {
let identifiers = value.simpleIdentifier_all();
let expressions = value.expression_all();
if identifiers.len() != expressions.len() {
ctx.error("struct literal has mismatched number of identifiers and expressions")
.at(value)
.emit();
return Self::default();
}
let mut fields = Vec::with_capacity(identifiers.len());
for (identifier, expression) in identifiers.into_iter().zip(expressions.into_iter()) {
let field_name =
ParsedSimpleIdentifier::from_cst_with_context(identifier.as_ref(), ctx);
let field_expression = Expression::from_cst_with_context(expression, ctx);
fields.push((field_name, Arc::new(field_expression)));
}
Self { fields }
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct FieldReference {
pub field_name: ParsedSimpleIdentifier,
}
#[derive(Debug, Clone, PartialEq)]
pub struct TemplateParameter {
pub name: ParsedSimpleIdentifier,
}
impl From<&str> for FieldReference {
fn from(s: &str) -> Self {
Self {
field_name: SimpleIdentifier::new(s).into(),
}
}
}
impl From<String> for FieldReference {
fn from(s: String) -> Self {
Self {
field_name: SimpleIdentifier::new(s).into(),
}
}
}
impl From<SimpleIdentifier> for FieldReference {
fn from(id: SimpleIdentifier) -> Self {
Self {
field_name: id.into(),
}
}
}
impl TryFromCst<&FieldReferenceAltContext<'static>> for FieldReference {
fn try_from_cst_with_context(
value: &FieldReferenceAltContext<'static>,
ctx: &mut ParseContext,
) -> Result<Self, Arc<TranslationError>> {
let field_ref_ctx = value.fieldReference().ok_or_else(|| {
ctx.error("missing identifier in field reference")
.at(value)
.emit()
})?;
let identifier_ctx = field_ref_ctx.simpleIdentifier().ok_or_else(|| {
ctx.error("missing identifier in field reference")
.at(value)
.emit()
})?;
let field_name =
ParsedSimpleIdentifier::from_cst_with_context(identifier_ctx.as_ref(), ctx);
Ok(Self { field_name })
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct UnaryPrefixOperator {
pub operator: UnaryPrefixOp,
pub operand: Arc<Expression>,
}
impl FromCst<&UnaryPrefixOperatorContext<'static>> for UnaryPrefixOperator {
fn from_cst_with_context(
value: &UnaryPrefixOperatorContext<'static>,
ctx: &mut ParseContext,
) -> Self {
let operator_token = value.operator.as_ref();
let operator = if let Some(token) = operator_token {
UnaryPrefixOp::from_cst(token)
} else {
None
};
let operator = match operator {
Some(op) => op,
None => {
let err = ctx
.error("missing or unknown prefix operator")
.at(value)
.emit();
return Self {
operator: UnaryPrefixOp::Not, operand: Arc::new(Expression::from_kind(ErrorExpression { error: err })),
};
}
};
let operand = if let Some(expr_ctx) = value.expression() {
Arc::new(Expression::from_cst_with_context(expr_ctx, ctx))
} else {
let err = ctx.error("missing operand").at(value).emit();
Arc::new(Expression::from_kind(ErrorExpression { error: err }))
};
Self { operator, operand }
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct UnaryPostfixOperator {
pub operand: Arc<Expression>,
pub operator: UnaryPostfixOp,
}
impl FromCst<&UnaryPostfixOperatorContext<'static>> for UnaryPostfixOperator {
fn from_cst_with_context(
value: &UnaryPostfixOperatorContext<'static>,
ctx: &mut ParseContext,
) -> Self {
let operator_token = value.operator.as_ref();
let operator = if let Some(token) = operator_token {
UnaryPostfixOp::from_cst(token)
} else {
None
};
let operator = match operator {
Some(op) => op,
None => {
let err = ctx
.error("missing or unknown postfix operator")
.at(value)
.emit();
return Self {
operand: Arc::new(Expression::from_kind(ErrorExpression {
error: err.clone(),
})),
operator: UnaryPostfixOp::Range, };
}
};
let operand = if let Some(expr_ctx) = value.expression() {
Arc::new(Expression::from_cst_with_context(expr_ctx, ctx))
} else {
let err = ctx.error("missing operand").at(value).emit();
Arc::new(Expression::from_kind(ErrorExpression { error: err }))
};
Self { operand, operator }
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct BinaryOperator {
pub left: Arc<Expression>,
pub operator: BinaryOp,
pub right: Arc<Expression>,
}
impl FromCst<&BinaryOperatorContext<'static>> for BinaryOperator {
fn from_cst_with_context(
value: &BinaryOperatorContext<'static>,
ctx: &mut ParseContext,
) -> Self {
let operator_token = value.operator.as_ref();
let operator = if let Some(token) = operator_token {
BinaryOp::from_cst(token)
} else {
None
};
let operator = match operator {
Some(op) => op,
None => {
let err = ctx
.error("missing or unknown binary operator")
.at(value)
.emit();
return Self {
left: Arc::new(Expression::from_kind(ErrorExpression {
error: err.clone(),
})),
operator: BinaryOp::Add, right: Arc::new(Expression::from_kind(ErrorExpression { error: err })),
};
}
};
let left = if let Some(left_ctx) = value.left.clone() {
Arc::new(Expression::from_cst_with_context(left_ctx, ctx))
} else {
let err = ctx.error("missing left operand").at(value).emit();
Arc::new(Expression::from_kind(ErrorExpression { error: err }))
};
let right = if let Some(right_ctx) = value.right.clone() {
Arc::new(Expression::from_cst_with_context(right_ctx, ctx))
} else {
let err = ctx.error("missing right operand").at(value).emit();
Arc::new(Expression::from_kind(ErrorExpression { error: err }))
};
Self {
left,
operator,
right,
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct FunctionCall {
pub name: ParsedSimpleIdentifier,
pub positional_args: Vec<Arc<Expression>>,
pub named_args: VecMap<ParsedSimpleIdentifier, Arc<Expression>>,
}
impl TryFromCst<&FunctionCallContext<'static>> for FunctionCall {
fn try_from_cst_with_context(
value: &FunctionCallContext<'static>,
ctx: &mut ParseContext,
) -> Result<Self, Arc<TranslationError>> {
let name_ctx = value.simpleIdentifier();
let name = if let Some(name_ctx) = name_ctx {
ParsedSimpleIdentifier::from_cst_with_context(name_ctx.as_ref(), ctx)
} else {
return Err(ctx.error("missing function name").at(value).emit());
};
let mut positional_args = Vec::new();
let mut named_args = VecMap::new();
for positional_arg in value.positionalArgument_all() {
if let Some(expr_ctx) = positional_arg.expression() {
let expression = Expression::from_cst_with_context(expr_ctx, ctx);
positional_args.push(Arc::new(expression));
}
}
for named_arg in value.namedArgument_all() {
if let (Some(identifier_ctx), Some(expr_ctx)) =
(named_arg.simpleIdentifier(), named_arg.expression())
{
let arg_name =
ParsedSimpleIdentifier::from_cst_with_context(identifier_ctx.as_ref(), ctx);
let expression = Expression::from_cst_with_context(expr_ctx, ctx);
named_args.insert(arg_name, Arc::new(expression));
}
}
let ret = Self {
name,
positional_args,
named_args,
};
Ok(ret)
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct IndexAccess {
pub value: Arc<Expression>,
pub index: Arc<Expression>,
}
impl FromCst<&IndexAccessContext<'static>> for IndexAccess {
fn from_cst_with_context(value: &IndexAccessContext<'static>, ctx: &mut ParseContext) -> Self {
let expr_value = if let Some(value_ctx) = value.value.clone() {
Arc::new(Expression::from_cst_with_context(value_ctx, ctx))
} else {
let err = ctx
.error("missing value expression in index access")
.at(value)
.emit();
Arc::new(Expression::from_kind(ErrorExpression { error: err }))
};
let index_expr = if let Some(index_ctx) = value.index.clone() {
Arc::new(Expression::from_cst_with_context(index_ctx, ctx))
} else {
let err = ctx
.error("missing index expression in index access")
.at(value)
.emit();
Arc::new(Expression::from_kind(ErrorExpression { error: err }))
};
Self {
value: expr_value,
index: index_expr,
}
}
}
#[derive(Debug, Clone)]
pub struct FieldLookup {
pub value: Arc<Expression>,
pub field_identifier: ParsedSimpleIdentifier,
}
impl PartialEq for FieldLookup {
fn eq(&self, other: &Self) -> bool {
self.value == other.value && self.field_identifier == other.field_identifier
}
}
impl TryFromCst<&FieldLookupContext<'static>> for FieldLookup {
fn try_from_cst_with_context(
value: &FieldLookupContext<'static>,
ctx: &mut ParseContext,
) -> Result<Self, Arc<TranslationError>> {
let expr_value = if let Some(left_ctx) = value.left.clone() {
Arc::new(Expression::from_cst_with_context(left_ctx, ctx))
} else {
return Err(ctx
.error("missing left expression in field lookup")
.at(value)
.emit());
};
let field_identifier = if let Some(field_cst) = value.right.clone() {
ParsedSimpleIdentifier::from_cst_with_context(field_cst.as_ref(), ctx)
} else {
return Err(ctx
.error("missing field name in field lookup")
.at(value)
.emit());
};
let ret = Self {
value: expr_value,
field_identifier,
};
Ok(ret)
}
}
#[derive(Debug, Clone)]
pub struct Cast {
pub expression: Arc<Expression>,
pub target_type: Arc<Type>,
}
impl PartialEq for Cast {
fn eq(&self, other: &Self) -> bool {
self.expression == other.expression && self.target_type == other.target_type
}
}
impl FromCst<&CastContext<'static>> for Cast {
fn from_cst_with_context(value: &CastContext<'static>, ctx: &mut ParseContext) -> Self {
let expression = if let Some(left_ctx) = value.left.clone() {
Arc::new(Expression::from_cst_with_context(left_ctx, ctx))
} else {
let err = ctx.error("missing expression in cast").at(value).emit();
Arc::new(Expression::from_kind(ErrorExpression { error: err }))
};
let target_type = if let Some(target_type_cst) = value.right.as_ref() {
match Type::try_from_cst_with_context(target_type_cst.as_ref(), ctx) {
Ok(target_type) => ctx.intern_type(target_type),
Err(_) => ctx.intern_type(Type::Unknown),
}
} else {
ctx.error("missing type in cast").at(value).emit();
ctx.intern_type(Type::Unknown)
};
Self {
expression,
target_type,
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct IntervalLiteral {
pub value: i64,
pub unit: IntervalUnit,
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub enum IntervalUnit {
Millisecond,
Second,
Minute,
Hour,
Day,
Week,
Month,
Quarter,
Year,
}
impl From<TruncUnit> for IntervalUnit {
fn from(value: TruncUnit) -> Self {
match value {
TruncUnit::Second => IntervalUnit::Second,
TruncUnit::Minute => IntervalUnit::Minute,
TruncUnit::Hour => IntervalUnit::Hour,
TruncUnit::Day => IntervalUnit::Day,
TruncUnit::Week => IntervalUnit::Week,
TruncUnit::Month => IntervalUnit::Month,
TruncUnit::Quarter => IntervalUnit::Quarter,
TruncUnit::Year => IntervalUnit::Year,
}
}
}
impl TryFromCst<&IntervalLiteralContext<'static>> for IntervalLiteral {
fn try_from_cst_with_context(
cst: &IntervalLiteralContext<'static>,
ctx: &mut ParseContext,
) -> Result<Self, Arc<TranslationError>> {
let unit = if cst.MILLISECOND_INTERVAL().is_some() {
IntervalUnit::Millisecond
} else if cst.SECOND_INTERVAL().is_some() {
IntervalUnit::Second
} else if cst.MINUTE_INTERVAL().is_some() {
IntervalUnit::Minute
} else if cst.HOUR_INTERVAL().is_some() {
IntervalUnit::Hour
} else if cst.DAY_INTERVAL().is_some() {
IntervalUnit::Day
} else if cst.WEEK_INTERVAL().is_some() {
IntervalUnit::Week
} else if cst.MONTH_INTERVAL().is_some() {
IntervalUnit::Month
} else if cst.QUARTER_INTERVAL().is_some() {
IntervalUnit::Quarter
} else if cst.YEAR_INTERVAL().is_some() {
IntervalUnit::Year
} else {
return Err(ctx.error("missing interval unit").at(cst).emit());
};
let value_text = cst.get_text();
let num_str: String = value_text
.chars()
.take_while(|c| c.is_ascii_digit() || *c == '-' || *c == '_')
.collect();
let value = strip_numeric_separators(&num_str)
.parse::<i64>()
.map_err(|e| {
ctx.error("invalid interval value")
.at(cst)
.with_source_boxed(Box::new(e))
.emit()
})?;
Ok(Self { value, unit })
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum TsTruncUnit {
Fixed { unit: TruncUnit, multiplier: u32 },
Template(TemplateParameter),
}
impl TsTruncUnit {
pub fn fixed(unit: TruncUnit, multiplier: u32) -> Self {
Self::Fixed { unit, multiplier }
}
pub fn is_template(&self) -> bool {
matches!(self, Self::Template(_))
}
pub fn as_fixed(&self) -> Option<(TruncUnit, u32)> {
match self {
Self::Fixed { unit, multiplier } => Some((*unit, *multiplier)),
Self::Template(_) => None,
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct TsTrunc {
pub expression: Arc<Expression>,
pub unit: TsTruncUnit,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Tsify)]
#[tsify(into_wasm_abi)]
pub enum TruncUnit {
Second,
Minute,
Hour,
Day,
Week,
Month,
Quarter,
Year,
}
impl IntervalLiteral {
pub fn trunc_unit_and_multiplier(&self) -> Option<(TruncUnit, u32)> {
let mult = u32::try_from(self.value).ok()?;
if mult == 0 {
return None;
}
let unit = match self.unit {
IntervalUnit::Millisecond => return None,
IntervalUnit::Second => TruncUnit::Second,
IntervalUnit::Minute => TruncUnit::Minute,
IntervalUnit::Hour => TruncUnit::Hour,
IntervalUnit::Day => TruncUnit::Day,
IntervalUnit::Week => TruncUnit::Week,
IntervalUnit::Month => TruncUnit::Month,
IntervalUnit::Quarter => TruncUnit::Quarter,
IntervalUnit::Year => TruncUnit::Year,
};
Some((unit, mult))
}
}
fn parse_trunc_multiplier(token_text: &str, ctx: &mut ParseContext, span: &impl Spannable) -> u32 {
let after_at = &token_text[1..];
let raw_digits: String = after_at
.chars()
.take_while(|c| c.is_ascii_digit() || *c == '_')
.collect();
if raw_digits.is_empty() {
return 1;
}
let digits = strip_numeric_separators(&raw_digits);
match digits.parse::<u32>() {
Ok(0) => {
ctx.error("truncation multiplier must be at least 1")
.at(span)
.emit();
1
}
Ok(n) => n,
Err(_) => {
ctx.error("invalid truncation multiplier").at(span).emit();
1
}
}
}
impl FromCst<&TsTruncContext<'static>> for TsTrunc {
fn from_cst_with_context(value: &TsTruncContext<'static>, ctx: &mut ParseContext) -> Self {
let expression = if let Some(expr_ctx) = value.expression() {
Arc::new(Expression::from_cst_with_context(expr_ctx, ctx))
} else {
let err = ctx.error("missing expression in ts_trunc").at(value).emit();
Arc::new(Expression::from_kind(ErrorExpression { error: err }))
};
let (unit, token_text) = if let Some(t) = value.SECOND_TRUNC() {
(TruncUnit::Second, t.get_text())
} else if let Some(t) = value.MINUTE_TRUNC() {
(TruncUnit::Minute, t.get_text())
} else if let Some(t) = value.HOUR_TRUNC() {
(TruncUnit::Hour, t.get_text())
} else if let Some(t) = value.DAY_TRUNC() {
(TruncUnit::Day, t.get_text())
} else if let Some(t) = value.WEEK_TRUNC() {
(TruncUnit::Week, t.get_text())
} else if let Some(t) = value.MONTH_TRUNC() {
(TruncUnit::Month, t.get_text())
} else if let Some(t) = value.QUARTER_TRUNC() {
(TruncUnit::Quarter, t.get_text())
} else if let Some(t) = value.YEAR_TRUNC() {
(TruncUnit::Year, t.get_text())
} else {
ctx.error("missing truncation unit").at(value).emit();
(TruncUnit::Second, "@s".to_string())
};
let multiplier = parse_trunc_multiplier(&token_text, ctx, value);
Self {
expression,
unit: TsTruncUnit::fixed(unit, multiplier),
}
}
}
impl FromCst<&TemplatedTsTruncContext<'static>> for TsTrunc {
fn from_cst_with_context(
value: &TemplatedTsTruncContext<'static>,
ctx: &mut ParseContext,
) -> Self {
let expression = if let Some(expr_ctx) = value.expression() {
Arc::new(Expression::from_cst_with_context(expr_ctx, ctx))
} else {
let err = ctx
.error("missing expression in templated ts_trunc")
.at(value)
.emit();
Arc::new(Expression::from_kind(ErrorExpression { error: err }))
};
let name = value
.simpleIdentifier()
.map(|si| ParsedSimpleIdentifier::from_cst_with_context(si.as_ref(), ctx))
.unwrap_or_else(|| {
ctx.error("missing template parameter name for truncation unit")
.at(value)
.emit()
.into()
});
Self {
expression,
unit: TsTruncUnit::Template(TemplateParameter { name }),
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct RowsLiteral {
pub value: i64,
}
impl TryFromCst<&RowsLiteralContext<'static>> for RowsLiteral {
fn try_from_cst_with_context(
value: &RowsLiteralContext<'static>,
ctx: &mut ParseContext,
) -> Result<Self, Arc<TranslationError>> {
let s = value.get_text();
let num_str = s
.strip_suffix("rows")
.or_else(|| s.strip_suffix("row"))
.or_else(|| s.strip_suffix("r"))
.ok_or_else(|| ctx.error("rows literal missing suffix").at(value).emit())?;
let parsed_value = strip_numeric_separators(num_str)
.parse::<i64>()
.map_err(|e| {
ctx.error("invalid rows literal value")
.at(value)
.with_source_boxed(Box::new(e))
.emit()
})?;
Ok(Self {
value: parsed_value,
})
}
}
#[derive(Debug, Clone, Default, PartialEq)]
pub struct UnboundRangeLiteral;
impl TryFromCst<&UnboundRangeLiteralContext<'static>> for UnboundRangeLiteral {
fn try_from_cst_with_context(
_value: &UnboundRangeLiteralContext<'static>,
_ctx: &mut ParseContext,
) -> Result<Self, Arc<TranslationError>> {
Ok(Self)
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct LambdaParameter {
pub name: SimpleIdentifier,
pub type_hint: Option<Arc<Type>>,
}
impl LambdaParameter {
pub fn new(name: impl Into<String>) -> Self {
Self {
name: SimpleIdentifier::new(&name.into()),
type_hint: None,
}
}
pub fn with_type(name: impl Into<String>, typ: Arc<Type>) -> Self {
Self {
name: SimpleIdentifier::new(&name.into()),
type_hint: Some(typ),
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct Lambda {
pub parameters: Vec<LambdaParameter>,
pub body: Arc<Expression>,
}
#[derive(Error, Debug, Clone, PartialEq)]
pub enum LambdaTypeHintError {
#[error("lambda has {expected} parameters but {got} type hints provided")]
ArityMismatch { expected: usize, got: usize },
#[error(
"conflicting type hint for parameter '{param}': hint is {hint}, inferred is {inferred}"
)]
ConflictingHint {
param: String,
hint: Arc<Type>,
inferred: Arc<Type>,
},
}
impl Lambda {
pub fn with_param_types(
&self,
param_types: &[Arc<Type>],
) -> Result<Lambda, LambdaTypeHintError> {
if self.parameters.len() != param_types.len() {
return Err(LambdaTypeHintError::ArityMismatch {
expected: self.parameters.len(),
got: param_types.len(),
});
}
let mut new_params = Vec::with_capacity(self.parameters.len());
for (param, typ) in self.parameters.iter().zip(param_types.iter()) {
if let Some(ref hint) = param.type_hint {
if hint != typ {
return Err(LambdaTypeHintError::ConflictingHint {
param: param.name.as_str().to_string(),
hint: hint.clone(),
inferred: typ.clone(),
});
}
}
new_params.push(LambdaParameter {
name: param.name.clone(),
type_hint: Some(typ.clone()),
});
}
Ok(Lambda {
parameters: new_params,
body: self.body.clone(),
})
}
}
impl FromCst<&LambdaExpressionContext<'static>> for Lambda {
fn from_cst_with_context(
value: &LambdaExpressionContext<'static>,
ctx: &mut ParseContext,
) -> Self {
let parameters = if let Some(params_ctx) = value.lambdaParams() {
params_ctx
.simpleIdentifier_all()
.into_iter()
.filter_map(|id_ctx| {
let parsed = ParsedSimpleIdentifier::from_cst_with_context(&*id_ctx, ctx);
match parsed.valid() {
Ok(id) => Some(LambdaParameter {
name: id,
type_hint: None,
}),
Err(err) => {
ctx.add_error(err);
None
}
}
})
.collect()
} else {
ctx.error("missing lambda parameters before '->'")
.at(value)
.emit();
vec![]
};
let body = if let Some(body_ctx) = value.expression() {
Arc::new(Expression::from_cst_with_context(body_ctx, ctx))
} else {
let err = ctx.error("missing lambda body after '->'").at(value).emit();
Arc::new(Expression::from_kind(ErrorExpression { error: err }))
};
Self { parameters, body }
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct ErrorExpression {
pub error: Arc<TranslationError>,
}
impl Expression {
pub fn subelements(&self) -> usize {
self.kind.subelements()
}
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: usize) -> fmt::Result {
self.kind.fmt_indented(f, indentation)
}
}
impl Display for Expression {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
self.fmt_indented(f, 0)
}
}
impl Expression {
pub fn contains_template_parameter(&self) -> bool {
self.kind.contains_template_parameter()
}
}
impl ExpressionKind {
pub fn contains_template_parameter(&self) -> bool {
match self {
ExpressionKind::TemplateParameter(_) => true,
ExpressionKind::IntLiteral(_)
| ExpressionKind::DecimalLiteral(_)
| ExpressionKind::ScientificLiteral(_)
| ExpressionKind::DoubleLiteral(_)
| ExpressionKind::BooleanLiteral(_)
| ExpressionKind::StringLiteral(_)
| ExpressionKind::BinaryLiteral(_)
| ExpressionKind::NullLiteral(_)
| ExpressionKind::RowsLiteral(_)
| ExpressionKind::UnboundRangeLiteral(_)
| ExpressionKind::FieldReference(_)
| ExpressionKind::IntervalLiteral(_)
| ExpressionKind::Error(_) => false,
ExpressionKind::ArrayLiteral(a) => {
a.elements.iter().any(|e| e.contains_template_parameter())
}
ExpressionKind::TupleLiteral(t) => {
t.elements.iter().any(|e| e.contains_template_parameter())
}
ExpressionKind::PairLiteral(p) => {
p.left.contains_template_parameter() || p.right.contains_template_parameter()
}
ExpressionKind::StructLiteral(s) => s
.fields
.iter()
.any(|(_, e)| e.contains_template_parameter()),
ExpressionKind::UnaryPrefixOperator(u) => u.operand.contains_template_parameter(),
ExpressionKind::UnaryPostfixOperator(u) => u.operand.contains_template_parameter(),
ExpressionKind::BinaryOperator(b) => {
b.left.contains_template_parameter() || b.right.contains_template_parameter()
}
ExpressionKind::FunctionCall(fc) => {
fc.positional_args
.iter()
.any(|e| e.contains_template_parameter())
|| fc
.named_args
.iter()
.any(|(_, e)| e.contains_template_parameter())
}
ExpressionKind::IndexAccess(ia) => {
ia.value.contains_template_parameter() || ia.index.contains_template_parameter()
}
ExpressionKind::FieldLookup(fl) => fl.value.contains_template_parameter(),
ExpressionKind::Cast(c) => c.expression.contains_template_parameter(),
ExpressionKind::TsTrunc(t) => {
t.expression.contains_template_parameter() || t.unit.is_template()
}
ExpressionKind::Lambda(l) => l.body.contains_template_parameter(),
}
}
pub fn subelements(&self) -> usize {
match self {
ExpressionKind::IntLiteral(_)
| ExpressionKind::DecimalLiteral(_)
| ExpressionKind::ScientificLiteral(_)
| ExpressionKind::DoubleLiteral(_)
| ExpressionKind::BooleanLiteral(_)
| ExpressionKind::StringLiteral(_)
| ExpressionKind::BinaryLiteral(_)
| ExpressionKind::NullLiteral(_)
| ExpressionKind::RowsLiteral(_)
| ExpressionKind::UnboundRangeLiteral(_)
| ExpressionKind::FieldReference(_)
| ExpressionKind::TemplateParameter(_)
| ExpressionKind::Error(_) => 0,
ExpressionKind::ArrayLiteral(a) => a.subelements(),
ExpressionKind::TupleLiteral(t) => t.subelements(),
ExpressionKind::PairLiteral(p) => p.subelements(),
ExpressionKind::StructLiteral(s) => s.subelements(),
ExpressionKind::UnaryPrefixOperator(u) => u.operand.subelements(),
ExpressionKind::UnaryPostfixOperator(u) => u.operand.subelements(),
ExpressionKind::BinaryOperator(b) => b.left.subelements() + b.right.subelements() + 1,
ExpressionKind::FunctionCall(fc) => fc.subelements(),
ExpressionKind::IndexAccess(ia) => ia.value.subelements() + ia.index.subelements() + 1,
ExpressionKind::FieldLookup(fl) => fl.value.subelements(),
ExpressionKind::Cast(c) => c.expression.subelements() + c.target_type.subfields() + 1,
ExpressionKind::IntervalLiteral(_) => 0,
ExpressionKind::TsTrunc(t) => t.expression.subelements(),
ExpressionKind::Lambda(l) => l.body.subelements() + l.parameters.len(),
}
}
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: usize) -> fmt::Result {
match self {
ExpressionKind::IntLiteral(lit) => write!(f, "{}", lit),
ExpressionKind::DecimalLiteral(lit) => write!(f, "{}", lit),
ExpressionKind::ScientificLiteral(lit) => write!(f, "{}", lit),
ExpressionKind::DoubleLiteral(lit) => write!(f, "{}", lit),
ExpressionKind::BooleanLiteral(lit) => write!(f, "{}", lit),
ExpressionKind::StringLiteral(lit) => write!(f, "{}", lit),
ExpressionKind::BinaryLiteral(lit) => write!(f, "{}", lit),
ExpressionKind::NullLiteral(_) => write!(f, "null"),
ExpressionKind::RowsLiteral(lit) => write!(f, "{}", lit),
ExpressionKind::UnboundRangeLiteral(_) => write!(f, ".."),
ExpressionKind::FieldReference(cr) => write!(f, "{}", cr.field_name),
ExpressionKind::TemplateParameter(p) => {
write!(f, "$")?;
write!(f, "{{")?;
write!(f, "{}", p.name)?;
write!(f, "}}")
}
ExpressionKind::ArrayLiteral(a) => a.fmt_indented(f, indentation),
ExpressionKind::TupleLiteral(t) => t.fmt_indented(f, indentation),
ExpressionKind::PairLiteral(p) => p.fmt_indented(f, indentation),
ExpressionKind::StructLiteral(s) => s.fmt_indented(f, indentation),
ExpressionKind::UnaryPrefixOperator(u) => u.fmt_indented(f, indentation),
ExpressionKind::UnaryPostfixOperator(u) => u.fmt_indented(f, indentation),
ExpressionKind::BinaryOperator(b) => b.fmt_indented(f, indentation),
ExpressionKind::FunctionCall(fc) => fc.fmt_indented(f, indentation),
ExpressionKind::IndexAccess(ia) => ia.fmt_indented(f, indentation),
ExpressionKind::FieldLookup(fl) => fl.fmt_indented(f, indentation),
ExpressionKind::Cast(c) => c.fmt_indented(f, indentation),
ExpressionKind::IntervalLiteral(il) => write!(f, "{}", il),
ExpressionKind::TsTrunc(t) => t.fmt_indented(f, indentation),
ExpressionKind::Lambda(l) => l.fmt_indented(f, indentation),
ExpressionKind::Error(_) => write!(f, "?!"),
}
}
}
impl Display for IntLiteral {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.int)
}
}
impl Display for DecimalLiteral {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
let s = self.unscaled_value.to_string();
if self.scale == 0 {
write!(f, "{}.0m", s)
} else {
let scale = self.scale as usize;
let abs_s = if s.starts_with('-') { &s[1..] } else { &s };
let neg = s.starts_with('-');
if abs_s.len() <= scale {
let zeros = scale - abs_s.len();
if neg {
write!(f, "-0.{}{}m", "0".repeat(zeros), abs_s)
} else {
write!(f, "0.{}{}m", "0".repeat(zeros), abs_s)
}
} else {
let split_pos = abs_s.len() - scale;
let integer_part = &abs_s[..split_pos];
let fractional_part = &abs_s[split_pos..];
if neg {
write!(f, "-{}.{}m", integer_part, fractional_part)
} else {
write!(f, "{}.{}m", integer_part, fractional_part)
}
}
}
}
}
impl Display for ScientificLiteral {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
write!(f, "{:e}", self.value)
}
}
impl Display for DoubleLiteral {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
if self.value.fract() == 0.0 && self.value.is_finite() {
write!(f, "{:.1}", self.value)
} else {
write!(f, "{}", self.value)
}
}
}
impl Display for BooleanLiteral {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
write!(f, "{}", if self.value { "true" } else { "false" })
}
}
impl Display for StringLiteral {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
let escaped = self.value.replace('\\', "\\\\").replace('\'', "\\'");
write!(f, "'{}'", escaped)
}
}
impl Display for BinaryLiteral {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
write!(f, "x'")?;
for byte in &self.value {
write!(f, "{:02x}", byte)?;
}
write!(f, "'")
}
}
impl Display for RowsLiteral {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
write!(f, "{}r", self.value)
}
}
impl Display for IntervalUnit {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
let suffix = match self {
IntervalUnit::Millisecond => "ms",
IntervalUnit::Second => "s",
IntervalUnit::Minute => "min",
IntervalUnit::Hour => "h",
IntervalUnit::Day => "d",
IntervalUnit::Week => "w",
IntervalUnit::Month => "mon",
IntervalUnit::Quarter => "q",
IntervalUnit::Year => "y",
};
write!(f, "{}", suffix)
}
}
impl Display for IntervalLiteral {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
write!(f, "{}{}", self.value, self.unit)
}
}
impl TruncUnit {
pub fn suffix(&self) -> &'static str {
match self {
TruncUnit::Second => "s",
TruncUnit::Minute => "min",
TruncUnit::Hour => "h",
TruncUnit::Day => "d",
TruncUnit::Week => "w",
TruncUnit::Month => "mon",
TruncUnit::Quarter => "q",
TruncUnit::Year => "y",
}
}
}
impl Display for TruncUnit {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
write!(f, "@{}", self.suffix())
}
}
impl TsTrunc {
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: usize) -> fmt::Result {
self.expression.fmt_indented(f, indentation)?;
match &self.unit {
TsTruncUnit::Fixed { unit, multiplier } => {
write!(f, "@{multiplier}{}", unit.suffix())
}
TsTruncUnit::Template(param) => {
write!(f, "@")?;
write!(f, "$")?;
write!(f, "{{")?;
write!(f, "{}", param.name)?;
write!(f, "}}")
}
}
}
}
impl ArrayLiteral {
pub fn subelements(&self) -> usize {
self.elements
.iter()
.map(|e| e.subelements() + 1)
.sum::<usize>()
}
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: usize) -> fmt::Result {
write!(f, "[")?;
let inner_indent = indentation + 1; write_comma_list(
f,
&self.elements,
inner_indent,
self.subelements(),
4,
|f, e, ind| e.fmt_indented(f, ind),
)?;
write!(f, "]")
}
}
impl TupleLiteral {
pub fn subelements(&self) -> usize {
self.elements
.iter()
.map(|e| e.subelements() + 1)
.sum::<usize>()
}
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: usize) -> fmt::Result {
write!(f, "(")?;
let inner_indent = indentation + 1; write_comma_list(
f,
&self.elements,
inner_indent,
self.subelements(),
4,
|f, e, ind| e.fmt_indented(f, ind),
)?;
write!(f, ")")
}
}
impl PairLiteral {
pub fn subelements(&self) -> usize {
self.left.subelements() + self.right.subelements() + 1
}
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: usize) -> fmt::Result {
self.left.fmt_indented(f, indentation)?;
write!(f, ": ")?;
self.right.fmt_indented(f, indentation)
}
}
impl UnaryPrefixOperator {
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: usize) -> fmt::Result {
let op_str = format!("{}", self.operator);
if op_str.chars().next().map_or(false, |c| c.is_alphabetic()) {
write!(f, "{} ", op_str)?;
} else {
write!(f, "{}", op_str)?;
}
self.operand.fmt_indented(f, indentation)
}
}
impl UnaryPostfixOperator {
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: usize) -> fmt::Result {
self.operand.fmt_indented(f, indentation)?;
write!(f, "{}", self.operator)
}
}
impl BinaryOperator {
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: usize) -> fmt::Result {
self.left.fmt_indented(f, indentation)?;
write!(f, " {} ", self.operator)?;
self.right.fmt_indented(f, indentation)
}
}
impl FunctionCall {
pub fn subelements(&self) -> usize {
let pos_count: usize = self
.positional_args
.iter()
.map(|a| a.subelements() + 1)
.sum();
let named_count: usize = self
.named_args
.iter()
.map(|(_, v)| v.subelements() + 1)
.sum();
pos_count + named_count
}
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: usize) -> fmt::Result {
let name_str = self.name.to_string();
write!(f, "{}(", name_str)?;
let total_args = self.positional_args.len() + self.named_args.len();
if total_args == 0 {
return write!(f, ")");
}
let flat = self.subelements() <= 4;
let args_indent = indentation + 4;
if !flat {
writeln!(f)?;
pad(f, args_indent)?;
}
let mut first = true;
for arg in &self.positional_args {
if !first {
if flat {
write!(f, ", ")?;
} else {
writeln!(f, ",")?;
pad(f, args_indent)?;
}
}
first = false;
arg.fmt_indented(f, args_indent)?;
}
for (name, val) in &self.named_args {
if !first {
if flat {
write!(f, ", ")?;
} else {
writeln!(f, ",")?;
pad(f, args_indent)?;
}
}
first = false;
write!(f, "{}=", name)?;
val.fmt_indented(f, args_indent)?;
}
if !flat {
writeln!(f)?;
pad(f, indentation)?;
}
write!(f, ")")
}
}
impl IndexAccess {
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: usize) -> fmt::Result {
self.value.fmt_indented(f, indentation)?;
write!(f, "[")?;
self.index.fmt_indented(f, indentation)?;
write!(f, "]")
}
}
impl FieldLookup {
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: usize) -> fmt::Result {
self.value.fmt_indented(f, indentation)?;
write!(f, ".{}", self.field_identifier)
}
}
impl Cast {
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: usize) -> fmt::Result {
self.expression.fmt_indented(f, indentation)?;
write!(f, " AS ")?;
self.target_type.fmt_indented(f, indentation, usize::MAX)
}
}
impl Lambda {
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: usize) -> fmt::Result {
if self.parameters.len() == 1 {
write!(f, "{}", self.parameters[0])?;
} else {
write!(f, "(")?;
for (i, param) in self.parameters.iter().enumerate() {
if i > 0 {
write!(f, ", ")?;
}
write!(f, "{}", param)?;
}
write!(f, ")")?;
}
write!(f, " -> ")?;
self.body.fmt_indented(f, indentation)
}
}
impl Display for LambdaParameter {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.name)?;
if let Some(ref typ) = self.type_hint {
write!(f, ": {}", typ)?;
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use rstest::rstest;
#[rstest]
#[case("1d", IntervalUnit::Day, 1)]
#[case("5h", IntervalUnit::Hour, 5)]
#[case("30min", IntervalUnit::Minute, 30)]
#[case("60s", IntervalUnit::Second, 60)]
#[case("3mon", IntervalUnit::Month, 3)]
#[case("2y", IntervalUnit::Year, 2)]
fn test_interval_literal_parsing(
#[case] input: &str,
#[case] expected_unit: IntervalUnit,
#[case] expected_value: i64,
) {
let expr = Expression::parse(input);
match expr.kind {
ExpressionKind::IntervalLiteral(ref lit) => {
assert_eq!(lit.value, expected_value, "Value mismatch for {}", input);
let unit_matches = match (&lit.unit, &expected_unit) {
(IntervalUnit::Day, IntervalUnit::Day) => true,
(IntervalUnit::Hour, IntervalUnit::Hour) => true,
(IntervalUnit::Minute, IntervalUnit::Minute) => true,
(IntervalUnit::Second, IntervalUnit::Second) => true,
(IntervalUnit::Month, IntervalUnit::Month) => true,
(IntervalUnit::Year, IntervalUnit::Year) => true,
_ => false,
};
assert!(
unit_matches,
"Unit mismatch for {}: expected {:?}, got {:?}",
input, expected_unit, lit.unit
);
}
_ => panic!(
"Expected IntervalLiteral for input '{}', got {:?}",
input, expr.kind
),
}
}
}