1use crate::expr::{BinaryOp, BoolOp, ExprNode, IntervalField, UnaryOp, Value, VectorBinaryOp};
2use crate::schema::ColumnRef;
3
4#[derive(Debug, Clone, PartialEq)]
5pub struct CompiledSql {
6 pub sql: String,
7 pub binds: Vec<Value>,
8}
9
10#[derive(Debug, Default)]
11pub struct SqlBuilder {
12 sql: String,
13 binds: Vec<Value>,
14}
15
16impl SqlBuilder {
17 pub fn new() -> Self {
18 Self::default()
19 }
20
21 pub fn push_sql(&mut self, fragment: &str) {
22 self.sql.push_str(fragment);
23 }
24
25 fn push_placeholder(&mut self, value: Value) {
26 let idx = if let Some(existing) = self.binds.iter().position(|item| item == &value) {
27 existing + 1
28 } else {
29 self.binds.push(value);
30 self.binds.len()
31 };
32 self.sql.push('$');
33 self.sql.push_str(&idx.to_string());
34 }
35
36 pub fn push_value(&mut self, value: Value) {
37 if value == Value::Null {
38 self.sql.push_str("NULL");
39 return;
40 }
41 let cast_as_vector = matches!(&value, Value::Vector(_));
42 let cast_as_interval = matches!(&value, Value::Interval(_));
43 let cast_as_enum = match &value {
44 Value::Enum { type_name, .. } => Some(*type_name),
45 _ => None,
46 };
47 self.push_placeholder(value);
48 if cast_as_vector {
49 self.sql.push_str("::vector");
50 } else if cast_as_interval {
51 self.sql.push_str("::interval");
52 } else if let Some(type_name) = cast_as_enum {
53 self.sql.push_str("::");
54 self.sql.push_str(type_name);
55 }
56 }
57
58 pub fn push_column(&mut self, col: ColumnRef) {
59 self.sql.push_str(&col.qualified_name());
60 }
61
62 pub fn push_compiled_sql(&mut self, compiled: &CompiledSql) {
63 let bytes = compiled.sql.as_bytes();
64 let mut idx = 0;
65 let mut segment_start = 0;
66 let mut in_quoted_identifier = false;
67
68 while idx < bytes.len() {
69 if bytes[idx] == b'"' {
70 if in_quoted_identifier && idx + 1 < bytes.len() && bytes[idx + 1] == b'"' {
73 idx += 2;
74 continue;
75 }
76
77 in_quoted_identifier = !in_quoted_identifier;
78 idx += 1;
79 continue;
80 }
81
82 if !in_quoted_identifier && bytes[idx] == b'$' {
83 let prev_is_ident = idx > 0 && is_bind_ident_char(bytes[idx - 1]);
86 let start = idx + 1;
87 let mut end = start;
88 while end < bytes.len() && bytes[end].is_ascii_digit() {
89 end += 1;
90 }
91 let next_is_ident = end < bytes.len() && is_bind_ident_char(bytes[end]);
92
93 if end > start && !prev_is_ident && !next_is_ident {
94 self.push_sql(&compiled.sql[segment_start..idx]);
95 let bind_idx = compiled.sql[start..end].parse::<usize>().expect("valid bind index");
96 let value = compiled.binds[bind_idx - 1].clone();
97 self.push_placeholder(value);
101 idx = end;
102 segment_start = end;
103 continue;
104 }
105 }
106
107 idx += 1;
108 }
109
110 self.push_sql(&compiled.sql[segment_start..]);
111 }
112
113 pub fn finish(self) -> CompiledSql {
114 CompiledSql {
115 sql: self.sql,
116 binds: self.binds,
117 }
118 }
119}
120
121fn is_bind_ident_char(byte: u8) -> bool {
122 !byte.is_ascii() || byte.is_ascii_alphanumeric() || byte == b'_' || byte == b'$'
123}
124
125pub trait ToSql {
126 fn to_sql(&self, builder: &mut SqlBuilder);
127}
128
129impl ToSql for ExprNode {
130 fn to_sql(&self, builder: &mut SqlBuilder) {
131 match self {
132 ExprNode::Column(col) => builder.push_column(*col),
133 ExprNode::Value(value) => builder.push_value(value.clone()),
134 ExprNode::Row { values } => {
135 builder.push_sql("(");
136 for (idx, value) in values.iter().enumerate() {
137 if idx > 0 {
138 builder.push_sql(", ");
139 }
140 value.to_sql(builder);
141 }
142 builder.push_sql(")");
143 }
144 ExprNode::Func { name, args } => {
145 builder.push_sql(name);
146 builder.push_sql("(");
147 for (idx, arg) in args.iter().enumerate() {
148 if idx > 0 {
149 builder.push_sql(", ");
150 }
151 arg.to_sql(builder);
152 }
153 builder.push_sql(")");
154 }
155 ExprNode::VectorBinary { left, op, right } => {
156 builder.push_sql("(");
157 left.to_sql(builder);
158 builder.push_sql(match op {
159 VectorBinaryOp::L2Distance => " <-> ",
160 VectorBinaryOp::CosineDistance => " <=> ",
161 VectorBinaryOp::InnerProductDistance => " <#> ",
162 VectorBinaryOp::L1Distance => " <+> ",
163 });
164 right.to_sql(builder);
165 builder.push_sql(")");
166 }
167 ExprNode::MakeInterval { field, value } => {
168 builder.push_sql("MAKE_INTERVAL(");
169 builder.push_sql(match field {
170 IntervalField::Days => "days => ",
171 IntervalField::Hours => "hours => ",
172 IntervalField::Minutes => "mins => ",
173 IntervalField::Seconds => "secs => ",
174 });
175 value.to_sql(builder);
176 builder.push_sql(")");
177 }
178 ExprNode::Binary { left, op, right } => {
179 builder.push_sql("(");
180 left.to_sql(builder);
181 builder.push_sql(match op {
182 BinaryOp::Add => " + ",
183 BinaryOp::Sub => " - ",
184 BinaryOp::Mul => " * ",
185 BinaryOp::Eq => " = ",
186 BinaryOp::Ne => " <> ",
187 BinaryOp::IsDistinctFrom => " IS DISTINCT FROM ",
188 BinaryOp::IsNotDistinctFrom => " IS NOT DISTINCT FROM ",
189 BinaryOp::Lt => " < ",
190 BinaryOp::Le => " <= ",
191 BinaryOp::Gt => " > ",
192 BinaryOp::Ge => " >= ",
193 });
194 right.to_sql(builder);
195 builder.push_sql(")");
196 }
197 ExprNode::Bool { left, op, right } => {
198 builder.push_sql("(");
199 left.to_sql(builder);
200 builder.push_sql(match op {
201 BoolOp::And => " AND ",
202 BoolOp::Or => " OR ",
203 });
204 right.to_sql(builder);
205 builder.push_sql(")");
206 }
207 ExprNode::Unary { op, expr } => {
208 builder.push_sql(match op {
209 UnaryOp::Not => "NOT ",
210 });
211 builder.push_sql("(");
212 expr.to_sql(builder);
213 builder.push_sql(")");
214 }
215 ExprNode::In { expr, values } => {
216 if values.is_empty() {
217 builder.push_sql("(FALSE)");
218 return;
219 }
220 builder.push_sql("(");
221 expr.to_sql(builder);
222 builder.push_sql(" IN (");
223 for (idx, value) in values.iter().enumerate() {
224 if idx > 0 {
225 builder.push_sql(", ");
226 }
227 builder.push_value(value.clone());
228 }
229 builder.push_sql("))");
230 }
231 ExprNode::RowIn { expr, rows } => {
232 if rows.is_empty() {
233 builder.push_sql("(FALSE)");
234 return;
235 }
236 builder.push_sql("(");
237 expr.to_sql(builder);
238 builder.push_sql(" IN (");
239 for (row_idx, row) in rows.iter().enumerate() {
240 if row_idx > 0 {
241 builder.push_sql(", ");
242 }
243 builder.push_sql("(");
244 for (value_idx, value) in row.iter().enumerate() {
245 if value_idx > 0 {
246 builder.push_sql(", ");
247 }
248 builder.push_value(value.clone());
249 }
250 builder.push_sql(")");
251 }
252 builder.push_sql("))");
253 }
254 ExprNode::IsNull { expr, negated } => {
255 builder.push_sql("(");
256 expr.to_sql(builder);
257 if *negated {
258 builder.push_sql(" IS NOT NULL)");
259 } else {
260 builder.push_sql(" IS NULL)");
261 }
262 }
263 ExprNode::Like {
264 expr,
265 pattern,
266 case_insensitive,
267 } => {
268 builder.push_sql("(");
269 expr.to_sql(builder);
270 builder.push_sql(if *case_insensitive { " ILIKE " } else { " LIKE " });
271 builder.push_value(pattern.clone());
272 builder.push_sql(")");
273 }
274 ExprNode::Exists { subquery } => {
275 builder.push_sql("EXISTS (");
276 builder.push_compiled_sql(subquery);
277 builder.push_sql(")");
278 }
279 }
280 }
281}