use drizzle_core::{
param::{OwnedParam, Param},
prepared::{
OwnedPreparedStatement as CoreOwnedPreparedStatement,
PreparedStatement as CorePreparedStatement,
},
traits::ToSQL,
};
use drizzle_sqlite::values::{OwnedSQLiteValue, SQLiteValue};
use std::{borrow::Cow, marker::PhantomData};
use ::worker::D1Database;
use super::{bind_statement, sqlite_value_to_js};
use drizzle_core::error::DrizzleError;
fn borrowed_values_to_js<'a, I>(bound: I) -> Vec<wasm_bindgen::JsValue>
where
I: IntoIterator<Item = SQLiteValue<'a>>,
{
bound.into_iter().map(|v| sqlite_value_to_js(&v)).collect()
}
fn owned_values_to_js<I>(bound: I) -> Vec<wasm_bindgen::JsValue>
where
I: IntoIterator<Item = OwnedSQLiteValue>,
{
bound
.into_iter()
.map(|v| sqlite_value_to_js(&SQLiteValue::from(v)))
.collect()
}
#[derive(Debug, Clone)]
pub struct PreparedStatement<'a, Marker = (), DecodedRow = ()> {
pub(crate) inner: CorePreparedStatement<'a, SQLiteValue<'a>>,
pub(crate) marker: PhantomData<(Marker, DecodedRow)>,
}
#[derive(Debug, Clone)]
pub struct OwnedPreparedStatement<Marker = (), DecodedRow = ()> {
pub(crate) inner: CoreOwnedPreparedStatement<OwnedSQLiteValue>,
pub(crate) marker: PhantomData<(Marker, DecodedRow)>,
}
impl<Marker, DecodedRow> From<OwnedPreparedStatement<Marker, DecodedRow>>
for PreparedStatement<'_, Marker, DecodedRow>
{
fn from(value: OwnedPreparedStatement<Marker, DecodedRow>) -> Self {
let sqlitevalue = value.inner.params.iter().map(|v| {
Param::new(
v.placeholder,
v.value.clone().map(|v| Cow::Owned(SQLiteValue::from(v))),
)
});
let inner = CorePreparedStatement {
text_segments: value.inner.text_segments,
params: sqlitevalue.collect::<Box<[_]>>(),
sql: value.inner.sql,
};
PreparedStatement {
inner,
marker: PhantomData,
}
}
}
impl<'a, Marker, DecodedRow> From<PreparedStatement<'a, Marker, DecodedRow>>
for OwnedPreparedStatement<Marker, DecodedRow>
{
fn from(value: PreparedStatement<'a, Marker, DecodedRow>) -> Self {
value.into_owned()
}
}
impl<'a, Marker, DecodedRow> PreparedStatement<'a, Marker, DecodedRow> {
pub(crate) fn new(inner: CorePreparedStatement<'a, SQLiteValue<'a>>) -> Self {
Self {
inner,
marker: PhantomData,
}
}
pub fn into_owned(self) -> OwnedPreparedStatement<Marker, DecodedRow> {
let owned_params = self.inner.params.iter().map(|p| OwnedParam {
placeholder: p.placeholder,
value: p
.value
.clone()
.map(|v| OwnedSQLiteValue::from(v.into_owned())),
});
let inner = CoreOwnedPreparedStatement {
text_segments: self.inner.text_segments.clone(),
params: owned_params.collect::<Box<[_]>>(),
sql: self.inner.sql.clone(),
};
OwnedPreparedStatement {
inner,
marker: PhantomData,
}
}
pub async fn execute<const N: usize>(
&self,
conn: &D1Database,
params: [drizzle_core::param::ParamBind<'a, SQLiteValue<'a>>; N],
) -> drizzle_core::error::Result<u64> {
debug_assert_eq!(
N,
self.inner.external_param_count(),
"parameter count mismatch: expected {} params but got {}",
self.inner.external_param_count(),
N
);
let (sql_str, bound) = self.inner.bind(params)?;
let values = borrowed_values_to_js(bound);
let stmt = bind_statement(conn.prepare(sql_str), &values)?;
run_execute(stmt).await
}
pub async fn all<T, const N: usize>(
&self,
conn: &D1Database,
params: [drizzle_core::param::ParamBind<'a, SQLiteValue<'a>>; N],
) -> drizzle_core::error::Result<Vec<T>>
where
T: for<'de> serde::Deserialize<'de>,
{
debug_assert_eq!(
N,
self.inner.external_param_count(),
"parameter count mismatch: expected {} params but got {}",
self.inner.external_param_count(),
N
);
let (sql_str, bound) = self.inner.bind(params)?;
let values = borrowed_values_to_js(bound);
let stmt = bind_statement(conn.prepare(sql_str), &values)?;
run_all::<T>(stmt).await
}
pub async fn get<T, const N: usize>(
&self,
conn: &D1Database,
params: [drizzle_core::param::ParamBind<'a, SQLiteValue<'a>>; N],
) -> drizzle_core::error::Result<T>
where
T: for<'de> serde::Deserialize<'de>,
{
debug_assert_eq!(
N,
self.inner.external_param_count(),
"parameter count mismatch: expected {} params but got {}",
self.inner.external_param_count(),
N
);
let (sql_str, bound) = self.inner.bind(params)?;
let values = borrowed_values_to_js(bound);
let stmt = bind_statement(conn.prepare(sql_str), &values)?;
run_get::<T>(stmt).await
}
}
impl<Marker, DecodedRow> OwnedPreparedStatement<Marker, DecodedRow> {
pub async fn execute<'a, const N: usize>(
&self,
conn: &D1Database,
params: [drizzle_core::param::ParamBind<'a, SQLiteValue<'a>>; N],
) -> drizzle_core::error::Result<u64> {
debug_assert_eq!(
N,
self.inner.external_param_count(),
"parameter count mismatch: expected {} params but got {}",
self.inner.external_param_count(),
N
);
let (sql_str, bound) = self.inner.bind(params)?;
let values = owned_values_to_js(bound);
let stmt = bind_statement(conn.prepare(sql_str), &values)?;
run_execute(stmt).await
}
pub async fn all<'a, T, const N: usize>(
&self,
conn: &D1Database,
params: [drizzle_core::param::ParamBind<'a, SQLiteValue<'a>>; N],
) -> drizzle_core::error::Result<Vec<T>>
where
T: for<'de> serde::Deserialize<'de>,
{
debug_assert_eq!(
N,
self.inner.external_param_count(),
"parameter count mismatch: expected {} params but got {}",
self.inner.external_param_count(),
N
);
let (sql_str, bound) = self.inner.bind(params)?;
let values = owned_values_to_js(bound);
let stmt = bind_statement(conn.prepare(sql_str), &values)?;
run_all::<T>(stmt).await
}
pub async fn get<'a, T, const N: usize>(
&self,
conn: &D1Database,
params: [drizzle_core::param::ParamBind<'a, SQLiteValue<'a>>; N],
) -> drizzle_core::error::Result<T>
where
T: for<'de> serde::Deserialize<'de>,
{
debug_assert_eq!(
N,
self.inner.external_param_count(),
"parameter count mismatch: expected {} params but got {}",
self.inner.external_param_count(),
N
);
let (sql_str, bound) = self.inner.bind(params)?;
let values = owned_values_to_js(bound);
let stmt = bind_statement(conn.prepare(sql_str), &values)?;
run_get::<T>(stmt).await
}
}
async fn run_execute(stmt: ::worker::D1PreparedStatement) -> drizzle_core::error::Result<u64> {
let result = stmt
.run()
.await
.map_err(|e| DrizzleError::Other(e.to_string().into()))?;
if !result.success() {
return Err(DrizzleError::Other(
result
.error()
.unwrap_or_else(|| "D1 statement failed".into())
.into(),
));
}
Ok(result
.meta()
.map_err(|e| DrizzleError::Other(e.to_string().into()))?
.and_then(|m| m.changes)
.unwrap_or(0) as u64)
}
async fn run_all<T>(stmt: ::worker::D1PreparedStatement) -> drizzle_core::error::Result<Vec<T>>
where
T: for<'de> serde::Deserialize<'de>,
{
let result = stmt
.all()
.await
.map_err(|e| DrizzleError::Other(e.to_string().into()))?;
if !result.success() {
return Err(DrizzleError::Other(
result
.error()
.unwrap_or_else(|| "D1 query failed".into())
.into(),
));
}
result
.results::<T>()
.map_err(|e| DrizzleError::Other(e.to_string().into()))
}
async fn run_get<T>(stmt: ::worker::D1PreparedStatement) -> drizzle_core::error::Result<T>
where
T: for<'de> serde::Deserialize<'de>,
{
stmt.first::<T>(None)
.await
.map_err(|e| DrizzleError::Other(e.to_string().into()))?
.ok_or(DrizzleError::NotFound)
}
impl<'a, Marker, DecodedRow> std::fmt::Display for PreparedStatement<'a, Marker, DecodedRow> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.inner)
}
}
impl<Marker, DecodedRow> std::fmt::Display for OwnedPreparedStatement<Marker, DecodedRow> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.inner)
}
}
impl<'a, Marker, DecodedRow> ToSQL<'a, SQLiteValue<'a>>
for PreparedStatement<'a, Marker, DecodedRow>
{
fn to_sql(&self) -> drizzle_core::sql::SQL<'a, SQLiteValue<'a>> {
self.inner.to_sql()
}
}
impl<'a, Marker, DecodedRow> ToSQL<'a, OwnedSQLiteValue>
for OwnedPreparedStatement<Marker, DecodedRow>
{
fn to_sql(&self) -> drizzle_core::sql::SQL<'a, OwnedSQLiteValue> {
self.inner.to_sql()
}
}
impl<'a, Marker, DecodedRow> ToSQL<'a, SQLiteValue<'a>>
for OwnedPreparedStatement<Marker, DecodedRow>
{
fn to_sql(&self) -> drizzle_core::sql::SQL<'a, SQLiteValue<'a>> {
self.inner.to_sql().map_params(SQLiteValue::from)
}
}