diff --git a/src/dialect/mod.rs b/src/dialect/mod.rs index ff83a4da6..d3c01e3a7 100644 --- a/src/dialect/mod.rs +++ b/src/dialect/mod.rs @@ -535,10 +535,32 @@ pub trait Dialect: Debug + Any { /// ```sql /// SELECT transform(array(1, 2, 3), x -> x + 1); -- returns [2,3,4] /// ``` + /// + /// This enables both the `->` spelling above and the `LAMBDA` keyword + /// spelling gated by [`Self::supports_lambda_keyword_syntax`]. A dialect + /// that uses `->` as a binary operator should override only the latter. fn supports_lambda_functions(&self) -> bool { false } + /// Returns true if the dialect supports the `LAMBDA` keyword spelling of + /// lambda functions, for example: + /// + /// ```sql + /// SELECT list_transform([1, 2, 3], lambda x : x + 1); -- returns [2, 3, 4] + /// ``` + /// + /// This spelling does not claim the `->` token, so it can be enabled by + /// dialects that already give `->` a different meaning, such as PostgreSQL + /// and its derivatives, where `->` is JSON member access. Defaults to + /// [`Self::supports_lambda_functions`], so dialects supporting the `->` + /// spelling accept the `LAMBDA` spelling too unless they say otherwise. + /// + /// See + fn supports_lambda_keyword_syntax(&self) -> bool { + self.supports_lambda_functions() + } + /// Returns true if the dialect supports multiple variable assignment /// using parentheses in a `SET` variable declaration. /// diff --git a/src/parser/mod.rs b/src/parser/mod.rs index 6af0fb776..5b621f9a9 100644 --- a/src/parser/mod.rs +++ b/src/parser/mod.rs @@ -1649,7 +1649,7 @@ impl<'a> Parser<'a> { Keyword::MAP if *self.peek_token_ref() == Token::LBrace && self.dialect.support_map_literal_syntax() => { Ok(Some(self.parse_duckdb_map_literal()?)) } - Keyword::LAMBDA if self.dialect.supports_lambda_functions() => { + Keyword::LAMBDA if self.dialect.supports_lambda_keyword_syntax() => { Ok(Some(self.parse_lambda_expr()?)) } _ if self.dialect.supports_geometric_types() => match w.keyword { diff --git a/tests/sqlparser_custom_dialect.rs b/tests/sqlparser_custom_dialect.rs index cee604aca..5bf38ddec 100644 --- a/tests/sqlparser_custom_dialect.rs +++ b/tests/sqlparser_custom_dialect.rs @@ -22,6 +22,7 @@ use sqlparser::{ dialect::Dialect, keywords::Keyword, parser::{Parser, ParserError}, + test_utils::{expr_from_projection, only}, tokenizer::Token, }; @@ -167,3 +168,119 @@ fn is_identifier_part(ch: char) -> bool { || ch == '$' || ch == '_' } + +#[test] +fn custom_dialect_lambda_keyword_syntax_without_arrow() { + // A dialect that gives `->` its own meaning can still support lambdas + // through the `LAMBDA` keyword spelling. + #[derive(Debug)] + struct MyDialect {} + + impl Dialect for MyDialect { + fn is_identifier_start(&self, ch: char) -> bool { + is_identifier_start(ch) + } + + fn is_identifier_part(&self, ch: char) -> bool { + is_identifier_part(ch) + } + + fn supports_lambda_keyword_syntax(&self) -> bool { + true + } + } + + let dialect = MyDialect {}; + + // The `LAMBDA` spelling parses. + let sql = "SELECT transform(xs, lambda x : x + 1)"; + assert_eq!( + sql, + &format!("{}", Parser::parse_sql(&dialect, sql).unwrap()[0]) + ); + + // `->` keeps whatever meaning the dialect gives it, rather than + // introducing a lambda parameter. + let sql = "SELECT a -> 'b'"; + let ast = Parser::parse_sql(&dialect, sql).unwrap(); + match &ast[0] { + Statement::Query(query) => { + let Expr::BinaryOp { op, .. } = + expr_from_projection(only(&query.body.as_select().unwrap().projection)) + else { + panic!("expected `->` to stay a binary operator"); + }; + assert_eq!(&BinaryOperator::Arrow, op); + } + stmt => panic!("unexpected statement {stmt}"), + } +} + +#[test] +fn custom_dialect_lambda_keyword_defaults_to_arrow_support() { + // Dialects that opt into the `->` spelling get the `LAMBDA` spelling too, + // so the new capability does not change any existing dialect. + #[derive(Debug)] + struct MyDialect {} + + impl Dialect for MyDialect { + fn is_identifier_start(&self, ch: char) -> bool { + is_identifier_start(ch) + } + + fn is_identifier_part(&self, ch: char) -> bool { + is_identifier_part(ch) + } + + fn supports_lambda_functions(&self) -> bool { + true + } + } + + let dialect = MyDialect {}; + assert!(dialect.supports_lambda_keyword_syntax()); + for sql in [ + "SELECT transform(xs, lambda x : x + 1)", + "SELECT transform(xs, x -> x + 1)", + ] { + assert_eq!( + sql, + &format!("{}", Parser::parse_sql(&dialect, sql).unwrap()[0]) + ); + } +} + +#[test] +fn custom_dialect_lambda_arrow_syntax_without_keyword() { + // Arrow lambdas stay on while the `LAMBDA` keyword spelling is off, + // as in engines like Spark and Snowflake. + #[derive(Debug)] + struct MyDialect {} + + impl Dialect for MyDialect { + fn is_identifier_start(&self, ch: char) -> bool { + is_identifier_start(ch) + } + + fn is_identifier_part(&self, ch: char) -> bool { + is_identifier_part(ch) + } + + fn supports_lambda_functions(&self) -> bool { + true + } + + fn supports_lambda_keyword_syntax(&self) -> bool { + false + } + } + + let dialect = MyDialect {}; + + let sql = "SELECT transform(xs, x -> x + 1)"; + assert_eq!( + sql, + &format!("{}", Parser::parse_sql(&dialect, sql).unwrap()[0]) + ); + assert!(Parser::parse_sql(&dialect, "SELECT transform(xs, lambda x : x + 1)").is_err()); +} diff --git a/tests/sqlparser_derive_dialect.rs b/tests/sqlparser_derive_dialect.rs index d60fa1e11..cc7508bf1 100644 --- a/tests/sqlparser_derive_dialect.rs +++ b/tests/sqlparser_derive_dialect.rs @@ -17,9 +17,13 @@ //! Tests for the `derive_dialect!` macro. +use sqlparser::ast::{ + BinaryOperator, Expr, FunctionArg, FunctionArgExpr, FunctionArguments, LambdaSyntax, Statement, +}; use sqlparser::derive_dialect; use sqlparser::dialect::{Dialect, GenericDialect, MySqlDialect, PostgreSqlDialect}; use sqlparser::parser::Parser; +use sqlparser::test_utils::{expr_from_projection, only}; #[test] fn test_method_overrides() { @@ -121,3 +125,65 @@ fn test_identifier_quote_style_overrides() { None ); } + +#[test] +fn test_lambda_keyword_syntax_on_postgres_derivative() { + // A PostgreSQL derivative can opt into the `LAMBDA` keyword spelling of + // lambda functions without giving up `->` as JSON member access. The two + // meet in a single expression below: a lambda whose body is a JSON access. + derive_dialect!( + LambdaPostgreSqlDialect, + PostgreSqlDialect, + overrides = { supports_lambda_keyword_syntax = true } + ); + let dialect = LambdaPostgreSqlDialect::new(); + + // Only the keyword spelling is enabled; the arrow spelling stays off. + assert!(dialect.supports_lambda_keyword_syntax()); + assert!(!dialect.supports_lambda_functions()); + + let sql = "SELECT transform(xs, lambda x : (x -> 'a')::INT + 1)"; + let ast = Parser::parse_sql(&dialect, sql).unwrap(); + assert_eq!(sql, ast[0].to_string()); + + // Round-tripping alone would not distinguish a JSON access from a nested + // lambda, since both print as `x -> 'a'`, so check the parsed shape. + let Statement::Query(query) = &ast[0] else { + panic!("unexpected statement {}", ast[0]); + }; + let Expr::Function(func) = + expr_from_projection(only(&query.body.as_select().unwrap().projection)) + else { + panic!("expected a function call"); + }; + let FunctionArguments::List(args) = &func.args else { + panic!("expected an argument list"); + }; + let [_, FunctionArg::Unnamed(FunctionArgExpr::Expr(Expr::Lambda(lambda)))] = &args.args[..] + else { + panic!("expected the second argument to be a lambda"); + }; + + // The lambda came from the `LAMBDA` keyword, not from `->`. + assert_eq!(LambdaSyntax::LambdaKeyword, lambda.syntax); + + // And the `->` in its body is still JSON member access. + let Expr::BinaryOp { + left, + op: BinaryOperator::Plus, + .. + } = lambda.body.as_ref() + else { + panic!("expected the lambda body to be an addition"); + }; + let Expr::Cast { expr, .. } = left.as_ref() else { + panic!("expected the left operand to be a cast"); + }; + let Expr::Nested(json_access) = expr.as_ref() else { + panic!("expected the cast operand to be parenthesized"); + }; + let Expr::BinaryOp { op, .. } = json_access.as_ref() else { + panic!("expected `->` to stay a binary operator"); + }; + assert_eq!(&BinaryOperator::Arrow, op); +}