diff --git a/derive/src/dialect.rs b/derive/src/dialect.rs index 9066bf964..0b5322021 100644 --- a/derive/src/dialect.rs +++ b/derive/src/dialect.rs @@ -23,7 +23,7 @@ use std::collections::HashSet; use syn::{ braced, parse::{Parse, ParseStream}, - Error, File, FnArg, Ident, Item, LitBool, LitChar, Pat, ReturnType, Signature, Token, + Error, Expr, File, FnArg, Ident, Item, LitBool, LitChar, Pat, ReturnType, Signature, Token, TraitItem, Type, }; @@ -31,7 +31,7 @@ use syn::{ pub(crate) enum Override { Bool(LitBool), Char(LitChar), - None, + Expr(Expr), } /// Parsed input for the `derive_dialect!` macro @@ -80,20 +80,8 @@ impl Parse for DeriveDialectInput { Override::Bool(content.parse()?) } else if content.peek(LitChar) { Override::Char(content.parse()?) - } else if content.peek(Ident) { - let ident: Ident = content.parse()?; - if ident == "None" { - Override::None - } else { - return Err(Error::new( - ident.span(), - format!("Expected `true`, `false`, a char, or `None`, found `{ident}`"), - )); - } } else { - return Err( - content.error("Expected `true`, `false`, a char, or `None`") - ); + Override::Expr(content.parse()?) }; overrides.push((key, value)); if content.peek(Token![,]) { @@ -136,6 +124,7 @@ fn derive_dialect_inner(input: DeriveDialectInput) -> syn::Result { let methods = extract_dialect_methods(&file)?; // Validate overrides + let method_names: HashSet<_> = methods.iter().map(|m| m.name.to_string()).collect(); let bool_names: HashSet<_> = methods .iter() .filter(|m| is_bool_method(&m.signature)) @@ -143,6 +132,12 @@ fn derive_dialect_inner(input: DeriveDialectInput) -> syn::Result { .collect(); for (key, value) in &input.overrides { let key_str = key.to_string(); + if !method_names.contains(&key_str) { + return Err(Error::new( + key.span(), + format!("Unknown method `{key_str}`"), + )); + } match value { Override::Bool(_) if !bool_names.contains(&key_str) => { return Err(Error::new( @@ -150,10 +145,10 @@ fn derive_dialect_inner(input: DeriveDialectInput) -> syn::Result { format!("Unknown boolean method `{key_str}`"), )); } - Override::Char(_) | Override::None if key_str != "identifier_quote_style" => { + Override::Char(_) if key_str != "identifier_quote_style" => { return Err(Error::new( key.span(), - format!("Char/None only valid for `identifier_quote_style`, not `{key_str}`"), + format!("Char only valid for `identifier_quote_style`, not `{key_str}`"), )); } _ => {} @@ -214,10 +209,9 @@ fn generate_derived_dialect(input: &DeriveDialectInput, methods: &[DialectMethod fn identifier_quote_style(&self, _: &str) -> Option { Some(#c) } } } - Some(Override::None) => { - quote_spanned! { method_name.span() => - fn identifier_quote_style(&self, _: &str) -> Option { None } - } + Some(Override::Expr(expr)) => { + let sig = &method.signature; + quote_spanned! { method_name.span() => #sig { #expr } } } None => delegate(method), } @@ -230,7 +224,7 @@ fn generate_derived_dialect(input: &DeriveDialectInput, methods: &[DialectMethod use ::core::iter::Peekable; use ::core::str::Chars; use sqlparser::ast::{ColumnOption, Expr, GranteesType, Ident, ObjectNamePart, Statement}; - use sqlparser::dialect::{Dialect, Precedence}; + use sqlparser::dialect::{BindPlaceholderStyle, Dialect, Precedence}; use sqlparser::keywords::Keyword; use sqlparser::parser::{Parser, ParserError}; diff --git a/src/dialect/mod.rs b/src/dialect/mod.rs index f99cbe2ea..102e805fe 100644 --- a/src/dialect/mod.rs +++ b/src/dialect/mod.rs @@ -167,6 +167,17 @@ macro_rules! dialect_is { } } +/// Placeholder spelling for ordered unnamed bind parameters. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)] +#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] +#[non_exhaustive] +pub enum BindPlaceholderStyle { + /// `?` + QuestionMark, + /// `$1`, `$2`, ... + DollarNumbered, +} + /// Encapsulates the differences between SQL implementations. /// /// # SQL Dialects @@ -254,6 +265,15 @@ pub trait Dialect: Debug + Any { None } + /// Return the placeholder syntax this server accepts for ordered unnamed bind parameters. + /// + /// This describes the server contract rather than every placeholder the parser can tokenize. + /// `None` means this dialect does not specify an ordered bind syntax. + /// See [`Dialect::supports_dollar_placeholder`] for SQLite-style named `$name` placeholders. + fn ordered_bind_placeholder_style(&self) -> Option { + None + } + /// Determine if a character is a valid start character for an unquoted identifier fn is_identifier_start(&self, ch: char) -> bool; @@ -1107,8 +1127,11 @@ pub trait Dialect: Debug + Any { false } - /// Returns true if this dialect allows dollar placeholders - /// e.g. `SELECT $var` (SQLite) + /// Returns true if this dialect allows SQLite-style named dollar placeholders, + /// for example `SELECT $name`. + /// + /// This does not describe PostgreSQL-style ordered binds such as `SELECT $1`. + /// See [`Dialect::ordered_bind_placeholder_style`]. fn supports_dollar_placeholder(&self) -> bool { false } @@ -1996,6 +2019,7 @@ mod tests { supports_order_by_all = true, supports_nested_comments = true, supports_triple_quoted_string = true, + ordered_bind_placeholder_style = Some(BindPlaceholderStyle::DollarNumbered), }, ); let dialect = EnhancedGenericDialect::new(); @@ -2003,6 +2027,10 @@ mod tests { assert!(dialect.supports_order_by_all()); assert!(dialect.supports_nested_comments()); assert!(dialect.supports_triple_quoted_string()); + assert_eq!( + dialect.ordered_bind_placeholder_style(), + Some(BindPlaceholderStyle::DollarNumbered) + ); let d: &dyn Dialect = &dialect; assert!(d.is::()); @@ -2036,6 +2064,23 @@ mod tests { } } + #[test] + fn ordered_bind_placeholder_style() { + let tests: Vec<(&dyn Dialect, Option)> = vec![ + (&GenericDialect {}, None), + (&MySqlDialect {}, Some(BindPlaceholderStyle::QuestionMark)), + ( + &PostgreSqlDialect {}, + Some(BindPlaceholderStyle::DollarNumbered), + ), + (&SQLiteDialect {}, Some(BindPlaceholderStyle::QuestionMark)), + ]; + + for (dialect, expected) in tests { + assert_eq!(dialect.ordered_bind_placeholder_style(), expected); + } + } + #[test] fn parse_with_wrapped_dialect() { /// Wrapper for a dialect. In a real-world example, this wrapper @@ -2072,6 +2117,10 @@ mod tests { self.0.identifier_quote_style(identifier) } + fn ordered_bind_placeholder_style(&self) -> Option { + self.0.ordered_bind_placeholder_style() + } + fn supports_string_literal_backslash_escape(&self) -> bool { self.0.supports_string_literal_backslash_escape() } @@ -2139,6 +2188,10 @@ mod tests { let statement = r#"SELECT 'Wayne\'s World'"#; let res1 = Parser::parse_sql(&MySqlDialect {}, statement); let res2 = Parser::parse_sql(&WrappedDialect(MySqlDialect {}), statement); + assert_eq!( + WrappedDialect(MySqlDialect {}).ordered_bind_placeholder_style(), + Some(BindPlaceholderStyle::QuestionMark) + ); assert!(res1.is_ok()); assert_eq!(res1, res2); } diff --git a/src/dialect/mysql.rs b/src/dialect/mysql.rs index f73219e8b..bd524fef4 100644 --- a/src/dialect/mysql.rs +++ b/src/dialect/mysql.rs @@ -20,7 +20,7 @@ use alloc::boxed::Box; use crate::{ ast::{BinaryOperator, Expr, LockTable, LockTableType, Statement}, - dialect::Dialect, + dialect::{BindPlaceholderStyle, Dialect}, keywords::Keyword, parser::{Parser, ParserError}, }; @@ -67,6 +67,10 @@ impl Dialect for MySqlDialect { Some('`') } + fn ordered_bind_placeholder_style(&self) -> Option { + Some(BindPlaceholderStyle::QuestionMark) + } + // See https://dev.mysql.com/doc/refman/8.0/en/string-literals.html#character-escape-sequences fn supports_string_literal_backslash_escape(&self) -> bool { true diff --git a/src/dialect/postgresql.rs b/src/dialect/postgresql.rs index d342276e4..949a51316 100644 --- a/src/dialect/postgresql.rs +++ b/src/dialect/postgresql.rs @@ -28,7 +28,7 @@ // limitations under the License. use log::debug; -use crate::dialect::{Dialect, Precedence}; +use crate::dialect::{BindPlaceholderStyle, Dialect, Precedence}; use crate::keywords::Keyword; use crate::parser::{Parser, ParserError}; use crate::tokenizer::Token; @@ -63,6 +63,10 @@ impl Dialect for PostgreSqlDialect { Some('"') } + fn ordered_bind_placeholder_style(&self) -> Option { + Some(BindPlaceholderStyle::DollarNumbered) + } + fn is_delimited_identifier_start(&self, ch: char) -> bool { ch == '"' // Postgres does not support backticks to quote identifiers } diff --git a/src/dialect/sqlite.rs b/src/dialect/sqlite.rs index d549c7507..a11b02d30 100644 --- a/src/dialect/sqlite.rs +++ b/src/dialect/sqlite.rs @@ -20,7 +20,7 @@ use alloc::boxed::Box; use crate::ast::BinaryOperator; use crate::ast::{Expr, Statement}; -use crate::dialect::Dialect; +use crate::dialect::{BindPlaceholderStyle, Dialect}; use crate::keywords::Keyword; use crate::parser::{Parser, ParserError}; @@ -46,6 +46,10 @@ impl Dialect for SQLiteDialect { Some('`') } + fn ordered_bind_placeholder_style(&self) -> Option { + Some(BindPlaceholderStyle::QuestionMark) + } + fn is_identifier_start(&self, ch: char) -> bool { // See https://www.sqlite.org/draft/tokenreq.html ch.is_ascii_lowercase()