1use crate::error::{Error, Result};
11use std::fmt::Write;
12
13#[derive(Debug, Clone, PartialEq)]
15#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
16#[cfg_attr(feature = "serde", serde(untagged))]
17pub enum SqlValue {
18 Null,
19 Integer(i64),
20 Real(f64),
21 Text(String),
22}
23
24impl SqlValue {
25 pub fn as_i64(&self) -> Option<i64> {
26 match self {
27 Self::Integer(v) => Some(*v),
28 Self::Text(v) => v.parse().ok(),
31 Self::Real(v) if v.fract() == 0.0 && v.abs() < 9.0e15 => Some(*v as i64),
32 _ => None,
33 }
34 }
35
36 pub fn as_f64(&self) -> Option<f64> {
37 match self {
38 Self::Integer(v) => Some(*v as f64),
39 Self::Real(v) => Some(*v),
40 Self::Text(v) => v.parse().ok(),
41 Self::Null => None,
42 }
43 }
44
45 pub fn as_str(&self) -> Option<&str> {
46 match self {
47 Self::Text(v) => Some(v),
48 _ => None,
49 }
50 }
51
52 pub fn into_string(self) -> Option<String> {
53 match self {
54 Self::Text(v) => Some(v),
55 Self::Integer(v) => Some(v.to_string()),
56 Self::Real(v) => Some(v.to_string()),
57 Self::Null => None,
58 }
59 }
60
61 pub fn is_null(&self) -> bool {
62 matches!(self, Self::Null)
63 }
64}
65
66#[derive(Debug, Clone, PartialEq)]
69#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
70pub struct Statement {
71 pub sql: String,
72 #[cfg_attr(
73 feature = "serde",
74 serde(default, skip_serializing_if = "Vec::is_empty")
75 )]
76 pub params: Vec<SqlValue>,
77}
78
79impl Statement {
80 pub fn new(sql: impl Into<String>) -> Self {
81 Self {
82 sql: sql.into(),
83 params: Vec::new(),
84 }
85 }
86}
87
88impl From<String> for Statement {
89 fn from(sql: String) -> Self {
90 Self::new(sql)
91 }
92}
93
94impl From<&str> for Statement {
95 fn from(sql: &str) -> Self {
96 Self::new(sql)
97 }
98}
99
100#[derive(Debug, Clone, Copy, PartialEq, Eq)]
102#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
103#[cfg_attr(feature = "serde", serde(rename_all = "lowercase"))]
104pub enum Mode {
105 Read,
107 Atomic,
109}
110
111#[derive(Debug, Clone, PartialEq)]
113#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
114pub struct Request {
115 pub statements: Vec<Statement>,
116 pub mode: Mode,
117}
118
119impl Request {
120 pub fn read(statements: Vec<Statement>) -> Self {
121 Self {
122 statements,
123 mode: Mode::Read,
124 }
125 }
126
127 pub fn atomic(statements: Vec<Statement>) -> Self {
128 Self {
129 statements,
130 mode: Mode::Atomic,
131 }
132 }
133}
134
135#[derive(Debug, Clone, Default, PartialEq)]
137#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
138pub struct ResultSet {
139 #[cfg_attr(feature = "serde", serde(default))]
140 pub rows: Vec<Vec<SqlValue>>,
141 #[cfg_attr(feature = "serde", serde(default))]
143 pub changes: u64,
144}
145
146pub type Response = Vec<ResultSet>;
148
149#[derive(Debug, Clone, PartialEq)]
151#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
152#[cfg_attr(feature = "serde", serde(default, rename_all = "camelCase"))]
153pub struct Capabilities {
154 pub max_sql_len: usize,
156 pub max_statements: usize,
158 pub udf: bool,
160 pub interactive_transactions: bool,
162 pub int64_as_text: bool,
164 pub max_compound_select: usize,
166 pub name: String,
168}
169
170impl Default for Capabilities {
171 fn default() -> Self {
172 Self::native()
173 }
174}
175
176impl Capabilities {
177 pub fn native() -> Self {
179 Self {
180 max_sql_len: 1_000_000,
181 max_statements: 10_000,
182 udf: false,
183 interactive_transactions: true,
184 int64_as_text: false,
185 max_compound_select: 500,
186 name: "sqlite".into(),
187 }
188 }
189
190 pub fn d1() -> Self {
192 Self {
193 max_sql_len: 90_000,
194 max_statements: 50,
195 udf: false,
196 interactive_transactions: false,
197 int64_as_text: true,
198 max_compound_select: 5,
199 name: "d1".into(),
200 }
201 }
202}
203
204pub fn quote_str(out: &mut String, s: &str) {
206 if s.contains('\0') {
207 out.push_str("CAST(X'");
208 for b in s.as_bytes() {
209 let _ = write!(out, "{b:02X}");
210 }
211 out.push_str("' AS TEXT)");
212 return;
213 }
214 out.push('\'');
215 for c in s.chars() {
216 if c == '\'' {
217 out.push('\'');
218 }
219 out.push(c);
220 }
221 out.push('\'');
222}
223
224pub fn sql_str(s: &str) -> String {
226 let mut out = String::with_capacity(s.len() + 2);
227 quote_str(&mut out, s);
228 out
229}
230
231pub fn sql_opt_str(s: Option<&str>) -> String {
233 s.map_or_else(|| "NULL".into(), sql_str)
234}
235
236pub fn sql_f64(v: f64) -> String {
238 if v.is_nan() {
239 "NULL".into()
240 } else if v == f64::INFINITY {
241 "9e999".into()
242 } else if v == f64::NEG_INFINITY {
243 "-9e999".into()
244 } else {
245 let s = format!("{v:?}");
246 if s.contains('.') || s.contains('e') || s.contains("inf") {
247 s
248 } else {
249 format!("{s}.0")
250 }
251 }
252}
253
254pub fn union_all(mut parts: Vec<String>, max_terms: usize) -> String {
257 let k = max_terms.max(2);
258 while parts.len() > k {
259 parts = parts
260 .chunks(k)
261 .map(|c| {
262 if c.len() == 1 {
263 c[0].clone()
264 } else {
265 format!("SELECT * FROM ({})", c.join(" UNION ALL "))
266 }
267 })
268 .collect();
269 }
270 parts.join(" UNION ALL ")
271}
272
273pub fn col(row: &[SqlValue], i: usize) -> Result<&SqlValue> {
275 row.get(i)
276 .ok_or_else(|| Error::corrupted(format!("missing column {i} in result row")))
277}
278
279pub fn expect_len(response: &Response, n: usize) -> Result<()> {
281 if response.len() < n {
282 return Err(Error::backend(format!(
283 "backend returned {} result sets, expected {n}",
284 response.len()
285 )));
286 }
287 Ok(())
288}
289
290#[cfg(test)]
291mod tests {
292 use super::*;
293
294 #[test]
295 fn quoting() {
296 assert_eq!(sql_str("it's"), "'it''s'");
297 assert_eq!(sql_str("a\0b"), "CAST(X'610062' AS TEXT)");
298 assert_eq!(sql_f64(1.0), "1.0");
299 assert_eq!(sql_f64(1e300), "1e300");
300 }
301}