|
1 | 1 | use crate::query::{Expression, BinaryOperator, LogicalOperator, LogicalPlan, Statement}; |
2 | 2 | use serde_json::Value; |
3 | | -use sqlparser::dialect::GenericDialect; |
| 3 | +use sqlparser::dialect::Dialect; |
4 | 4 | use sqlparser::parser::Parser; |
5 | | -use sqlparser::ast::{self, SetExpr, TableFactor, BinaryOperator as SqlBinaryOperator, Expr, LimitClause}; |
6 | | -use nom::{ |
7 | | - bytes::complete::{tag_no_case, take_while1}, |
8 | | - character::complete::{multispace0, multispace1, char}, |
9 | | - sequence::tuple, |
10 | | - multi::separated_list1, |
11 | | - IResult, |
12 | | -}; |
| 5 | +use sqlparser::ast::{self, SetExpr, TableFactor, BinaryOperator as SqlBinaryOperator, Expr, LimitClause, Values}; |
13 | 6 |
|
14 | | -pub fn parse(sql: &str) -> Result<Statement, String> { |
15 | | - let trimmed = sql.trim(); |
16 | | - if trimmed.to_uppercase().starts_with("INSERT") { |
17 | | - parse_insert(trimmed).map_err(|e| format!("Insert parse error: {}", e)) |
18 | | - } else { |
19 | | - parse_select(trimmed) |
20 | | - } |
21 | | -} |
22 | | - |
23 | | -// --- INSERT Parsing (Custom using Nom) --- |
| 7 | +#[derive(Debug)] |
| 8 | +struct ArgusDialect; |
24 | 9 |
|
25 | | -fn parse_insert(input: &str) -> Result<Statement, String> { |
26 | | - match insert_statement(input) { |
27 | | - Ok((_, stmt)) => Ok(stmt), |
28 | | - Err(e) => Err(format!("{}", e)), |
| 10 | +impl Dialect for ArgusDialect { |
| 11 | + fn is_identifier_start(&self, ch: char) -> bool { |
| 12 | + (ch >= 'a' && ch <= 'z') || (ch >= 'A' && ch <= 'Z') || ch == '_' |
29 | 13 | } |
30 | | -} |
31 | 14 |
|
32 | | -fn insert_statement(input: &str) -> IResult<&str, Statement> { |
33 | | - // INSERT INTO <collection> RECORDS <json>... |
34 | | - let (input, _) = tag_no_case("INSERT")(input)?; |
35 | | - let (input, _) = multispace1(input)?; |
36 | | - let (input, _) = tag_no_case("INTO")(input)?; |
37 | | - let (input, _) = multispace1(input)?; |
38 | | - |
39 | | - let (input, collection) = identifier(input)?; |
40 | | - |
41 | | - let (input, _) = multispace1(input)?; |
42 | | - let (input, _) = tag_no_case("RECORDS")(input)?; |
43 | | - let (input, _) = multispace0(input)?; |
44 | | - |
45 | | - let (input, documents) = separated_list1( |
46 | | - tuple((multispace0, char(','), multispace0)), |
47 | | - json_object |
48 | | - )(input)?; |
49 | | - |
50 | | - Ok((input, Statement::Insert { |
51 | | - collection: collection.to_string(), |
52 | | - documents, |
53 | | - })) |
54 | | -} |
55 | | - |
56 | | -fn identifier(input: &str) -> IResult<&str, &str> { |
57 | | - take_while1(|c: char| c.is_alphanumeric() || c == '_')(input) |
58 | | -} |
59 | | - |
60 | | -fn json_object(input: &str) -> IResult<&str, Value> { |
61 | | - let mut depth = 0; |
62 | | - let mut len = 0; |
63 | | - let mut found_start = false; |
64 | | - |
65 | | - // Skip leading whitespace |
66 | | - let leading_ws = input.chars().take_while(|c| c.is_whitespace()).count(); |
67 | | - let trimmed_input = &input[leading_ws..]; |
68 | | - |
69 | | - if !trimmed_input.starts_with('{') { |
70 | | - return Err(nom::Err::Error(nom::error::Error::new(input, nom::error::ErrorKind::Tag))); |
| 15 | + fn is_identifier_part(&self, ch: char) -> bool { |
| 16 | + (ch >= 'a' && ch <= 'z') |
| 17 | + || (ch >= 'A' && ch <= 'Z') |
| 18 | + || (ch >= '0' && ch <= '9') |
| 19 | + || ch == '_' |
71 | 20 | } |
72 | 21 |
|
73 | | - // Iterate to find the matching brace |
74 | | - for (i, c) in trimmed_input.char_indices() { |
75 | | - if c == '{' { |
76 | | - depth += 1; |
77 | | - found_start = true; |
78 | | - } else if c == '}' { |
79 | | - depth -= 1; |
80 | | - } |
81 | | - |
82 | | - if found_start && depth == 0 { |
83 | | - len = i + 1; |
84 | | - break; |
85 | | - } |
| 22 | + fn is_delimited_identifier_start(&self, ch: char) -> bool { |
| 23 | + ch == '`' |
86 | 24 | } |
87 | | - |
88 | | - if depth != 0 { |
89 | | - return Err(nom::Err::Error(nom::error::Error::new(input, nom::error::ErrorKind::Complete))); |
90 | | - } |
91 | | - |
92 | | - let json_str = &trimmed_input[0..len]; |
93 | | - let value: Value = serde_json::from_str(json_str).map_err(|_| { |
94 | | - nom::Err::Error(nom::error::Error::new(input, nom::error::ErrorKind::MapRes)) |
95 | | - })?; |
96 | | - |
97 | | - Ok((&trimmed_input[len..], value)) |
98 | 25 | } |
99 | 26 |
|
| 27 | +pub fn parse(sql: &str) -> Result<Statement, String> { |
| 28 | + let dialect = ArgusDialect {}; |
| 29 | + |
| 30 | + // Hack: Replace RECORDS with VALUES to satisfy sqlparser |
| 31 | + // We strictly assume "RECORDS" is used for INSERT. |
| 32 | + let sql_to_parse = if sql.trim().to_uppercase().starts_with("INSERT") { |
| 33 | + sql.replacen("RECORDS", "VALUES", 1).replacen("records", "VALUES", 1) |
| 34 | + } else { |
| 35 | + sql.to_string() |
| 36 | + }; |
100 | 37 |
|
101 | | -// --- SELECT Parsing (sqlparser) --- |
102 | | - |
103 | | -fn parse_select(sql: &str) -> Result<Statement, String> { |
104 | | - let dialect = GenericDialect {}; |
105 | | - let ast = Parser::parse_sql(&dialect, sql).map_err(|e| e.to_string())?; |
| 38 | + let ast = Parser::parse_sql(&dialect, &sql_to_parse).map_err(|e| e.to_string())?; |
106 | 39 |
|
107 | 40 | if ast.len() != 1 { |
108 | 41 | return Err("Expected exactly one statement".to_string()); |
109 | 42 | } |
110 | 43 |
|
111 | 44 | match &ast[0] { |
| 45 | + ast::Statement::Insert { table_name, source, .. } => { |
| 46 | + let collection = table_name.to_string(); |
| 47 | + let documents = convert_insert_source(source)?; |
| 48 | + Ok(Statement::Insert { collection, documents }) |
| 49 | + } |
112 | 50 | ast::Statement::Query(query) => { |
113 | 51 | let logical_plan = convert_query(query)?; |
114 | 52 | Ok(Statement::Select(logical_plan)) |
115 | 53 | } |
116 | | - _ => Err("Only SELECT statements are supported".to_string()), |
| 54 | + _ => Err("Only SELECT and INSERT statements are supported".to_string()), |
| 55 | + } |
| 56 | +} |
| 57 | + |
| 58 | +fn convert_insert_source(source: &Option<Box<ast::Query>>) -> Result<Vec<Value>, String> { |
| 59 | + let query = source.as_ref().ok_or("Insert must have a source")?; |
| 60 | + |
| 61 | + match &*query.body { |
| 62 | + SetExpr::Values(Values { rows, .. }) => { |
| 63 | + let mut docs = Vec::new(); |
| 64 | + for row in rows { |
| 65 | + if row.len() != 1 { |
| 66 | + return Err("Each record must contain exactly one JSON object".to_string()); |
| 67 | + } |
| 68 | + let expr = &row[0]; |
| 69 | + match expr { |
| 70 | + Expr::Identifier(ident) => { |
| 71 | + // We expect a backtick-quoted identifier which contains the JSON |
| 72 | + let json_str = &ident.value; |
| 73 | + let value: Value = serde_json::from_str(json_str).map_err(|e| format!("Invalid JSON in INSERT: {}", e))?; |
| 74 | + docs.push(value); |
| 75 | + } |
| 76 | + _ => return Err("Expected a JSON object enclosed in backticks".to_string()), |
| 77 | + } |
| 78 | + } |
| 79 | + Ok(docs) |
| 80 | + } |
| 81 | + _ => Err("INSERT expects VALUES (RECORDS) clause".to_string()), |
117 | 82 | } |
118 | 83 | } |
119 | 84 |
|
@@ -198,7 +163,6 @@ fn convert_select(select: &ast::Select) -> Result<LogicalPlan, String> { |
198 | 163 | projections.push(convert_expr(expr)?); |
199 | 164 | } |
200 | 165 | ast::SelectItem::ExprWithAlias { expr, alias: _ } => { |
201 | | - // Ignore alias for now as LogicalPlan doesn't support renaming explicitly yet |
202 | 166 | projections.push(convert_expr(expr)?); |
203 | 167 | } |
204 | 168 | ast::SelectItem::Wildcard(_) => { |
@@ -280,7 +244,7 @@ mod tests { |
280 | 244 |
|
281 | 245 | #[test] |
282 | 246 | fn test_parse_insert() { |
283 | | - let sql = r#"INSERT INTO users RECORDS {"name": "Alice", "age": 30}, {"name": "Bob"}"#; |
| 247 | + let sql = r#"INSERT INTO users RECORDS `{"name": "Alice", "age": 30}`, `{"name": "Bob"}`"#; |
284 | 248 | let stmt = parse(sql).unwrap(); |
285 | 249 | match stmt { |
286 | 250 | Statement::Insert { collection, documents } => { |
|
0 commit comments