use std::fmt;
use std::sync::Arc;
use crate::ast::precedence::{Assoc, BindingPower};
use crate::ast::render::{
Render, RenderConfig, RenderCtx, RenderExt as _, RenderMode, render_extension_infix,
render_extension_prefix,
};
use crate::ast::{BinaryOperator, Expr, SelectItem, SetExpr, Span, Spanned, Statement};
use crate::error::ParseResult;
use crate::parser::{Dialect, Parsed, Parser, parse_with};
use crate::tokenizer::{Operator, Token, TokenKind};
const MATCH_BP: BindingPower = BindingPower {
left: 64,
right: 65,
assoc: Assoc::Left,
};
const CMP_BP: BindingPower = BindingPower {
left: 40,
right: 41,
assoc: Assoc::NonAssoc,
};
const NEG_PREFIX_BP: u8 = 82;
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
enum OpExt {
Match {
left: Box<Expr<OpExt>>,
right: Box<Expr<OpExt>>,
span: Span,
},
Cmp {
left: Box<Expr<OpExt>>,
right: Box<Expr<OpExt>>,
span: Span,
},
Neg {
operand: Box<Expr<OpExt>>,
span: Span,
},
}
impl Spanned for OpExt {
fn span(&self) -> Span {
match self {
OpExt::Match { span, .. } | OpExt::Cmp { span, .. } | OpExt::Neg { span, .. } => *span,
}
}
}
impl Render for OpExt {
fn render(&self, ctx: &RenderCtx<'_>, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
OpExt::Match { left, right, .. } => {
render_extension_infix(ctx, f, MATCH_BP, (left, right), |f| f.write_str(" ^ "))
}
OpExt::Cmp { left, right, .. } => {
render_extension_infix(ctx, f, CMP_BP, (left, right), |f| f.write_str(" & "))
}
OpExt::Neg { operand, .. } => {
render_extension_prefix(ctx, f, NEG_PREFIX_BP, |f| f.write_str("~"), operand)
}
}
}
fn operand_binding_power(&self) -> Option<BindingPower> {
Some(match self {
OpExt::Match { .. } => MATCH_BP,
OpExt::Cmp { .. } => CMP_BP,
OpExt::Neg { .. } => BindingPower {
left: NEG_PREFIX_BP,
right: NEG_PREFIX_BP,
assoc: Assoc::Right,
},
})
}
}
#[derive(Clone, Copy)]
struct OpDialect;
impl Dialect for OpDialect {
type Ext = OpExt;
fn features(&self) -> &crate::ast::dialect::FeatureSet {
&crate::ast::dialect::FeatureSet::ANSI
}
fn peek_infix_operator_hook<'a>(
parser: &mut Parser<'a, Self>,
) -> ParseResult<Option<BindingPower>> {
Ok(match parser.peek()? {
Some(token) if token.kind == TokenKind::Operator(Operator::Caret) => Some(MATCH_BP),
Some(token) if token.kind == TokenKind::Operator(Operator::Amp) => Some(CMP_BP),
_ => None,
})
}
fn build_infix_operator<'a>(
parser: &mut Parser<'a, Self>,
op: Token,
left: Expr<OpExt>,
right: Expr<OpExt>,
) -> ParseResult<Expr<OpExt>> {
let span = left.span().union(right.span());
let (left, right) = (Box::new(left), Box::new(right));
let ext = match op.kind {
TokenKind::Operator(Operator::Caret) => OpExt::Match { left, right, span },
TokenKind::Operator(Operator::Amp) => OpExt::Cmp { left, right, span },
other => {
unreachable!("peek_infix_operator_hook only recognizes `^` and `&`, got {other:?}")
}
};
let meta = parser.make_meta(span);
Ok(Expr::Other { ext, meta })
}
fn extension_operand_binding_power(ext: &OpExt) -> Option<BindingPower> {
ext.operand_binding_power()
}
fn peek_prefix_operator_hook<'a>(parser: &mut Parser<'a, Self>) -> ParseResult<Option<u8>> {
match parser.peek()? {
Some(token) if token.kind == TokenKind::Operator(Operator::Tilde) => {
Ok(Some(NEG_PREFIX_BP))
}
_ => Ok(None),
}
}
fn build_prefix_operator<'a>(
parser: &mut Parser<'a, Self>,
op: Token,
operand: Expr<OpExt>,
) -> ParseResult<Expr<OpExt>> {
debug_assert_eq!(op.kind, TokenKind::Operator(Operator::Tilde));
let span = op.span.union(operand.span());
let ext = OpExt::Neg {
operand: Box::new(operand),
span,
};
let meta = parser.make_meta(span);
Ok(Expr::Other { ext, meta })
}
}
fn parse(src: &str) -> Parsed<Arc<str>, OpExt> {
parse_with(src, OpDialect).expect("OpDialect parses the test expression")
}
fn projection_expr(parsed: &Parsed<Arc<str>, OpExt>) -> &Expr<OpExt> {
let Statement::Query { query, .. } = &parsed.statements()[0] else {
panic!("expected a query statement");
};
let SetExpr::Select { select, .. } = &query.body else {
panic!("expected a SELECT body");
};
let SelectItem::Expr { expr, .. } = &select.projection[0] else {
panic!("expected a bare projection expression");
};
expr
}
fn canonical(src: &str) -> String {
parse_with(src, OpDialect)
.expect("OpDialect parses the test expression")
.to_string()
}
fn parenthesized(src: &str) -> String {
let parsed = parse_with(src, OpDialect).expect("OpDialect parses the test expression");
let config = RenderConfig {
mode: RenderMode::Parenthesized,
..RenderConfig::default()
};
let ctx = RenderCtx::new(parsed.resolver(), parsed.source(), &config);
parsed.statements()[0].displayed(&ctx).to_string()
}
#[test]
fn infix_hook_binds_at_its_reported_right_power() {
let parsed = parse("SELECT a ^ b + c");
let expr = projection_expr(&parsed);
let Expr::BinaryOp { left, op, .. } = expr else {
panic!("expected a top-level `+`, got {expr:?}");
};
assert_eq!(*op, BinaryOperator::Plus);
assert!(
matches!(
&**left,
Expr::Other {
ext: OpExt::Match { .. },
..
}
),
"the `^` must have bound only `a ^ b`, leaving `+ c` to the outer climb",
);
let mirror = parse("SELECT a + b ^ c");
let Expr::BinaryOp { right, op, .. } = projection_expr(&mirror) else {
panic!("expected a top-level `+`");
};
assert_eq!(*op, BinaryOperator::Plus);
assert!(matches!(
&**right,
Expr::Other {
ext: OpExt::Match { .. },
..
}
));
}
#[test]
fn infix_hook_left_associates() {
let parsed = parse("SELECT a ^ b ^ c");
let Expr::Other {
ext: OpExt::Match { left, right, .. },
..
} = projection_expr(&parsed)
else {
panic!("expected a top-level `^`");
};
assert!(
matches!(
&**left,
Expr::Other {
ext: OpExt::Match { .. },
..
}
),
"left operand should be the nested `a ^ b`",
);
assert!(
matches!(&**right, Expr::Column { .. }),
"right operand should be the bare column `c`",
);
}
#[test]
fn prefix_hook_binds_tighter_than_infix() {
let parsed = parse("SELECT ~ a ^ b");
let Expr::Other {
ext: OpExt::Match { left, .. },
..
} = projection_expr(&parsed)
else {
panic!("expected a top-level `^`");
};
assert!(
matches!(
&**left,
Expr::Other {
ext: OpExt::Neg { .. },
..
}
),
"the `^` left operand should be the prefix `~a`",
);
}
#[test]
fn renders_minimal_parens_from_the_same_binding_power() {
assert_eq!(canonical("SELECT a ^ b + c"), "SELECT a ^ b + c");
assert_eq!(canonical("SELECT a + b ^ c"), "SELECT a + b ^ c");
assert_eq!(canonical("SELECT a * b ^ c"), "SELECT a * b ^ c");
assert_eq!(canonical("SELECT a ^ b * c"), "SELECT a ^ b * c");
assert_eq!(canonical("SELECT a ^ b = c"), "SELECT a ^ b = c");
assert_eq!(canonical("SELECT a ^ b ^ c"), "SELECT a ^ b ^ c");
assert_eq!(canonical("SELECT ~ a ^ b"), "SELECT ~a ^ b");
assert_eq!(canonical("SELECT ~ a + b"), "SELECT ~a + b");
assert_eq!(canonical("SELECT a ^ (b ^ c)"), "SELECT a ^ (b ^ c)");
assert_eq!(canonical("SELECT (a + b) ^ c"), "SELECT (a + b) ^ c");
}
#[test]
fn custom_operator_trees_round_trip() {
for src in [
"SELECT a ^ b + c",
"SELECT a + b ^ c",
"SELECT a ^ b ^ c",
"SELECT a ^ (b ^ c)",
"SELECT (a + b) ^ c",
"SELECT ~ a ^ b",
"SELECT ~ (a ^ b)",
"SELECT a ^ b = c AND d",
] {
let once = canonical(src);
let twice = canonical(&once);
assert_eq!(once, twice, "canonical render of `{src}` is not a fixpoint");
}
}
#[test]
fn parenthesized_mode_fully_wraps_custom_operators() {
assert_eq!(parenthesized("SELECT a ^ b ^ c"), "SELECT ((a ^ b) ^ c)");
assert_eq!(parenthesized("SELECT a ^ b + c"), "SELECT ((a ^ b) + c)");
assert_eq!(parenthesized("SELECT ~ a ^ b"), "SELECT ((~a) ^ b)");
}
#[test]
fn nonassoc_infix_hook_rejects_chains() {
let err = parse_with("SELECT a & b & c", OpDialect)
.expect_err("the non-associative `&` does not chain");
assert_eq!(err.expected.as_str(), "the end of the operator chain");
assert_eq!(err.span, Span::new(13, 14));
let parsed = parse_with("SELECT a & b", OpDialect).expect("a single `&` parses");
assert!(matches!(
projection_expr(&parsed),
Expr::Other {
ext: OpExt::Cmp { .. },
..
}
));
}
#[test]
fn parenthesized_nonassoc_extension_resets_chain_detection() {
for (src, expected) in [
("SELECT (a & b) & c", "SELECT (a & b) & c"),
("SELECT a & (b & c)", "SELECT a & (b & c)"),
] {
let canon = canonical(src);
assert_eq!(canon, expected);
assert_eq!(
canonical(&canon),
canon,
"grouped `&` render is not a fixpoint"
);
}
let parsed = parse_with("SELECT (a & b) & c", OpDialect).expect("left grouping parses");
let Expr::Other {
ext: OpExt::Cmp { left, .. },
..
} = projection_expr(&parsed)
else {
panic!("expected the outer `&`");
};
assert!(
matches!(
&**left,
Expr::Other {
ext: OpExt::Cmp { .. },
..
}
),
"the parenthesized `a & b` is the left operand",
);
}
#[test]
fn nonassoc_extension_and_builtin_comparison_do_not_chain_across_the_boundary() {
for src in ["SELECT a & b < c", "SELECT a < b & c"] {
parse_with(src, OpDialect)
.expect_err("a non-associative extension operator does not chain with a comparison");
}
assert_eq!(canonical("SELECT (a & b) < c"), "SELECT (a & b) < c");
assert_eq!(canonical("SELECT a & (b < c)"), "SELECT a & (b < c)");
}