Skip to main content

dbkit_core/
compile.rs

1use crate::expr::{BinaryOp, BoolOp, ExprNode, IntervalField, TrimDirection, 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                // Quoted identifiers may legally contain `$1` text. Treat doubled quotes as
71                // escaped identifier content and avoid placeholder scanning until the closing `"`.
72                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                // Scan bytewise and only interpret ASCII placeholder syntax (`$` + digits).
84                // Everything else is copied through verbatim below as UTF-8 string slices.
85                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                    // Rebind only the placeholder token. Any suffix text such as `::vector`,
98                    // `::interval`, or `::schema.enum_type` remains in `compiled.sql` and is
99                    // copied verbatim by the fallback branch after this placeholder is emitted.
100                    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                if (*name == "CONCAT" && args.is_empty()) || (*name == "CONCAT_WS" && args.len() == 1) {
154                    if !args.is_empty() {
155                        builder.push_sql(", ");
156                    }
157                    builder.push_sql("VARIADIC ARRAY[]::TEXT[]");
158                }
159                builder.push_sql(")");
160            }
161            ExprNode::Normalize { expr, form } => {
162                builder.push_sql("NORMALIZE(");
163                expr.to_sql(builder);
164                builder.push_sql(", ");
165                builder.push_sql(match form {
166                    crate::func::NormalizationForm::Nfc => "NFC",
167                    crate::func::NormalizationForm::Nfd => "NFD",
168                    crate::func::NormalizationForm::Nfkc => "NFKC",
169                    crate::func::NormalizationForm::Nfkd => "NFKD",
170                });
171                builder.push_sql(")");
172            }
173            ExprNode::Trim {
174                direction,
175                expr,
176                characters,
177            } => {
178                builder.push_sql("TRIM(");
179                builder.push_sql(match direction {
180                    TrimDirection::Both => "BOTH",
181                    TrimDirection::Leading => "LEADING",
182                    TrimDirection::Trailing => "TRAILING",
183                });
184                if let Some(characters) = characters {
185                    builder.push_sql(" ");
186                    characters.to_sql(builder);
187                }
188                builder.push_sql(" FROM ");
189                expr.to_sql(builder);
190                builder.push_sql(")");
191            }
192            ExprNode::AggregateFilter { aggregate, predicate } => {
193                aggregate.to_sql(builder);
194                builder.push_sql(" FILTER (WHERE ");
195                predicate.to_sql(builder);
196                builder.push_sql(")");
197            }
198            ExprNode::VectorBinary { left, op, right } => {
199                builder.push_sql("(");
200                left.to_sql(builder);
201                builder.push_sql(match op {
202                    VectorBinaryOp::L2Distance => " <-> ",
203                    VectorBinaryOp::CosineDistance => " <=> ",
204                    VectorBinaryOp::InnerProductDistance => " <#> ",
205                    VectorBinaryOp::L1Distance => " <+> ",
206                });
207                right.to_sql(builder);
208                builder.push_sql(")");
209            }
210            ExprNode::MakeInterval { field, value } => {
211                builder.push_sql("MAKE_INTERVAL(");
212                builder.push_sql(match field {
213                    IntervalField::Days => "days => ",
214                    IntervalField::Hours => "hours => ",
215                    IntervalField::Minutes => "mins => ",
216                    IntervalField::Seconds => "secs => ",
217                });
218                value.to_sql(builder);
219                builder.push_sql(")");
220            }
221            ExprNode::Binary { left, op, right } => {
222                builder.push_sql("(");
223                left.to_sql(builder);
224                builder.push_sql(match op {
225                    BinaryOp::Add => " + ",
226                    BinaryOp::Sub => " - ",
227                    BinaryOp::Mul => " * ",
228                    BinaryOp::Eq => " = ",
229                    BinaryOp::Ne => " <> ",
230                    BinaryOp::IsDistinctFrom => " IS DISTINCT FROM ",
231                    BinaryOp::IsNotDistinctFrom => " IS NOT DISTINCT FROM ",
232                    BinaryOp::Lt => " < ",
233                    BinaryOp::Le => " <= ",
234                    BinaryOp::Gt => " > ",
235                    BinaryOp::Ge => " >= ",
236                });
237                right.to_sql(builder);
238                builder.push_sql(")");
239            }
240            ExprNode::Bool { left, op, right } => {
241                builder.push_sql("(");
242                left.to_sql(builder);
243                builder.push_sql(match op {
244                    BoolOp::And => " AND ",
245                    BoolOp::Or => " OR ",
246                });
247                right.to_sql(builder);
248                builder.push_sql(")");
249            }
250            ExprNode::Unary { op, expr } => {
251                builder.push_sql(match op {
252                    UnaryOp::Not => "NOT ",
253                });
254                builder.push_sql("(");
255                expr.to_sql(builder);
256                builder.push_sql(")");
257            }
258            ExprNode::In { expr, values } => {
259                if values.is_empty() {
260                    builder.push_sql("(FALSE)");
261                    return;
262                }
263                builder.push_sql("(");
264                expr.to_sql(builder);
265                builder.push_sql(" IN (");
266                for (idx, value) in values.iter().enumerate() {
267                    if idx > 0 {
268                        builder.push_sql(", ");
269                    }
270                    builder.push_value(value.clone());
271                }
272                builder.push_sql("))");
273            }
274            ExprNode::RowIn { expr, rows } => {
275                if rows.is_empty() {
276                    builder.push_sql("(FALSE)");
277                    return;
278                }
279                builder.push_sql("(");
280                expr.to_sql(builder);
281                builder.push_sql(" IN (");
282                for (row_idx, row) in rows.iter().enumerate() {
283                    if row_idx > 0 {
284                        builder.push_sql(", ");
285                    }
286                    builder.push_sql("(");
287                    for (value_idx, value) in row.iter().enumerate() {
288                        if value_idx > 0 {
289                            builder.push_sql(", ");
290                        }
291                        builder.push_value(value.clone());
292                    }
293                    builder.push_sql(")");
294                }
295                builder.push_sql("))");
296            }
297            ExprNode::IsNull { expr, negated } => {
298                builder.push_sql("(");
299                expr.to_sql(builder);
300                if *negated {
301                    builder.push_sql(" IS NOT NULL)");
302                } else {
303                    builder.push_sql(" IS NULL)");
304                }
305            }
306            ExprNode::Like {
307                expr,
308                pattern,
309                case_insensitive,
310            } => {
311                builder.push_sql("(");
312                expr.to_sql(builder);
313                builder.push_sql(if *case_insensitive { " ILIKE " } else { " LIKE " });
314                builder.push_value(pattern.clone());
315                builder.push_sql(")");
316            }
317            ExprNode::Exists { subquery } => {
318                builder.push_sql("EXISTS (");
319                builder.push_compiled_sql(subquery);
320                builder.push_sql(")");
321            }
322        }
323    }
324}