use crate::error::{Error, Result};
use std::fmt::Write;
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
#[cfg_attr(feature = "serde", serde(untagged))]
pub enum SqlValue {
Null,
Integer(i64),
Real(f64),
Text(String),
}
#[cfg(feature = "serde")]
impl<'de> serde::Deserialize<'de> for SqlValue {
fn deserialize<D: serde::Deserializer<'de>>(d: D) -> std::result::Result<Self, D::Error> {
struct V;
impl<'de> serde::de::Visitor<'de> for V {
type Value = SqlValue;
fn expecting(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("null, a number or a string")
}
fn visit_unit<E>(self) -> std::result::Result<SqlValue, E> {
Ok(SqlValue::Null)
}
fn visit_none<E>(self) -> std::result::Result<SqlValue, E> {
Ok(SqlValue::Null)
}
fn visit_some<D: serde::Deserializer<'de>>(
self,
d: D,
) -> std::result::Result<SqlValue, D::Error> {
d.deserialize_any(self)
}
fn visit_bool<E>(self, v: bool) -> std::result::Result<SqlValue, E> {
Ok(SqlValue::Integer(i64::from(v)))
}
fn visit_i64<E>(self, v: i64) -> std::result::Result<SqlValue, E> {
Ok(SqlValue::Integer(v))
}
fn visit_u64<E: serde::de::Error>(self, v: u64) -> std::result::Result<SqlValue, E> {
i64::try_from(v)
.map(SqlValue::Integer)
.or(Ok(SqlValue::Real(v as f64)))
}
fn visit_f64<E>(self, v: f64) -> std::result::Result<SqlValue, E> {
Ok(SqlValue::Real(v))
}
fn visit_str<E>(self, v: &str) -> std::result::Result<SqlValue, E> {
Ok(SqlValue::Text(v.to_owned()))
}
fn visit_string<E>(self, v: String) -> std::result::Result<SqlValue, E> {
Ok(SqlValue::Text(v))
}
fn visit_map<A: serde::de::MapAccess<'de>>(
self,
mut map: A,
) -> std::result::Result<SqlValue, A::Error> {
let Some((_, n)) = map.next_entry::<String, String>()? else {
return Err(serde::de::Error::custom("empty map is not a SQL value"));
};
if let Ok(i) = n.parse::<i64>() {
return Ok(SqlValue::Integer(i));
}
n.parse::<f64>()
.map(SqlValue::Real)
.map_err(serde::de::Error::custom)
}
}
d.deserialize_any(V)
}
}
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 compound_recursive_cte: bool,
pub vectors: bool,
pub vector_index_methods: bool,
pub name: String,
pub versioning: crate::version::Versioning,
}
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,
compound_recursive_cte: true,
vectors: false,
vector_index_methods: false,
name: "sqlite".into(),
versioning: crate::version::Versioning::Off,
}
}
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,
compound_recursive_cte: false,
vectors: false,
vector_index_methods: false,
name: "d1".into(),
versioning: crate::version::Versioning::Off,
}
}
}
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");
}
#[cfg(feature = "serde")]
#[test]
fn sql_values_survive_arbitrary_precision_json() {
let plain: Vec<SqlValue> = serde_json::from_str(r#"[null, 7, 1.5, "x"]"#).unwrap();
assert_eq!(
plain,
[
SqlValue::Null,
SqlValue::Integer(7),
SqlValue::Real(1.5),
SqlValue::Text("x".into())
]
);
let token = r#"[{"$serde_json::private::Number": "1.5"}, {"$serde_json::private::Number": "9007199254740993"}]"#;
let v: Vec<SqlValue> = serde_json::from_str(token).unwrap();
assert_eq!(
v,
[
SqlValue::Real(1.5),
SqlValue::Integer(9_007_199_254_740_993)
]
);
}
}