Skip to main content

valence_backend_sql/
query.rs

1//! Execute compiled SQL queries and map rows to JSON.
2
3use serde_json::{Map, Value};
4use valence_core::compiled_query::CompiledQuery;
5use valence_core::error::{Error, Result};
6
7/// Bind `$param_key` placeholders in SQL to `?` for SQLite positional binding.
8pub fn sql_with_positional_placeholders(
9    query: &str,
10    params: &[(String, Value)],
11) -> (String, Vec<Value>) {
12    let mut out = String::with_capacity(query.len());
13    let mut values = Vec::new();
14    let mut rest = query;
15    while let Some(dollar) = rest.find('$') {
16        out.push_str(&rest[..dollar]);
17        rest = &rest[dollar + 1..];
18        let key_len = rest
19            .chars()
20            .take_while(|c| c.is_ascii_alphanumeric() || *c == '_')
21            .count();
22        let key = &rest[..key_len];
23        rest = &rest[key_len..];
24        if let Some((_, value)) = params.iter().find(|(k, _)| k == key) {
25            out.push('?');
26            values.push(value.clone());
27        } else {
28            out.push('$');
29            out.push_str(key);
30        }
31    }
32    out.push_str(rest);
33    (out, values)
34}
35
36/// Decode a SQL row `(id, body)` into Valence JSON record shape.
37pub fn row_to_json(table: &str, id: &str, body_text: &str) -> Result<Value> {
38    let body: Value = serde_json::from_str(body_text).unwrap_or_else(|_| Value::Object(Map::new()));
39    Ok(super::document::row_from_body(table, id, body))
40}
41
42/// Parse SELECT results from generic JSON rows returned by driver layer.
43pub fn decode_select_rows(rows: Vec<Value>, default_table: &str) -> Result<Vec<Value>> {
44    let mut out = Vec::new();
45    for row in rows {
46        if let Some(obj) = row.as_object() {
47            if let (Some(id), Some(body)) = (obj.get("id"), obj.get("body")) {
48                let id_str = id.as_str().unwrap_or_default();
49                let body_val = if let Some(s) = body.as_str() {
50                    serde_json::from_str(s).unwrap_or(Value::Object(Map::new()))
51                } else {
52                    body.clone()
53                };
54                out.push(super::document::row_from_body(
55                    default_table,
56                    id_str,
57                    body_val,
58                ));
59                continue;
60            }
61        }
62        out.push(row);
63    }
64    Ok(out)
65}
66
67/// Extract count from first row.
68pub fn first_count(rows: &[Value]) -> i64 {
69    rows.first()
70        .and_then(|v| {
71            v.as_i64()
72                .or_else(|| v.get("count").and_then(|c| c.as_i64()))
73                .or_else(|| v.as_f64().map(|f| f as i64))
74        })
75        .unwrap_or(0)
76}
77
78/// Extract id strings from SELECT id queries.
79pub fn extract_ids(rows: &[Value]) -> Vec<String> {
80    rows.iter()
81        .filter_map(|v| {
82            v.as_str()
83                .map(str::to_string)
84                .or_else(|| v.get("id").and_then(|id| id.as_str().map(str::to_string)))
85        })
86        .collect()
87}
88
89/// Validate compiled query is read-only SELECT.
90pub fn ensure_read_only(query: &str) -> Result<()> {
91    let upper = query.trim().to_uppercase();
92    if upper.starts_with("SELECT ") {
93        Ok(())
94    } else {
95        Err(Error::Internal(format!(
96            "unsupported SQL in execute_compiled_query: {query}"
97        )))
98    }
99}
100
101/// Normalize compiled query for SQLite execution (`?` placeholders).
102pub fn prepare_compiled(compiled: &CompiledQuery) -> Result<(String, Vec<Value>)> {
103    ensure_read_only(&compiled.query_string)?;
104    Ok(sql_with_positional_placeholders(
105        &compiled.query_string,
106        &compiled.params,
107    ))
108}
109
110/// Translate SQLite-style `json_extract(expr, '$.a.b')` into Postgres jsonb operators.
111pub fn rewrite_json_extract_for_postgres(sql: &str) -> String {
112    let mut out = String::with_capacity(sql.len());
113    let mut rest = sql;
114    while let Some(start) = rest.find("json_extract(") {
115        out.push_str(&rest[..start]);
116        rest = &rest[start + "json_extract(".len()..];
117        let Some(comma) = rest.find(',') else {
118            out.push_str("json_extract(");
119            break;
120        };
121        let expr = rest[..comma].trim();
122        rest = rest[comma + 1..].trim_start();
123        let path = if let Some(stripped) = rest.strip_prefix("'$.") {
124            let Some(end_q) = stripped.find('\'') else {
125                out.push_str("json_extract(");
126                out.push_str(expr);
127                out.push_str(", ");
128                break;
129            };
130            let path = &stripped[..end_q];
131            rest = stripped[end_q + 1..].trim_start();
132            if let Some(r) = rest.strip_prefix(')') {
133                rest = r;
134            }
135            path
136        } else {
137            out.push_str("json_extract(");
138            out.push_str(expr);
139            out.push_str(", ");
140            continue;
141        };
142        let parts: Vec<&str> = path.split('.').filter(|p| !p.is_empty()).collect();
143        if parts.len() <= 1 {
144            let seg = parts.first().copied().unwrap_or("");
145            out.push_str(&format!("({expr}->>'{seg}')"));
146        } else {
147            out.push_str(&format!("({expr}#>>'{{{}}}')", parts.join(",")));
148        }
149    }
150    out.push_str(rest);
151    out
152}
153
154/// Normalize compiled query for Postgres execution (`$1`, `$2`, … placeholders).
155pub fn prepare_compiled_postgres(compiled: &CompiledQuery) -> Result<(String, Vec<Value>)> {
156    ensure_read_only(&compiled.query_string)?;
157    let rewritten = rewrite_json_extract_for_postgres(&compiled.query_string);
158    Ok(sql_with_postgres_placeholders(&rewritten, &compiled.params))
159}
160
161/// Bind `$param_key` placeholders to Postgres numbered params.
162pub fn sql_with_postgres_placeholders(
163    query: &str,
164    params: &[(String, Value)],
165) -> (String, Vec<Value>) {
166    let mut out = String::with_capacity(query.len());
167    let mut values = Vec::new();
168    let mut rest = query;
169    let mut idx = 1usize;
170    while let Some(dollar) = rest.find('$') {
171        out.push_str(&rest[..dollar]);
172        rest = &rest[dollar + 1..];
173        let key_len = rest
174            .chars()
175            .take_while(|c| c.is_ascii_alphanumeric() || *c == '_')
176            .count();
177        let key = &rest[..key_len];
178        rest = &rest[key_len..];
179        if let Some((_, value)) = params.iter().find(|(k, _)| k == key) {
180            out.push_str(&format!("${idx}"));
181            values.push(value.clone());
182            idx += 1;
183        } else if key.chars().all(|c| c.is_ascii_digit()) {
184            // Already positional ($1) — keep as-is.
185            out.push('$');
186            out.push_str(key);
187        } else {
188            out.push('$');
189            out.push_str(key);
190        }
191    }
192    out.push_str(rest);
193    (out, values)
194}
195
196#[cfg(test)]
197mod tests {
198    use super::*;
199    use serde_json::json;
200
201    #[test]
202    fn positional_preserves_json_path_and_binds_params() {
203        let q = "SELECT id, body FROM task WHERE (json_extract(body, '$.project') = $param_0 OR json_extract(body, '$.project') = $param_1 OR json_extract(body, '$.project.id') = $param_1)";
204        let params = vec![
205            ("param_0".into(), json!("project:mem-1")),
206            ("param_1".into(), json!("mem-1")),
207        ];
208        let (sql, values) = sql_with_positional_placeholders(q, &params);
209        assert_eq!(
210            sql,
211            "SELECT id, body FROM task WHERE (json_extract(body, '$.project') = ? OR json_extract(body, '$.project') = ? OR json_extract(body, '$.project.id') = ?)"
212        );
213        assert_eq!(values.len(), 3);
214        assert_eq!(values[0], json!("project:mem-1"));
215        assert_eq!(values[1], json!("mem-1"));
216        assert_eq!(values[2], json!("mem-1"));
217    }
218
219    #[test]
220    fn postgres_rewrites_json_extract_paths() {
221        let q = "SELECT id FROM task WHERE json_extract(body, '$.project') = $p OR json_extract(t.body, '$.project.id') = $p";
222        let out = rewrite_json_extract_for_postgres(q);
223        assert_eq!(
224            out,
225            "SELECT id FROM task WHERE (body->>'project') = $p OR (t.body#>>'{project,id}') = $p"
226        );
227    }
228
229    #[test]
230    fn prepare_compiled_postgres_rewrites_json_extract() {
231        let compiled = CompiledQuery {
232            query_string: "SELECT id, body FROM project WHERE json_extract(body, '$.name') = $param_0 LIMIT 10".into(),
233            params: vec![("param_0".into(), json!("alpha"))],
234        };
235        let (sql, values) = prepare_compiled_postgres(&compiled).expect("prepare");
236        assert!(
237            !sql.contains("json_extract"),
238            "raw json_extract must be rewritten: {sql}"
239        );
240        assert!(sql.contains("body->>'name'") || sql.contains("(body->>'name')"));
241        assert_eq!(values, vec![json!("alpha")]);
242    }
243}