1use crate::error::{Error, Result};
11use std::fmt::Write;
12
13#[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 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 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#[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#[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,
172 Atomic,
174}
175
176#[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#[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 #[cfg_attr(feature = "serde", serde(default))]
208 pub changes: u64,
209}
210
211pub type Response = Vec<ResultSet>;
213
214#[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 pub max_sql_len: usize,
221 pub max_statements: usize,
223 pub udf: bool,
225 pub interactive_transactions: bool,
227 pub int64_as_text: bool,
229 pub max_compound_select: usize,
231 pub compound_recursive_cte: bool,
235 pub name: String,
237}
238
239impl Default for Capabilities {
240 fn default() -> Self {
241 Self::native()
242 }
243}
244
245impl Capabilities {
246 pub fn native() -> Self {
248 Self {
249 max_sql_len: 1_000_000,
250 max_statements: 10_000,
251 udf: false,
252 interactive_transactions: true,
253 int64_as_text: false,
254 max_compound_select: 500,
255 compound_recursive_cte: true,
257 name: "sqlite".into(),
258 }
259 }
260
261 pub fn d1() -> Self {
263 Self {
264 max_sql_len: 90_000,
265 max_statements: 50,
266 udf: false,
267 interactive_transactions: false,
268 int64_as_text: true,
269 max_compound_select: 5,
270 compound_recursive_cte: false,
272 name: "d1".into(),
273 }
274 }
275}
276
277pub fn quote_str(out: &mut String, s: &str) {
279 if s.contains('\0') {
280 out.push_str("CAST(X'");
281 for b in s.as_bytes() {
282 let _ = write!(out, "{b:02X}");
283 }
284 out.push_str("' AS TEXT)");
285 return;
286 }
287 out.push('\'');
288 for c in s.chars() {
289 if c == '\'' {
290 out.push('\'');
291 }
292 out.push(c);
293 }
294 out.push('\'');
295}
296
297pub fn sql_str(s: &str) -> String {
299 let mut out = String::with_capacity(s.len() + 2);
300 quote_str(&mut out, s);
301 out
302}
303
304pub fn sql_opt_str(s: Option<&str>) -> String {
306 s.map_or_else(|| "NULL".into(), sql_str)
307}
308
309pub fn sql_f64(v: f64) -> String {
311 if v.is_nan() {
312 "NULL".into()
313 } else if v == f64::INFINITY {
314 "9e999".into()
315 } else if v == f64::NEG_INFINITY {
316 "-9e999".into()
317 } else {
318 let s = format!("{v:?}");
319 if s.contains('.') || s.contains('e') || s.contains("inf") {
320 s
321 } else {
322 format!("{s}.0")
323 }
324 }
325}
326
327pub fn union_all(mut parts: Vec<String>, max_terms: usize) -> String {
330 let k = max_terms.max(2);
331 while parts.len() > k {
332 parts = parts
333 .chunks(k)
334 .map(|c| {
335 if c.len() == 1 {
336 c[0].clone()
337 } else {
338 format!("SELECT * FROM ({})", c.join(" UNION ALL "))
339 }
340 })
341 .collect();
342 }
343 parts.join(" UNION ALL ")
344}
345
346pub fn col(row: &[SqlValue], i: usize) -> Result<&SqlValue> {
348 row.get(i)
349 .ok_or_else(|| Error::corrupted(format!("missing column {i} in result row")))
350}
351
352pub fn expect_len(response: &Response, n: usize) -> Result<()> {
354 if response.len() < n {
355 return Err(Error::backend(format!(
356 "backend returned {} result sets, expected {n}",
357 response.len()
358 )));
359 }
360 Ok(())
361}
362
363#[cfg(test)]
364mod tests {
365 use super::*;
366
367 #[test]
368 fn quoting() {
369 assert_eq!(sql_str("it's"), "'it''s'");
370 assert_eq!(sql_str("a\0b"), "CAST(X'610062' AS TEXT)");
371 assert_eq!(sql_f64(1.0), "1.0");
372 assert_eq!(sql_f64(1e300), "1e300");
373 }
374
375 #[cfg(feature = "serde")]
377 #[test]
378 fn sql_values_survive_arbitrary_precision_json() {
379 let plain: Vec<SqlValue> = serde_json::from_str(r#"[null, 7, 1.5, "x"]"#).unwrap();
380 assert_eq!(
381 plain,
382 [
383 SqlValue::Null,
384 SqlValue::Integer(7),
385 SqlValue::Real(1.5),
386 SqlValue::Text("x".into())
387 ]
388 );
389 let token = r#"[{"$serde_json::private::Number": "1.5"}, {"$serde_json::private::Number": "9007199254740993"}]"#;
391 let v: Vec<SqlValue> = serde_json::from_str(token).unwrap();
392 assert_eq!(
393 v,
394 [
395 SqlValue::Real(1.5),
396 SqlValue::Integer(9_007_199_254_740_993)
397 ]
398 );
399 }
400}