Skip to main content

oxilite_core/
sql.rs

1//! The sans-IO contract between oxilite and a SQLite-compatible engine.
2//!
3//! oxilite never talks to SQLite directly. Every operation produces [`Request`]s (lists of SQL
4//! statements) and consumes [`Response`]s. A backend only has to run statements — natively
5//! through rusqlite or a dlopen'ed `libsqlite3`, or remotely through Cloudflare D1's
6//! `batch()` API.
7//!
8// @lat: [[architecture#Sans-IO core]]
9
10use crate::error::{Error, Result};
11use std::fmt::Write;
12
13/// A SQL value, the subset of SQLite storage classes oxilite uses.
14///
15/// In JSON (JavaScript drivers, D1 over HTTP) it is `null`, a number or a string. The
16/// deserializer is written by hand: an untagged derive breaks when another crate in the build
17/// enables serde_json's `arbitrary_precision` (as the `ssi` crates behind `oxilite-vc` do).
18#[derive(Debug, Clone, PartialEq)]
19#[cfg_attr(feature = "serde", derive(serde::Serialize))]
20#[cfg_attr(feature = "serde", serde(untagged))]
21pub enum SqlValue {
22    Null,
23    Integer(i64),
24    Real(f64),
25    Text(String),
26}
27
28#[cfg(feature = "serde")]
29impl<'de> serde::Deserialize<'de> for SqlValue {
30    fn deserialize<D: serde::Deserializer<'de>>(d: D) -> std::result::Result<Self, D::Error> {
31        struct V;
32        impl<'de> serde::de::Visitor<'de> for V {
33            type Value = SqlValue;
34            fn expecting(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
35                f.write_str("null, a number or a string")
36            }
37            fn visit_unit<E>(self) -> std::result::Result<SqlValue, E> {
38                Ok(SqlValue::Null)
39            }
40            fn visit_none<E>(self) -> std::result::Result<SqlValue, E> {
41                Ok(SqlValue::Null)
42            }
43            fn visit_some<D: serde::Deserializer<'de>>(
44                self,
45                d: D,
46            ) -> std::result::Result<SqlValue, D::Error> {
47                d.deserialize_any(self)
48            }
49            fn visit_bool<E>(self, v: bool) -> std::result::Result<SqlValue, E> {
50                Ok(SqlValue::Integer(i64::from(v)))
51            }
52            fn visit_i64<E>(self, v: i64) -> std::result::Result<SqlValue, E> {
53                Ok(SqlValue::Integer(v))
54            }
55            fn visit_u64<E: serde::de::Error>(self, v: u64) -> std::result::Result<SqlValue, E> {
56                i64::try_from(v)
57                    .map(SqlValue::Integer)
58                    .or(Ok(SqlValue::Real(v as f64)))
59            }
60            fn visit_f64<E>(self, v: f64) -> std::result::Result<SqlValue, E> {
61                Ok(SqlValue::Real(v))
62            }
63            fn visit_str<E>(self, v: &str) -> std::result::Result<SqlValue, E> {
64                Ok(SqlValue::Text(v.to_owned()))
65            }
66            fn visit_string<E>(self, v: String) -> std::result::Result<SqlValue, E> {
67                Ok(SqlValue::Text(v))
68            }
69            // serde_json with `arbitrary_precision` hands numbers over as a one-entry map.
70            fn visit_map<A: serde::de::MapAccess<'de>>(
71                self,
72                mut map: A,
73            ) -> std::result::Result<SqlValue, A::Error> {
74                let Some((_, n)) = map.next_entry::<String, String>()? else {
75                    return Err(serde::de::Error::custom("empty map is not a SQL value"));
76                };
77                if let Ok(i) = n.parse::<i64>() {
78                    return Ok(SqlValue::Integer(i));
79                }
80                n.parse::<f64>()
81                    .map(SqlValue::Real)
82                    .map_err(serde::de::Error::custom)
83            }
84        }
85        d.deserialize_any(V)
86    }
87}
88
89impl SqlValue {
90    pub fn as_i64(&self) -> Option<i64> {
91        match self {
92            Self::Integer(v) => Some(*v),
93            // Backends that cannot transport 64-bit integers (D1 → JavaScript numbers) receive
94            // ids as TEXT (see `Capabilities::int64_as_text`).
95            Self::Text(v) => v.parse().ok(),
96            Self::Real(v) if v.fract() == 0.0 && v.abs() < 9.0e15 => Some(*v as i64),
97            _ => None,
98        }
99    }
100
101    pub fn as_f64(&self) -> Option<f64> {
102        match self {
103            Self::Integer(v) => Some(*v as f64),
104            Self::Real(v) => Some(*v),
105            Self::Text(v) => v.parse().ok(),
106            Self::Null => None,
107        }
108    }
109
110    pub fn as_str(&self) -> Option<&str> {
111        match self {
112            Self::Text(v) => Some(v),
113            _ => None,
114        }
115    }
116
117    pub fn into_string(self) -> Option<String> {
118        match self {
119            Self::Text(v) => Some(v),
120            Self::Integer(v) => Some(v.to_string()),
121            Self::Real(v) => Some(v.to_string()),
122            Self::Null => None,
123        }
124    }
125
126    pub fn is_null(&self) -> bool {
127        matches!(self, Self::Null)
128    }
129}
130
131/// A single SQL statement. oxilite inlines all constants as SQL literals, so `params` is
132/// almost always empty; this sidesteps D1's 100-bound-parameter limit.
133#[derive(Debug, Clone, PartialEq)]
134#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
135pub struct Statement {
136    pub sql: String,
137    #[cfg_attr(
138        feature = "serde",
139        serde(default, skip_serializing_if = "Vec::is_empty")
140    )]
141    pub params: Vec<SqlValue>,
142}
143
144impl Statement {
145    pub fn new(sql: impl Into<String>) -> Self {
146        Self {
147            sql: sql.into(),
148            params: Vec::new(),
149        }
150    }
151}
152
153impl From<String> for Statement {
154    fn from(sql: String) -> Self {
155        Self::new(sql)
156    }
157}
158
159impl From<&str> for Statement {
160    fn from(sql: &str) -> Self {
161        Self::new(sql)
162    }
163}
164
165/// How a request must be executed.
166#[derive(Debug, Clone, Copy, PartialEq, Eq)]
167#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
168#[cfg_attr(feature = "serde", serde(rename_all = "lowercase"))]
169pub enum Mode {
170    /// Read-only statements; may run outside a transaction.
171    Read,
172    /// All statements succeed or none does (one SQLite transaction, one D1 `batch()`).
173    Atomic,
174}
175
176/// A group of statements sent to the backend in one round-trip.
177#[derive(Debug, Clone, PartialEq)]
178#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
179pub struct Request {
180    pub statements: Vec<Statement>,
181    pub mode: Mode,
182}
183
184impl Request {
185    pub fn read(statements: Vec<Statement>) -> Self {
186        Self {
187            statements,
188            mode: Mode::Read,
189        }
190    }
191
192    pub fn atomic(statements: Vec<Statement>) -> Self {
193        Self {
194            statements,
195            mode: Mode::Atomic,
196        }
197    }
198}
199
200/// The result of one statement.
201#[derive(Debug, Clone, Default, PartialEq)]
202#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
203pub struct ResultSet {
204    #[cfg_attr(feature = "serde", serde(default))]
205    pub rows: Vec<Vec<SqlValue>>,
206    /// Rows modified by a write statement.
207    #[cfg_attr(feature = "serde", serde(default))]
208    pub changes: u64,
209}
210
211/// One [`ResultSet`] per statement of the [`Request`], in order.
212pub type Response = Vec<ResultSet>;
213
214/// What a backend can do. The compiler adapts SQL generation to it.
215#[derive(Debug, Clone, PartialEq)]
216#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
217#[cfg_attr(feature = "serde", serde(default, rename_all = "camelCase"))]
218pub struct Capabilities {
219    /// Maximum length of one SQL statement in bytes (D1: 100 KB).
220    pub max_sql_len: usize,
221    /// Maximum number of statements in one request.
222    pub max_statements: usize,
223    /// The `oxilite_*` user-defined functions (regex, replace, hashes, Unicode case) exist.
224    pub udf: bool,
225    /// The backend supports interactive (read-then-write) transactions.
226    pub interactive_transactions: bool,
227    /// 64-bit integers must be returned as TEXT (JavaScript numbers lose precision above 2^53).
228    pub int64_as_text: bool,
229    /// Maximum number of terms in one compound SELECT (`UNION ALL` chain); D1 allows 5.
230    pub max_compound_select: usize,
231    /// Name of the backend, for `explain()`.
232    pub name: String,
233}
234
235impl Default for Capabilities {
236    fn default() -> Self {
237        Self::native()
238    }
239}
240
241impl Capabilities {
242    /// A native SQLite linked in-process.
243    pub fn native() -> Self {
244        Self {
245            max_sql_len: 1_000_000,
246            max_statements: 10_000,
247            udf: false,
248            interactive_transactions: true,
249            int64_as_text: false,
250            max_compound_select: 500,
251            name: "sqlite".into(),
252        }
253    }
254
255    /// Cloudflare D1 limits.
256    pub fn d1() -> Self {
257        Self {
258            max_sql_len: 90_000,
259            max_statements: 50,
260            udf: false,
261            interactive_transactions: false,
262            int64_as_text: true,
263            max_compound_select: 5,
264            name: "d1".into(),
265        }
266    }
267}
268
269/// Writes a SQL string literal (single quotes doubled; NUL-containing strings as hex blobs).
270pub fn quote_str(out: &mut String, s: &str) {
271    if s.contains('\0') {
272        out.push_str("CAST(X'");
273        for b in s.as_bytes() {
274            let _ = write!(out, "{b:02X}");
275        }
276        out.push_str("' AS TEXT)");
277        return;
278    }
279    out.push('\'');
280    for c in s.chars() {
281        if c == '\'' {
282            out.push('\'');
283        }
284        out.push(c);
285    }
286    out.push('\'');
287}
288
289/// Returns a SQL string literal.
290pub fn sql_str(s: &str) -> String {
291    let mut out = String::with_capacity(s.len() + 2);
292    quote_str(&mut out, s);
293    out
294}
295
296/// Formats an optional string as SQL.
297pub fn sql_opt_str(s: Option<&str>) -> String {
298    s.map_or_else(|| "NULL".into(), sql_str)
299}
300
301/// Formats a float as a SQL literal.
302pub fn sql_f64(v: f64) -> String {
303    if v.is_nan() {
304        "NULL".into()
305    } else if v == f64::INFINITY {
306        "9e999".into()
307    } else if v == f64::NEG_INFINITY {
308        "-9e999".into()
309    } else {
310        let s = format!("{v:?}");
311        if s.contains('.') || s.contains('e') || s.contains("inf") {
312            s
313        } else {
314            format!("{s}.0")
315        }
316    }
317}
318
319/// `UNION ALL` of several SELECTs within the backend's compound-select limit: longer chains
320/// are nested (`SELECT * FROM (a UNION ALL b …) UNION ALL …`).
321pub fn union_all(mut parts: Vec<String>, max_terms: usize) -> String {
322    let k = max_terms.max(2);
323    while parts.len() > k {
324        parts = parts
325            .chunks(k)
326            .map(|c| {
327                if c.len() == 1 {
328                    c[0].clone()
329                } else {
330                    format!("SELECT * FROM ({})", c.join(" UNION ALL "))
331                }
332            })
333            .collect();
334    }
335    parts.join(" UNION ALL ")
336}
337
338/// Reads a column from a row.
339pub fn col(row: &[SqlValue], i: usize) -> Result<&SqlValue> {
340    row.get(i)
341        .ok_or_else(|| Error::corrupted(format!("missing column {i} in result row")))
342}
343
344/// Checks that a response has the expected number of result sets.
345pub fn expect_len(response: &Response, n: usize) -> Result<()> {
346    if response.len() < n {
347        return Err(Error::backend(format!(
348            "backend returned {} result sets, expected {n}",
349            response.len()
350        )));
351    }
352    Ok(())
353}
354
355#[cfg(test)]
356mod tests {
357    use super::*;
358
359    #[test]
360    fn quoting() {
361        assert_eq!(sql_str("it's"), "'it''s'");
362        assert_eq!(sql_str("a\0b"), "CAST(X'610062' AS TEXT)");
363        assert_eq!(sql_f64(1.0), "1.0");
364        assert_eq!(sql_f64(1e300), "1e300");
365    }
366
367    // @lat: [[tests#D1#SQL values survive arbitrary-precision JSON]]
368    #[cfg(feature = "serde")]
369    #[test]
370    fn sql_values_survive_arbitrary_precision_json() {
371        let plain: Vec<SqlValue> = serde_json::from_str(r#"[null, 7, 1.5, "x"]"#).unwrap();
372        assert_eq!(
373            plain,
374            [
375                SqlValue::Null,
376                SqlValue::Integer(7),
377                SqlValue::Real(1.5),
378                SqlValue::Text("x".into())
379            ]
380        );
381        // The shape serde_json hands numbers over in when `arbitrary_precision` is enabled.
382        let token = r#"[{"$serde_json::private::Number": "1.5"}, {"$serde_json::private::Number": "9007199254740993"}]"#;
383        let v: Vec<SqlValue> = serde_json::from_str(token).unwrap();
384        assert_eq!(
385            v,
386            [
387                SqlValue::Real(1.5),
388                SqlValue::Integer(9_007_199_254_740_993)
389            ]
390        );
391    }
392}