valence_backend_sql/
query.rs1use serde_json::{Map, Value};
4use valence_core::compiled_query::CompiledQuery;
5use valence_core::error::{Error, Result};
6
7pub 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
36pub 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
42pub 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
67pub 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
78pub 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
89pub 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
101pub 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
110pub 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
154pub 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
161pub 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 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, ¶ms);
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}