use crate::error::{Error, Result};
use std::fmt::Write;
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[cfg_attr(feature = "serde", serde(untagged))]
pub enum SqlValue {
Null,
Integer(i64),
Real(f64),
Text(String),
}
impl SqlValue {
pub fn as_i64(&self) -> Option<i64> {
match self {
Self::Integer(v) => Some(*v),
Self::Text(v) => v.parse().ok(),
Self::Real(v) if v.fract() == 0.0 && v.abs() < 9.0e15 => Some(*v as i64),
_ => None,
}
}
pub fn as_f64(&self) -> Option<f64> {
match self {
Self::Integer(v) => Some(*v as f64),
Self::Real(v) => Some(*v),
Self::Text(v) => v.parse().ok(),
Self::Null => None,
}
}
pub fn as_str(&self) -> Option<&str> {
match self {
Self::Text(v) => Some(v),
_ => None,
}
}
pub fn into_string(self) -> Option<String> {
match self {
Self::Text(v) => Some(v),
Self::Integer(v) => Some(v.to_string()),
Self::Real(v) => Some(v.to_string()),
Self::Null => None,
}
}
pub fn is_null(&self) -> bool {
matches!(self, Self::Null)
}
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct Statement {
pub sql: String,
#[cfg_attr(
feature = "serde",
serde(default, skip_serializing_if = "Vec::is_empty")
)]
pub params: Vec<SqlValue>,
}
impl Statement {
pub fn new(sql: impl Into<String>) -> Self {
Self {
sql: sql.into(),
params: Vec::new(),
}
}
}
impl From<String> for Statement {
fn from(sql: String) -> Self {
Self::new(sql)
}
}
impl From<&str> for Statement {
fn from(sql: &str) -> Self {
Self::new(sql)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[cfg_attr(feature = "serde", serde(rename_all = "lowercase"))]
pub enum Mode {
Read,
Atomic,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct Request {
pub statements: Vec<Statement>,
pub mode: Mode,
}
impl Request {
pub fn read(statements: Vec<Statement>) -> Self {
Self {
statements,
mode: Mode::Read,
}
}
pub fn atomic(statements: Vec<Statement>) -> Self {
Self {
statements,
mode: Mode::Atomic,
}
}
}
#[derive(Debug, Clone, Default, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct ResultSet {
#[cfg_attr(feature = "serde", serde(default))]
pub rows: Vec<Vec<SqlValue>>,
#[cfg_attr(feature = "serde", serde(default))]
pub changes: u64,
}
pub type Response = Vec<ResultSet>;
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[cfg_attr(feature = "serde", serde(default, rename_all = "camelCase"))]
pub struct Capabilities {
pub max_sql_len: usize,
pub max_statements: usize,
pub udf: bool,
pub interactive_transactions: bool,
pub int64_as_text: bool,
pub max_compound_select: usize,
pub name: String,
}
impl Default for Capabilities {
fn default() -> Self {
Self::native()
}
}
impl Capabilities {
pub fn native() -> Self {
Self {
max_sql_len: 1_000_000,
max_statements: 10_000,
udf: false,
interactive_transactions: true,
int64_as_text: false,
max_compound_select: 500,
name: "sqlite".into(),
}
}
pub fn d1() -> Self {
Self {
max_sql_len: 90_000,
max_statements: 50,
udf: false,
interactive_transactions: false,
int64_as_text: true,
max_compound_select: 5,
name: "d1".into(),
}
}
}
pub fn quote_str(out: &mut String, s: &str) {
if s.contains('\0') {
out.push_str("CAST(X'");
for b in s.as_bytes() {
let _ = write!(out, "{b:02X}");
}
out.push_str("' AS TEXT)");
return;
}
out.push('\'');
for c in s.chars() {
if c == '\'' {
out.push('\'');
}
out.push(c);
}
out.push('\'');
}
pub fn sql_str(s: &str) -> String {
let mut out = String::with_capacity(s.len() + 2);
quote_str(&mut out, s);
out
}
pub fn sql_opt_str(s: Option<&str>) -> String {
s.map_or_else(|| "NULL".into(), sql_str)
}
pub fn sql_f64(v: f64) -> String {
if v.is_nan() {
"NULL".into()
} else if v == f64::INFINITY {
"9e999".into()
} else if v == f64::NEG_INFINITY {
"-9e999".into()
} else {
let s = format!("{v:?}");
if s.contains('.') || s.contains('e') || s.contains("inf") {
s
} else {
format!("{s}.0")
}
}
}
pub fn union_all(mut parts: Vec<String>, max_terms: usize) -> String {
let k = max_terms.max(2);
while parts.len() > k {
parts = parts
.chunks(k)
.map(|c| {
if c.len() == 1 {
c[0].clone()
} else {
format!("SELECT * FROM ({})", c.join(" UNION ALL "))
}
})
.collect();
}
parts.join(" UNION ALL ")
}
pub fn col(row: &[SqlValue], i: usize) -> Result<&SqlValue> {
row.get(i)
.ok_or_else(|| Error::corrupted(format!("missing column {i} in result row")))
}
pub fn expect_len(response: &Response, n: usize) -> Result<()> {
if response.len() < n {
return Err(Error::backend(format!(
"backend returned {} result sets, expected {n}",
response.len()
)));
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn quoting() {
assert_eq!(sql_str("it's"), "'it''s'");
assert_eq!(sql_str("a\0b"), "CAST(X'610062' AS TEXT)");
assert_eq!(sql_f64(1.0), "1.0");
assert_eq!(sql_f64(1e300), "1e300");
}
}