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 vectors: bool,
238 pub vector_index_methods: bool,
241 pub name: String,
243 pub versioning: crate::version::Versioning,
246}
247
248impl Default for Capabilities {
249 fn default() -> Self {
250 Self::native()
251 }
252}
253
254impl Capabilities {
255 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 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 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 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
292pub 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
312pub 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
319pub fn sql_opt_str(s: Option<&str>) -> String {
321 s.map_or_else(|| "NULL".into(), sql_str)
322}
323
324pub 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
342pub 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
361pub 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
367pub 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 #[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 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}