Skip to main content

lance_datafusion/
sql.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright The Lance Authors
3
4//! SQL Parser utility
5
6use std::any::TypeId;
7
8use datafusion::sql::sqlparser::{
9    ast::{Expr, SelectItem, SetExpr, Statement},
10    dialect::{Dialect, GenericDialect},
11    parser::Parser,
12    tokenizer::{Token, Tokenizer},
13};
14
15use lance_core::{Error, Result};
16#[derive(Debug, Default)]
17struct LanceDialect(GenericDialect);
18
19impl LanceDialect {
20    fn new() -> Self {
21        Self(GenericDialect {})
22    }
23}
24
25impl Dialect for LanceDialect {
26    fn dialect(&self) -> TypeId {
27        self.0.dialect()
28    }
29
30    fn is_identifier_start(&self, ch: char) -> bool {
31        self.0.is_identifier_start(ch)
32    }
33
34    fn is_identifier_part(&self, ch: char) -> bool {
35        self.0.is_identifier_part(ch)
36    }
37
38    fn is_delimited_identifier_start(&self, ch: char) -> bool {
39        ch == '`'
40    }
41
42    fn supports_bitwise_shift_operators(&self) -> bool {
43        self.0.supports_bitwise_shift_operators()
44    }
45}
46
47/// Parse sql filter to Expression.
48pub(crate) fn parse_sql_filter(filter: &str) -> Result<Expr> {
49    let sql = format!("SELECT 1 FROM t WHERE {filter}");
50    let statement = parse_statement(&sql)?;
51
52    if let Statement::Query(query) = statement
53        && let SetExpr::Select(select) = *query.body
54        && let Some(expr) = select.selection
55    {
56        Ok(expr)
57    } else {
58        Err(Error::invalid_input(format!(
59            "Filter is not valid: {filter}"
60        )))
61    }
62}
63
64/// Parse a SQL expression to Expression. This is more lenient than parse_sql_filter
65/// as it can be used for projection expressions as well.
66pub(crate) fn parse_sql_expr(expr: &str) -> Result<Expr> {
67    let sql = format!("SELECT {expr} FROM t");
68    let statement = parse_statement(&sql)?;
69
70    if let Statement::Query(query) = statement
71        && let SetExpr::Select(select) = *query.body
72        && let Some(SelectItem::UnnamedExpr(expr)) = select.projection.into_iter().next()
73    {
74        Ok(expr)
75    } else {
76        Err(Error::invalid_input(format!(
77            "Expression is not valid: {expr}"
78        )))
79    }
80}
81
82fn parse_statement(statement: &str) -> Result<Statement> {
83    let dialect = LanceDialect::new();
84
85    // Hack to allow == as equals
86    // This is used to parse PyArrow expressions from strings.
87    // See: https://github.com/sqlparser-rs/sqlparser-rs/pull/815#issuecomment-1450714278
88    let mut tokenizer = Tokenizer::new(&dialect, statement);
89    let mut tokens = Vec::new();
90    let mut token_iter = tokenizer
91        .tokenize()
92        .map_err(|e| {
93            Error::invalid_input(format!("Error tokenizing statement: {statement} ({e})"))
94        })?
95        .into_iter();
96    let mut prev_token = token_iter.next().unwrap();
97    for next_token in token_iter {
98        if let (Token::Eq, Token::Eq) = (&prev_token, &next_token) {
99            continue; // skip second equals
100        }
101        let token = std::mem::replace(&mut prev_token, next_token);
102        tokens.push(token);
103    }
104    tokens.push(prev_token);
105
106    Parser::new(&dialect)
107        .with_tokens(tokens)
108        .parse_statement()
109        .map_err(|e| Error::invalid_input(format!("Error parsing statement: {statement} ({e})")))
110}
111
112#[cfg(test)]
113mod tests {
114    use super::*;
115
116    use datafusion::sql::sqlparser::{
117        ast::{BinaryOperator, Ident, Value, ValueWithSpan},
118        tokenizer::Span,
119    };
120
121    #[test]
122    fn test_double_equal() {
123        let expr = parse_sql_filter("a == b").unwrap();
124        assert_eq!(
125            Expr::BinaryOp {
126                left: Box::new(Expr::Identifier(Ident::new("a"))),
127                op: BinaryOperator::Eq,
128                right: Box::new(Expr::Identifier(Ident::new("b")))
129            },
130            expr
131        );
132    }
133
134    #[test]
135    fn test_like() {
136        let expr = parse_sql_filter("a LIKE 'abc%'").unwrap();
137        assert_eq!(
138            Expr::Like {
139                negated: false,
140                expr: Box::new(Expr::Identifier(Ident::new("a"))),
141                pattern: Box::new(Expr::Value(ValueWithSpan {
142                    value: Value::SingleQuotedString("abc%".to_string()),
143                    span: Span::empty(),
144                })),
145                escape_char: None,
146                any: false,
147            },
148            expr
149        );
150    }
151
152    #[test]
153    fn test_quoted_ident() {
154        // CUBE is a SQL keyword, so it must be quoted.
155        let expr = parse_sql_filter("`a:Test_Something` == `CUBE`").unwrap();
156        assert_eq!(
157            Expr::BinaryOp {
158                left: Box::new(Expr::Identifier(Ident::with_quote('`', "a:Test_Something"))),
159                op: BinaryOperator::Eq,
160                right: Box::new(Expr::Identifier(Ident::with_quote('`', "CUBE")))
161            },
162            expr
163        );
164
165        let expr = parse_sql_filter("`outer field`.`inner field` == 1").unwrap();
166        assert_eq!(
167            Expr::BinaryOp {
168                left: Box::new(Expr::CompoundIdentifier(vec![
169                    Ident::with_quote('`', "outer field"),
170                    Ident::with_quote('`', "inner field")
171                ])),
172                op: BinaryOperator::Eq,
173                right: Box::new(Expr::Value(ValueWithSpan {
174                    value: Value::Number("1".to_string(), false),
175                    span: Span::empty(),
176                })),
177            },
178            expr
179        );
180    }
181}