1use 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
47pub(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
64pub(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 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; }
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 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}