Skip to main content

rust_store_core/dialect/
raw.rs

1//! 原生 SQL 语句编译:命名占位符(`:name`)→ 方言占位符 + 读写推断(纯逻辑,无 IO)。
2//!
3//! 为宿主 `execute_raw` 的「两档」能力提供 core 侧唯一实现:
4//! - **位置档**(params 为数组/null):SQL 原样透传,占位符由调用方手写方言原生风格
5//!   (对标 SQLAlchemy `exec_driver_sql()`);
6//! - **命名档**(params 为对象):`:name` 按出现顺序编译为 [`Backend::placeholder`],
7//!   参数按引用顺序重排、同名复用(对标 SQLAlchemy `text()`)。
8//!
9//! 跳过边界(R4)保证字符串/注释/PG cast 内的冒号不被误编译;读写推断(R7)以
10//! 「默认写」为安全方向——误判为写至多路由主库,误判为读会读到旧数据。
11
12use std::collections::HashSet;
13
14use serde_json::Value;
15
16use super::Backend;
17
18/// 编译产物:最终 SQL + 重排后参数 + 读写标记
19#[derive(Debug)]
20pub struct RawStmt {
21    /// 位置档 = 原文;命名档 = 编译后文本
22    pub sql: String,
23    /// 位置档 = 原数组(顺序保持);命名档 = 按 `:name` 出现顺序重排
24    pub params: Vec<Value>,
25    /// 显式指定优先;否则按首词推断(R7)
26    pub is_write: bool,
27}
28
29/// 原生 SQL 语句编译。
30///
31/// - `params`:`Array` = 位置档(透传);`Object` = 命名档(`:name` 编译);
32///   `Null` = 位置档空参;其余 → Err(R1)。
33/// - `is_write`:`Some(b)` 显式采用;`None` 按首词推断(R7)。
34pub 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        // R2 位置档:原样透传(Null 视为空参数组)
47        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        // R3–R6 命名档:`:name` 编译 + 参数重排 + 一致性校验
58        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
73/// R3–R6:命名档扫描编译。
74///
75/// 状态机逐字符扫描:引号字符串 / 行注释 / 块注释 / `::` cast 内部原样复制(R4);
76/// 裸 `?` / `$n` 不在编译范围,保持原样(PG `jsonb ? 'k'` 等操作符合法,误用时由驱动
77/// 报参数数错,不静默)。缺名 / 多余名显式 Err(R6,禁静默)。
78fn 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        // R4:单引号 / 双引号字符串('' / "" 翻倍转义)
92        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                        // 翻倍转义:两个字面引号,仍在字符串内
100                        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); // 收尾引号(未闭合引号原样保留,交由驱动报错)
112                i += 1;
113            }
114            continue;
115        }
116        // R4:`-- …` 行注释(至行尾)
117        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        // R4:`/* … */` 块注释
125        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        // R4:`::` cast(PG),双冒号一并跳过
142        if c == ':' && i + 1 < n && chars[i + 1] == ':' {
143            out.push(':');
144            out.push(':');
145            i += 2;
146            continue;
147        }
148        // R3:`:name` 命名占位符(首字符字母/下划线,续字母/数字/下划线)
149        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    // R6:params 中 text 未使用的名字 → Err
171    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
186/// R7:读写推断。trim 后跳过一层前导 `(` 与空白,取首词大写化;
187/// ∈ {SELECT, WITH, EXPLAIN, SHOW, PRAGMA, TABLE} → 读,其余 → 写(安全方向)。
188fn 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}