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
43pub(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
60pub(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 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; }
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 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}