use std::sync::Arc;
use pg_query::NodeEnum;
use pg_query::protobuf::ObjectType;
use systemprompt_extension::{Extension, LoaderError};
use tracing::info;
use super::phase::execute_phase;
use crate::services::{DatabaseProvider, SqlExecutor};
struct Retirement {
extension: String,
statements: Vec<String>,
}
fn refused(extension: &str, message: &str) -> LoaderError {
LoaderError::SchemaInstallationFailed {
extension: extension.to_owned(),
message: format!("retirement: {message}"),
}
}
fn check_statement(extension: &str, statement: &str) -> Result<(), LoaderError> {
let parsed = pg_query::parse(statement)
.map_err(|e| refused(extension, &format!("SQL parse failed: {e}")))?;
for raw in &parsed.protobuf.stmts {
let node = raw.stmt.as_ref().and_then(|s| s.node.as_ref());
let Some(NodeEnum::DropStmt(drop)) = node else {
return Err(refused(
extension,
&format!(
"only DROP … IF EXISTS is allowed, got `{}`",
statement.trim()
),
));
};
let allowed = matches!(
ObjectType::try_from(drop.remove_type),
Ok(ObjectType::ObjectTrigger
| ObjectType::ObjectFunction
| ObjectType::ObjectProcedure
| ObjectType::ObjectView)
);
if !allowed {
return Err(refused(
extension,
&format!(
"only triggers, functions, procedures and views can be retired, got `{}`",
statement.trim()
),
));
}
if !drop.missing_ok {
return Err(refused(
extension,
&format!("a retirement must say IF EXISTS: `{}`", statement.trim()),
));
}
}
Ok(())
}
fn prepare(extensions: &[Arc<dyn Extension>]) -> Result<Vec<Retirement>, LoaderError> {
let mut out = Vec::new();
for ext in extensions {
let extension = ext.id().to_owned();
let mut statements = Vec::new();
for retirement in ext.retirements() {
let parsed = SqlExecutor::parse_sql_statements(&retirement.sql)
.map_err(|e| refused(&extension, &format!("SQL split failed: {e}")))?;
for statement in parsed {
check_statement(&extension, &statement)?;
statements.push(statement);
}
}
if !statements.is_empty() {
out.push(Retirement {
extension,
statements,
});
}
}
Ok(out)
}
pub(super) fn check_retirements(extensions: &[Arc<dyn Extension>]) -> Result<(), LoaderError> {
prepare(extensions).map(|_| ())
}
pub(super) async fn apply_retirements(
db: &dyn DatabaseProvider,
extensions: &[Arc<dyn Extension>],
) -> Result<(), LoaderError> {
for retirement in prepare(extensions)? {
execute_phase(db, &retirement.statements, &[], &retirement.extension).await?;
info!(
extension = %retirement.extension,
statements = retirement.statements.len(),
"Retirements applied"
);
}
Ok(())
}