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    /// A recursive CTE may have a compound recursive term (SQLite 3.34.0, 2020-12-01). A
232    /// backend whose version cannot be established MUST leave this false: mutual recursion
233    /// then takes a strategy that does not need it, instead of emitting SQL that would fail.
234    pub compound_recursive_cte: bool,
235    /// Vector types and distance functions (`vector32`, `vector_distance_cos`…; Turso): vector
236    /// indexes can be built and searched (see [`crate::vector`]).
237    pub vectors: bool,
238    /// Index methods (`CREATE INDEX … USING method`; Turso): sparse vector indexes get an
239    /// inverted-file index.
240    pub vector_index_methods: bool,
241    /// Name of the backend, for `explain()`.
242    pub name: String,
243    /// The versioning level of the store (set by the store at open, not by the backend):
244    /// from `stamped` on, writers record the current tick in `quads.t`.
245    pub versioning: crate::version::Versioning,
246}
247
248impl Default for Capabilities {
249    fn default() -> Self {
250        Self::native()
251    }
252}
253
254impl Capabilities {
255    /// A native SQLite linked in-process.
256    pub fn native() -> Self {
257        Self {
258            max_sql_len: 1_000_000,
259            max_statements: 10_000,
260            udf: false,
261            interactive_transactions: true,
262            int64_as_text: false,
263            max_compound_select: 500,
264            // rusqlite bundles a current SQLite, and a dlopen'ed library is probed on open.
265            compound_recursive_cte: true,
266            vectors: false,
267            vector_index_methods: false,
268            name: "sqlite".into(),
269            versioning: crate::version::Versioning::Off,
270        }
271    }
272
273    /// Cloudflare D1 limits.
274    pub fn d1() -> Self {
275        Self {
276            max_sql_len: 90_000,
277            max_statements: 50,
278            udf: false,
279            interactive_transactions: false,
280            int64_as_text: true,
281            max_compound_select: 5,
282            // D1's SQLite version is not ours to assume.
283            compound_recursive_cte: false,
284            vectors: false,
285            vector_index_methods: false,
286            name: "d1".into(),
287            versioning: crate::version::Versioning::Off,
288        }
289    }
290}
291
292/// Writes a SQL string literal (single quotes doubled; NUL-containing strings as hex blobs).
293pub fn quote_str(out: &mut String, s: &str) {
294    if s.contains('\0') {
295        out.push_str("CAST(X'");
296        for b in s.as_bytes() {
297            let _ = write!(out, "{b:02X}");
298        }
299        out.push_str("' AS TEXT)");
300        return;
301    }
302    out.push('\'');
303    for c in s.chars() {
304        if c == '\'' {
305            out.push('\'');
306        }
307        out.push(c);
308    }
309    out.push('\'');
310}
311
312/// Returns a SQL string literal.
313pub fn sql_str(s: &str) -> String {
314    let mut out = String::with_capacity(s.len() + 2);
315    quote_str(&mut out, s);
316    out
317}
318
319/// Formats an optional string as SQL.
320pub fn sql_opt_str(s: Option<&str>) -> String {
321    s.map_or_else(|| "NULL".into(), sql_str)
322}
323
324/// Formats a float as a SQL literal.
325pub fn sql_f64(v: f64) -> String {
326    if v.is_nan() {
327        "NULL".into()
328    } else if v == f64::INFINITY {
329        "9e999".into()
330    } else if v == f64::NEG_INFINITY {
331        "-9e999".into()
332    } else {
333        let s = format!("{v:?}");
334        if s.contains('.') || s.contains('e') || s.contains("inf") {
335            s
336        } else {
337            format!("{s}.0")
338        }
339    }
340}
341
342/// `UNION ALL` of several SELECTs within the backend's compound-select limit: longer chains
343/// are nested (`SELECT * FROM (a UNION ALL b …) UNION ALL …`).
344pub fn union_all(mut parts: Vec<String>, max_terms: usize) -> String {
345    let k = max_terms.max(2);
346    while parts.len() > k {
347        parts = parts
348            .chunks(k)
349            .map(|c| {
350                if c.len() == 1 {
351                    c[0].clone()
352                } else {
353                    format!("SELECT * FROM ({})", c.join(" UNION ALL "))
354                }
355            })
356            .collect();
357    }
358    parts.join(" UNION ALL ")
359}
360
361/// Reads a column from a row.
362pub fn col(row: &[SqlValue], i: usize) -> Result<&SqlValue> {
363    row.get(i)
364        .ok_or_else(|| Error::corrupted(format!("missing column {i} in result row")))
365}
366
367/// Checks that a response has the expected number of result sets.
368pub fn expect_len(response: &Response, n: usize) -> Result<()> {
369    if response.len() < n {
370        return Err(Error::backend(format!(
371            "backend returned {} result sets, expected {n}",
372            response.len()
373        )));
374    }
375    Ok(())
376}
377
378#[cfg(test)]
379mod tests {
380    use super::*;
381
382    #[test]
383    fn quoting() {
384        assert_eq!(sql_str("it's"), "'it''s'");
385        assert_eq!(sql_str("a\0b"), "CAST(X'610062' AS TEXT)");
386        assert_eq!(sql_f64(1.0), "1.0");
387        assert_eq!(sql_f64(1e300), "1e300");
388    }
389
390    // @lat: [[tests#D1#SQL values survive arbitrary-precision JSON]]
391    #[cfg(feature = "serde")]
392    #[test]
393    fn sql_values_survive_arbitrary_precision_json() {
394        let plain: Vec<SqlValue> = serde_json::from_str(r#"[null, 7, 1.5, "x"]"#).unwrap();
395        assert_eq!(
396            plain,
397            [
398                SqlValue::Null,
399                SqlValue::Integer(7),
400                SqlValue::Real(1.5),
401                SqlValue::Text("x".into())
402            ]
403        );
404        // The shape serde_json hands numbers over in when `arbitrary_precision` is enabled.
405        let token = r#"[{"$serde_json::private::Number": "1.5"}, {"$serde_json::private::Number": "9007199254740993"}]"#;
406        let v: Vec<SqlValue> = serde_json::from_str(token).unwrap();
407        assert_eq!(
408            v,
409            [
410                SqlValue::Real(1.5),
411                SqlValue::Integer(9_007_199_254_740_993)
412            ]
413        );
414    }
415}