Skip to main content

dbkit_core/
compile.rs

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                // 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                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}