use systemprompt_extension::{Extension, LoaderError, Migration};
use systemprompt_identifiers::{ExtensionId, ToDbValue};
use systemprompt_traits::BoxedSource;
use tracing::info;
use super::exec::{TrackingWrite, check_cross_extension_alters, execute_statements_transactional};
use super::triggers::{self, Target};
use super::{MigrationService, RECORD_MIGRATION_SQL, budget};
use crate::services::SqlExecutor;
impl MigrationService<'_> {
pub(super) async fn execute_migration(
&self,
extension: &dyn Extension,
migration: &Migration,
) -> Result<(), LoaderError> {
let ext_id = &ExtensionId::new(extension.metadata().id);
check_cross_extension_alters(extension, migration)?;
info!(
extension = %ext_id,
version = migration.version,
name = %migration.name,
no_transaction = migration.no_transaction,
"Running migration"
);
let id = format!("{}_{:03}", ext_id, migration.version);
let checksum = migration.checksum();
let record_params: [&dyn ToDbValue; 5] =
[&id, &ext_id, &migration.version, &migration.name, &checksum];
if migration.no_transaction {
self.run_without_transaction(ext_id, migration, &record_params)
.await?;
} else {
let statements = SqlExecutor::parse_sql_statements(migration.sql).map_err(|e| {
LoaderError::MigrationStepFailed {
extension: ext_id.clone(),
context: format!(
"Failed to parse migration {} ({})",
migration.version, migration.name
),
source: Box::new(e),
}
})?;
execute_statements_transactional(
self.db,
&statements,
ext_id,
migration,
Some(TrackingWrite {
sql: RECORD_MIGRATION_SQL,
params: &record_params,
}),
)
.await?;
}
Ok(())
}
async fn run_without_transaction(
&self,
ext_id: &ExtensionId,
migration: &Migration,
record_params: &[&dyn ToDbValue],
) -> Result<(), LoaderError> {
self.apply_timeouts(ext_id, migration).await?;
let failed = |context: String, source: BoxedSource| LoaderError::MigrationStepFailed {
extension: ext_id.clone(),
context,
source,
};
let suspended = triggers::suspend(&mut Target::Pool(self.db), migration)
.await
.map_err(|e| failed("Failed to suspend row triggers".to_owned(), Box::new(e)))?;
if !suspended.is_empty() {
info!(
extension = %ext_id,
version = migration.version,
name = %migration.name,
triggers = %suspended.describe(),
"Row triggers suspended for migration",
);
}
let outcome = SqlExecutor::execute_statements_parsed(self.db, migration.sql)
.await
.map_err(|e| {
failed(
format!(
"Failed to execute migration {} ({})",
migration.version, migration.name
),
Box::new(e),
)
});
let restored = suspended.restore(&mut Target::Pool(self.db)).await;
self.clear_timeouts(ext_id).await?;
outcome?;
restored.map_err(|e| failed("Failed to restore row triggers".to_owned(), Box::new(e)))?;
self.db
.execute(&RECORD_MIGRATION_SQL, record_params)
.await
.map_err(|e| LoaderError::MigrationStepFailed {
extension: ext_id.clone(),
context: "Failed to record migration".to_owned(),
source: Box::new(e),
})?;
Ok(())
}
async fn apply_timeouts(
&self,
ext_id: &ExtensionId,
migration: &Migration,
) -> Result<(), LoaderError> {
let timeout = budget::statement_timeout(migration);
self.set_timeouts(ext_id, &budget::timeout_statements(timeout, false))
.await
}
async fn clear_timeouts(&self, ext_id: &ExtensionId) -> Result<(), LoaderError> {
self.set_timeouts(
ext_id,
&[
"SET statement_timeout = DEFAULT".to_owned(),
"SET lock_timeout = DEFAULT".to_owned(),
],
)
.await
}
async fn set_timeouts(
&self,
ext_id: &ExtensionId,
statements: &[String],
) -> Result<(), LoaderError> {
for statement in statements {
self.db
.execute(&statement.as_str(), &[])
.await
.map_err(|e| LoaderError::MigrationStepFailed {
extension: ext_id.clone(),
context: format!("Failed to run `{statement}`"),
source: Box::new(e),
})?;
}
Ok(())
}
}