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 name: String,
233}
234
235impl Default for Capabilities {
236 fn default() -> Self {
237 Self::native()
238 }
239}
240
241impl Capabilities {
242 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 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
269pub 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
289pub 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
296pub fn sql_opt_str(s: Option<&str>) -> String {
298 s.map_or_else(|| "NULL".into(), sql_str)
299}
300
301pub 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
319pub 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
338pub 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
344pub 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 #[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 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}