use std::marker::PhantomData;
use turso_orm_driver::ConnectionTrait;
use turso_sql::{Build, Expr, IntoCondition, Returning, Statement, Value};
use crate::entity::{
ActiveModelTrait, ActiveValue, EntityTrait, FromQueryResult, IdenStatic, Iterable,
};
use crate::query::select::pk_condition;
use crate::{DbErr, Result};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct UpdateResult {
pub rows_affected: u64,
}
#[derive(Clone, Debug)]
pub struct UpdateOne<A: ActiveModelTrait> {
model: A,
}
impl<A: ActiveModelTrait> UpdateOne<A> {
pub(crate) fn new(model: A) -> Self {
Self { model }
}
fn statement(&self) -> Result<Option<turso_sql::Update>> {
let pk = self
.model
.get_primary_key_value()
.ok_or(DbErr::PrimaryKeyNotSet)?;
let mut update = turso_sql::Query::update().table(<A::Entity as EntityTrait>::TABLE_NAME);
let mut any = false;
for c in <<A::Entity as EntityTrait>::Column as Iterable>::iter() {
if let ActiveValue::Set(v) = self.model.get(c) {
update = update.value(c.as_str(), Expr::val(v));
any = true;
}
}
if !any {
return Ok(None);
}
update = update
.and_where(pk_condition::<A::Entity>(pk))
.returning(Returning::All);
Ok(Some(update))
}
pub async fn exec<C: ConnectionTrait>(
self,
db: &C,
) -> Result<<A::Entity as EntityTrait>::Model> {
if let Some(update) = self.statement()? {
let row = db
.query_one(update.to_statement())
.await?
.ok_or(DbErr::RecordNotUpdated)?;
<A::Entity as EntityTrait>::Model::from_query_result(&row, "")
} else {
let pk = self
.model
.get_primary_key_value()
.ok_or(DbErr::PrimaryKeyNotSet)?;
crate::query::Select::<A::Entity>::new()
.filter_by_pk(pk)
.one(db)
.await?
.ok_or(DbErr::RecordNotUpdated)
}
}
}
#[derive(Clone, Debug)]
pub struct UpdateMany<E: EntityTrait> {
query: turso_sql::Update,
_e: PhantomData<E>,
}
impl<E: EntityTrait> UpdateMany<E> {
pub(crate) fn new() -> Self {
Self {
query: turso_sql::Query::update().table(E::TABLE_NAME),
_e: PhantomData,
}
}
#[must_use]
pub fn col_expr(mut self, column: E::Column, value: impl Into<Expr>) -> Self {
self.query = self.query.value(column.as_str(), value);
self
}
#[must_use]
pub fn col(self, column: E::Column, value: impl Into<Value>) -> Self {
self.col_expr(column, Expr::val(value))
}
#[must_use]
pub fn set<A: ActiveModelTrait<Entity = E>>(mut self, model: &A) -> Self {
for c in E::Column::iter() {
if let ActiveValue::Set(v) = model.get(c) {
self.query = self.query.value(c.as_str(), Expr::val(v));
}
}
self
}
#[must_use]
pub fn filter(mut self, cond: impl IntoCondition) -> Self {
self.query = self.query.and_where(cond);
self
}
pub fn build(&self) -> Statement {
self.query.to_statement()
}
pub async fn exec<C: ConnectionTrait>(self, db: &C) -> Result<UpdateResult> {
if !self.query.has_sets() {
return Ok(UpdateResult { rows_affected: 0 });
}
let result = db.execute(self.build()).await?;
Ok(UpdateResult {
rows_affected: result.rows_affected,
})
}
pub async fn exec_with_returning<C: ConnectionTrait>(
mut self,
db: &C,
) -> Result<Vec<E::Model>> {
if !self.query.has_sets() {
return Ok(Vec::new());
}
self.query = self.query.returning(Returning::All);
db.query_all(self.build())
.await?
.iter()
.map(|r| E::Model::from_query_result(r, ""))
.collect()
}
}