rust_store_core/dialect/
raw.rs1use std::collections::HashSet;
13
14use serde_json::Value;
15
16use super::Backend;
17
18#[derive(Debug)]
20pub struct RawStmt {
21 pub sql: String,
23 pub params: Vec<Value>,
25 pub is_write: bool,
27}
28
29pub fn compile_raw_stmt(
35 backend: Backend,
36 text: &str,
37 params: Value,
38 is_write: Option<bool>,
39) -> Result<RawStmt, String> {
40 let trimmed = text.trim();
41 if trimmed.is_empty() {
42 return Err("原生 SQL 文本为空".to_string());
43 }
44 let is_write = is_write.unwrap_or_else(|| infer_is_write(trimmed));
45 match params {
46 Value::Array(items) => Ok(RawStmt {
48 sql: text.to_string(),
49 params: items,
50 is_write,
51 }),
52 Value::Null => Ok(RawStmt {
53 sql: text.to_string(),
54 params: Vec::new(),
55 is_write,
56 }),
57 Value::Object(names) => {
59 let (sql, ordered) = compile_named(backend, text, &names)?;
60 Ok(RawStmt {
61 sql,
62 params: ordered,
63 is_write,
64 })
65 }
66 other => Err(format!(
67 "原生 SQL params 仅支持数组(位置档)或对象(命名档),收到 {}",
68 json_type_name(&other)
69 )),
70 }
71}
72
73fn compile_named(
79 backend: Backend,
80 text: &str,
81 params: &serde_json::Map<String, Value>,
82) -> Result<(String, Vec<Value>), String> {
83 let chars: Vec<char> = text.chars().collect();
84 let n = chars.len();
85 let mut out = String::with_capacity(text.len() + 8);
86 let mut ordered: Vec<Value> = Vec::new();
87 let mut used: HashSet<String> = HashSet::new();
88 let mut i = 0;
89 while i < n {
90 let c = chars[i];
91 if c == '\'' || c == '"' {
93 let quote = c;
94 out.push(c);
95 i += 1;
96 while i < n {
97 if chars[i] == quote {
98 if i + 1 < n && chars[i + 1] == quote {
99 out.push(quote);
101 out.push(quote);
102 i += 2;
103 continue;
104 }
105 break;
106 }
107 out.push(chars[i]);
108 i += 1;
109 }
110 if i < n {
111 out.push(quote); i += 1;
113 }
114 continue;
115 }
116 if c == '-' && i + 1 < n && chars[i + 1] == '-' {
118 while i < n && chars[i] != '\n' {
119 out.push(chars[i]);
120 i += 1;
121 }
122 continue;
123 }
124 if c == '/' && i + 1 < n && chars[i + 1] == '*' {
126 out.push(c);
127 out.push('*');
128 i += 2;
129 while i < n {
130 if chars[i] == '*' && i + 1 < n && chars[i + 1] == '/' {
131 out.push('*');
132 out.push('/');
133 i += 2;
134 break;
135 }
136 out.push(chars[i]);
137 i += 1;
138 }
139 continue;
140 }
141 if c == ':' && i + 1 < n && chars[i + 1] == ':' {
143 out.push(':');
144 out.push(':');
145 i += 2;
146 continue;
147 }
148 if c == ':' && i + 1 < n && (chars[i + 1].is_ascii_alphabetic() || chars[i + 1] == '_') {
150 let start = i + 1;
151 let mut j = start;
152 while j < n && (chars[j].is_ascii_alphanumeric() || chars[j] == '_') {
153 j += 1;
154 }
155 let name: String = chars[start..j].iter().collect();
156 match params.get(&name) {
157 Some(v) => {
158 out.push_str(&backend.placeholder(ordered.len()));
159 ordered.push(v.clone());
160 used.insert(name);
161 }
162 None => return Err(format!("原生 SQL 命名参数 :{} 未在 params 中提供", name)),
163 }
164 i = j;
165 continue;
166 }
167 out.push(c);
168 i += 1;
169 }
170 let mut unused: Vec<String> = params
172 .keys()
173 .filter(|k| !used.contains(*k))
174 .map(|k| format!(":{}", k))
175 .collect();
176 if !unused.is_empty() {
177 unused.sort();
178 return Err(format!(
179 "原生 SQL params 中存在未使用的命名参数: {}",
180 unused.join(", ")
181 ));
182 }
183 Ok((out, ordered))
184}
185
186fn infer_is_write(trimmed: &str) -> bool {
189 let mut s = trimmed.trim_start();
190 if let Some(rest) = s.strip_prefix('(') {
191 s = rest.trim_start();
192 }
193 let word: String = s
194 .chars()
195 .take_while(|c| c.is_ascii_alphanumeric() || *c == '_')
196 .collect();
197 !matches!(
198 word.to_ascii_uppercase().as_str(),
199 "SELECT" | "WITH" | "EXPLAIN" | "SHOW" | "PRAGMA" | "TABLE"
200 )
201}
202
203fn json_type_name(v: &Value) -> &'static str {
204 match v {
205 Value::Null => "null",
206 Value::Bool(_) => "bool",
207 Value::Number(_) => "number",
208 Value::String(_) => "string",
209 Value::Array(_) => "array",
210 Value::Object(_) => "object",
211 }
212}