use std::marker::PhantomData;
use turso_orm_driver::ConnectionTrait;
use turso_sql::{Build, Expr, OnConflict, Returning, Statement};
use crate::entity::{
ActiveModelTrait, ColumnTrait, EntityTrait, FromQueryResult, IdenStatic, Iterable,
PrimaryKeyTrait,
};
use crate::types::TryFromU64;
use crate::{DbErr, Result};
#[derive(Clone, Debug)]
pub struct InsertResult<E: EntityTrait> {
pub last_insert_id: <E::PrimaryKey as PrimaryKeyTrait>::ValueType,
}
#[derive(Clone, Debug)]
pub struct Insert<A: ActiveModelTrait> {
model: A,
on_conflict: Option<OnConflict>,
}
impl<A: ActiveModelTrait> Insert<A> {
pub(crate) fn one(model: A) -> Self {
Self {
model,
on_conflict: None,
}
}
#[must_use]
pub fn on_conflict(mut self, on_conflict: OnConflict) -> Self {
self.on_conflict = Some(on_conflict);
self
}
fn statement(&self, returning: bool) -> turso_sql::Insert {
let mut columns = Vec::new();
let mut values = Vec::new();
for c in <<A::Entity as EntityTrait>::Column as Iterable>::iter() {
if let crate::entity::ActiveValue::Set(v) = self.model.get(c) {
columns.push(c.as_str());
values.push(Expr::val(v));
}
}
let mut insert =
turso_sql::Query::insert().into_table(<A::Entity as EntityTrait>::TABLE_NAME);
if columns.is_empty() {
insert = insert.default_values();
} else {
insert = insert.columns(columns).values(values);
}
if let Some(oc) = &self.on_conflict {
insert = insert.on_conflict(oc.clone());
}
if returning {
insert = insert.returning(Returning::All);
}
insert
}
pub fn build(&self) -> Statement {
self.statement(false).to_statement()
}
pub async fn exec<C: ConnectionTrait>(self, db: &C) -> Result<InsertResult<A::Entity>> {
let result = db.execute(self.build()).await?;
if result.rows_affected == 0 {
return Err(DbErr::RecordNotInserted);
}
let last_insert_id = if let Some(values) = self.model.get_primary_key_value()
&& !<A::Entity as EntityTrait>::PrimaryKey::auto_increment()
{
decode_pk::<A::Entity>(&values)?
} else {
let id = u64::try_from(result.last_insert_id).unwrap_or_default();
<<A::Entity as EntityTrait>::PrimaryKey as PrimaryKeyTrait>::ValueType::try_from_u64(
id,
)?
};
Ok(InsertResult { last_insert_id })
}
pub async fn exec_with_returning<C: ConnectionTrait>(
self,
db: &C,
) -> Result<<A::Entity as EntityTrait>::Model> {
let row = db
.query_one(self.statement(true).to_statement())
.await?
.ok_or(DbErr::RecordNotInserted)?;
<A::Entity as EntityTrait>::Model::from_query_result(&row, "")
}
}
fn decode_pk<E: EntityTrait>(
values: &[turso_sql::Value],
) -> Result<<E::PrimaryKey as PrimaryKeyTrait>::ValueType> {
match values.first() {
Some(turso_sql::Value::Integer(n)) if values.len() == 1 => {
let id = u64::try_from(*n).map_err(|_| DbErr::Type("negative primary key".into()))?;
<E::PrimaryKey as PrimaryKeyTrait>::ValueType::try_from_u64(id)
}
_ => Err(DbErr::Type(
"non-integer primary keys cannot be returned by exec(); use exec_with_returning()"
.into(),
)),
}
}
#[derive(Clone, Debug)]
pub struct InsertMany<A: ActiveModelTrait> {
models: Vec<A>,
on_conflict: Option<OnConflict>,
_a: PhantomData<A>,
}
impl<A: ActiveModelTrait> InsertMany<A> {
pub(crate) fn many(models: impl IntoIterator<Item = A>) -> Self {
Self {
models: models.into_iter().collect(),
on_conflict: None,
_a: PhantomData,
}
}
#[must_use]
pub fn on_conflict(mut self, on_conflict: OnConflict) -> Self {
self.on_conflict = Some(on_conflict);
self
}
pub fn is_empty(&self) -> bool {
self.models.is_empty()
}
fn statement(&self, returning: bool) -> Option<turso_sql::Insert> {
if self.models.is_empty() {
return None;
}
let columns: Vec<<A::Entity as EntityTrait>::Column> =
<<A::Entity as EntityTrait>::Column as Iterable>::iter()
.filter(|c| self.models.iter().any(|m| m.get(*c).is_set()))
.collect();
let mut insert = turso_sql::Query::insert()
.into_table(<A::Entity as EntityTrait>::TABLE_NAME)
.columns(columns.iter().map(IdenStatic::as_str));
for m in &self.models {
let row: Vec<Expr> = columns
.iter()
.map(|c| match m.get(*c).into_value() {
Some(v) => Expr::val(v),
None => c
.def()
.default
.unwrap_or_else(|| Expr::val(turso_sql::Value::Null)),
})
.collect();
insert = insert.values(row);
}
if let Some(oc) = &self.on_conflict {
insert = insert.on_conflict(oc.clone());
}
if returning {
insert = insert.returning(Returning::All);
}
Some(insert)
}
pub async fn exec<C: ConnectionTrait>(self, db: &C) -> Result<u64> {
match self.statement(false) {
None => Ok(0),
Some(stmt) => Ok(db.execute(stmt.to_statement()).await?.rows_affected),
}
}
pub async fn exec_with_returning<C: ConnectionTrait>(
self,
db: &C,
) -> Result<Vec<<A::Entity as EntityTrait>::Model>> {
match self.statement(true) {
None => Ok(Vec::new()),
Some(stmt) => db
.query_all(stmt.to_statement())
.await?
.iter()
.map(|r| <A::Entity as EntityTrait>::Model::from_query_result(r, ""))
.collect(),
}
}
}