use sea_orm::entity::prelude::*;
use sea_orm::{Condition, QueryOrder, QuerySelect, SqlErr};
use time::OffsetDateTime;
use toolkit_db::secure::{AccessScope, DBRunner, SecureEntityExt};
use toolkit_db_macros::Scopable;
use uuid::Uuid;
use crate::domain::error::ChatEngineError;
use crate::infra::db::migrations::{UQ_VARIANT_INDEX, UQ_VARIANT_INDEX_ROOT};
pub const VARIANT_INDEX_MAX_RETRIES: u32 = 3;
#[derive(Clone, Debug, PartialEq, DeriveEntityModel, Scopable)]
#[sea_orm(table_name = "messages")]
#[secure(unrestricted)]
#[allow(clippy::struct_excessive_bools)]
pub struct Model {
#[sea_orm(primary_key, auto_increment = false)]
pub message_id: Uuid,
pub session_id: Uuid,
pub tenant_id: Option<String>,
pub user_id: Option<String>,
pub parent_message_id: Option<Uuid>,
pub role: MessageRole,
#[sea_orm(column_type = "JsonBinary", nullable)]
pub file_ids: Option<serde_json::Value>,
pub variant_index: i32,
pub is_active: bool,
pub is_complete: bool,
pub is_hidden_from_user: bool,
pub is_hidden_from_backend: bool,
#[sea_orm(column_type = "JsonBinary", nullable)]
pub metadata: Option<serde_json::Value>,
pub created_at: OffsetDateTime,
}
#[derive(Clone, Debug, PartialEq, Eq, EnumIter, DeriveActiveEnum)]
#[sea_orm(rs_type = "String", db_type = "String(StringLen::N(16))")]
pub enum MessageRole {
#[sea_orm(string_value = "user")]
User,
#[sea_orm(string_value = "assistant")]
Assistant,
#[sea_orm(string_value = "system")]
System,
}
#[derive(Copy, Clone, Debug, EnumIter, DeriveRelation)]
pub enum Relation {
#[sea_orm(
belongs_to = "super::session::Entity",
from = "Column::SessionId",
to = "super::session::Column::SessionId",
on_update = "NoAction",
on_delete = "Cascade"
)]
Session,
#[sea_orm(
belongs_to = "Entity",
from = "Column::ParentMessageId",
to = "Column::MessageId",
on_update = "NoAction",
on_delete = "Restrict"
)]
Parent,
#[sea_orm(has_many = "super::message_reaction::Entity")]
Reaction,
#[sea_orm(has_many = "super::message_part::Entity")]
Part,
}
impl Related<super::session::Entity> for Entity {
fn to() -> RelationDef {
Relation::Session.def()
}
}
impl Related<super::message_reaction::Entity> for Entity {
fn to() -> RelationDef {
Relation::Reaction.def()
}
}
impl Related<super::message_part::Entity> for Entity {
fn to() -> RelationDef {
Relation::Part.def()
}
}
impl ActiveModelBehavior for ActiveModel {}
pub async fn compute_next_variant_index<R>(
runner: &R,
session_id: Uuid,
parent: Option<Uuid>,
) -> Result<i32, ChatEngineError>
where
R: DBRunner,
{
let scope = AccessScope::allow_all();
let parent_filter = match parent {
Some(p) => Condition::all().add(Column::ParentMessageId.eq(p)),
None => Condition::all().add(Column::ParentMessageId.is_null()),
};
let row = Entity::find()
.order_by_desc(Column::VariantIndex)
.limit(1)
.secure()
.scope_with(&scope)
.filter(Condition::all().add(Column::SessionId.eq(session_id)))
.filter(parent_filter)
.one(runner)
.await?;
Ok(match row {
Some(row) => row.variant_index + 1,
None => 0,
})
}
pub fn is_variant_unique_violation(err: &DbErr) -> bool {
let Some(SqlErr::UniqueConstraintViolation(message)) = err.sql_err() else {
return false;
};
if message.contains(UQ_VARIANT_INDEX) || message.contains(UQ_VARIANT_INDEX_ROOT) {
return true;
}
message.contains("messages.session_id") && message.contains("messages.variant_index")
}
#[cfg(test)]
#[path = "message_tests.rs"]
mod message_tests;