#![allow(dead_code)]
use std::sync::Arc;
use chat_engine::domain::ports::MessageRepo;
use chat_engine::domain::ports::PluginConfigRepo;
use chat_engine::domain::ports::SessionRepo;
use chat_engine::domain::ports::SessionTypeRepo;
use chat_engine::infra::db::Migrator;
use chat_engine::infra::db::entity::{file_citation, message, message_part, session};
use chat_engine::infra::db::repo::ChatEngineDb;
use chat_engine::infra::db::repo::message_repo::SeaMessageRepo;
use chat_engine::infra::db::repo::plugin_config_repo::SeaPluginConfigRepo;
use chat_engine::infra::db::repo::session_repo::SeaSessionRepo;
use chat_engine::infra::db::repo::session_type_repo::SeaSessionTypeRepo;
use sea_orm::{ActiveValue::Set, ColumnTrait, Condition, EntityTrait, QueryOrder};
use time::OffsetDateTime;
use toolkit_db::secure::{AccessScope, SecureEntityExt, SecureInsertExt};
use toolkit_db::{ConnectOpts, DBProvider, connect_db};
use uuid::Uuid;
pub struct DbHarness {
pub db: Arc<ChatEngineDb>,
pub sessions: Arc<dyn SessionRepo>,
pub session_types: Arc<dyn SessionTypeRepo>,
pub messages: Arc<dyn MessageRepo>,
pub plugin_configs: Arc<dyn PluginConfigRepo>,
}
pub async fn setup_sqlite() -> DbHarness {
use sea_orm_migration::MigratorTrait;
let opts = ConnectOpts {
max_conns: Some(1),
min_conns: Some(1),
..Default::default()
};
let db = connect_db("sqlite::memory:", opts)
.await
.expect("connect sqlite::memory:");
toolkit_db::migration_runner::run_migrations_for_testing(&db, Migrator::migrations())
.await
.expect("apply chat-engine migrations");
let provider: Arc<ChatEngineDb> = Arc::new(DBProvider::new(db));
let sessions: Arc<dyn SessionRepo> = Arc::new(SeaSessionRepo::new(Arc::clone(&provider)));
let session_types: Arc<dyn SessionTypeRepo> =
Arc::new(SeaSessionTypeRepo::new(Arc::clone(&provider)));
let messages: Arc<dyn MessageRepo> = Arc::new(SeaMessageRepo::new(Arc::clone(&provider)));
let plugin_configs: Arc<dyn PluginConfigRepo> =
Arc::new(SeaPluginConfigRepo::new(Arc::clone(&provider)));
DbHarness {
db: provider,
sessions,
session_types,
messages,
plugin_configs,
}
}
pub async fn seed_session_type(h: &DbHarness, plugin_instance_id: &str) -> Uuid {
let id = Uuid::new_v4();
let now = OffsetDateTime::now_utc();
h.session_types
.insert(chat_engine::domain::ports::NewSessionType {
session_type_id: id,
name: "integration-test".to_owned(),
plugin_instance_id: Some(plugin_instance_id.to_owned()),
created_at: now,
updated_at: now,
})
.await
.expect("insert session_type row");
id
}
pub async fn seed_active_session(
h: &DbHarness,
tenant_id: &str,
user_id: &str,
session_type_id: Uuid,
) -> Uuid {
let id = Uuid::new_v4();
let now = OffsetDateTime::now_utc();
h.sessions
.insert(chat_engine::domain::ports::NewSession {
session_id: id,
tenant_id: tenant_id.to_owned(),
user_id: user_id.to_owned(),
client_id: None,
session_type_id: Some(session_type_id),
metadata: None,
created_at: now,
updated_at: now,
})
.await
.expect("insert session row");
id
}
pub async fn seed_message(
h: &DbHarness,
session_id: Uuid,
parent_message_id: Option<Uuid>,
variant_index: i32,
) -> Uuid {
let conn = h.db.conn().expect("conn for seed_message");
let scope = AccessScope::allow_all();
let id = Uuid::new_v4();
let now = OffsetDateTime::now_utc();
let am = message::ActiveModel {
message_id: Set(id),
session_id: Set(session_id),
tenant_id: Set(None),
user_id: Set(None),
parent_message_id: Set(parent_message_id),
role: Set(message::MessageRole::User),
file_ids: Set(None),
variant_index: Set(variant_index),
is_active: Set(true),
is_complete: Set(true),
is_hidden_from_user: Set(false),
is_hidden_from_backend: Set(false),
metadata: Set(None),
created_at: Set(now),
};
message::Entity::insert(am)
.secure()
.scope_unchecked(&scope)
.expect("scope message insert")
.exec(&conn)
.await
.expect("insert message row");
id
}
pub async fn find_message(db: &Arc<ChatEngineDb>, message_id: Uuid) -> Option<message::Model> {
let conn = db.conn().expect("conn for find_message");
let scope = AccessScope::allow_all();
message::Entity::find()
.secure()
.scope_with(&scope)
.filter(Condition::all().add(message::Column::MessageId.eq(message_id)))
.one(&conn)
.await
.expect("read messages row")
}
pub async fn list_messages(db: &Arc<ChatEngineDb>, session_id: Uuid) -> Vec<message::Model> {
let conn = db.conn().expect("conn for list_messages");
let scope = AccessScope::allow_all();
message::Entity::find()
.order_by_asc(message::Column::CreatedAt)
.secure()
.scope_with(&scope)
.filter(Condition::all().add(message::Column::SessionId.eq(session_id)))
.all(&conn)
.await
.expect("list messages")
}
pub async fn find_assistant_message(
db: &Arc<ChatEngineDb>,
session_id: Uuid,
) -> Option<message::Model> {
let conn = db.conn().expect("conn for find_assistant_message");
let scope = AccessScope::allow_all();
message::Entity::find()
.secure()
.scope_with(&scope)
.filter(
Condition::all()
.add(message::Column::SessionId.eq(session_id))
.add(message::Column::Role.eq("assistant")),
)
.one(&conn)
.await
.expect("find assistant row")
}
pub async fn message_parts_ordered(
db: &Arc<ChatEngineDb>,
session_id: Uuid,
role: &str,
) -> Vec<(String, i32, serde_json::Value)> {
let conn = db.conn().expect("conn for message_parts_ordered");
let scope = AccessScope::allow_all();
let msgs = message::Entity::find()
.secure()
.scope_with(&scope)
.filter(
Condition::all()
.add(message::Column::SessionId.eq(session_id))
.add(message::Column::Role.eq(role)),
)
.all(&conn)
.await
.expect("load messages by role");
let ids: Vec<Uuid> = msgs.iter().map(|m| m.message_id).collect();
let rows = message_part::Entity::find()
.order_by_asc(message_part::Column::Number)
.secure()
.scope_with(&scope)
.filter(Condition::all().add(message_part::Column::MessageId.is_in(ids)))
.all(&conn)
.await
.expect("load message parts");
rows.into_iter()
.map(|p| (part_type_str(&p.r#type).to_owned(), p.number, p.content))
.collect()
}
pub async fn file_citations_for_message(
db: &Arc<ChatEngineDb>,
message_id: Uuid,
) -> Vec<serde_json::Value> {
let conn = db.conn().expect("conn for file_citations_for_message");
let scope = AccessScope::allow_all();
let parts = message_part::Entity::find()
.secure()
.scope_with(&scope)
.filter(Condition::all().add(message_part::Column::MessageId.eq(message_id)))
.all(&conn)
.await
.expect("load parts");
let pids: Vec<Uuid> = parts.iter().map(|p| p.id).collect();
let rows = file_citation::Entity::find()
.order_by_asc(file_citation::Column::Number)
.secure()
.scope_with(&scope)
.filter(Condition::all().add(file_citation::Column::MessagePartId.is_in(pids)))
.all(&conn)
.await
.expect("load file citations");
rows.into_iter().map(|r| r.content).collect()
}
fn part_type_str(t: &message_part::MessagePartType) -> &'static str {
match t {
message_part::MessagePartType::Text => "text",
message_part::MessagePartType::Code => "code",
message_part::MessagePartType::Images => "images",
message_part::MessagePartType::Videos => "videos",
message_part::MessagePartType::Links => "links",
message_part::MessagePartType::Statuses => "statuses",
}
}
pub async fn message_text(db: &Arc<ChatEngineDb>, message_id: Uuid) -> String {
let conn = db.conn().expect("conn for message_text");
let scope = AccessScope::allow_all();
let rows = message_part::Entity::find()
.order_by_asc(message_part::Column::Number)
.secure()
.scope_with(&scope)
.filter(
Condition::all()
.add(message_part::Column::MessageId.eq(message_id))
.add(message_part::Column::Type.eq("text")),
)
.all(&conn)
.await
.expect("load message parts");
rows.iter()
.filter_map(|p| p.content.get("text").and_then(|v| v.as_str()))
.collect::<Vec<_>>()
.join("\n")
}
pub async fn wait_for_finalize(
db: &Arc<ChatEngineDb>,
session_id: Uuid,
deadline: std::time::Duration,
) -> message::Model {
let started = std::time::Instant::now();
loop {
let row = find_assistant_message(db, session_id).await;
if let Some(m) = row {
if m.is_complete || m.metadata.is_some() {
return m;
}
assert!(
started.elapsed() < deadline,
"assistant row for session {session_id} not finalised within \
{deadline:?}; last row = {m:?}",
);
} else if started.elapsed() >= deadline {
panic!("no assistant row appeared for session {session_id} within {deadline:?}");
}
tokio::time::sleep(std::time::Duration::from_millis(5)).await;
}
}
pub fn db_provider(h: &DbHarness) -> &Arc<ChatEngineDb> {
&h.db
}
pub async fn force_lifecycle_state(db: &Arc<ChatEngineDb>, session_id: Uuid, new_state: &str) {
use sea_orm::sea_query::Expr;
use toolkit_db::secure::SecureUpdateExt;
let conn = db.conn().expect("conn for force_lifecycle_state");
let scope = AccessScope::allow_all();
session::Entity::update_many()
.secure()
.scope_with(&scope)
.filter(Condition::all().add(session::Column::SessionId.eq(session_id)))
.col_expr(
session::Column::LifecycleState,
Expr::value(new_state.to_owned()),
)
.col_expr(
session::Column::UpdatedAt,
Expr::value(OffsetDateTime::now_utc()),
)
.exec(&conn)
.await
.expect("flip lifecycle state");
}