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