use crate::error::{TViewError, TViewResult};
use crate::utils::quote_identifier;
use pgrx::datum::DatumWithOid;
use pgrx::prelude::*;
const ROW_HANDLER: &str = "pg_tview_trigger_handler";
const FLUSH_HANDLER: &str = "pg_tview_flush_trigger";
const TRIGGERS: [(&str, &str, &str); 2] = [
(ROW_HANDLER, "ROW", "row"),
(FLUSH_HANDLER, "STATEMENT", "flush"),
];
fn trigger_name(tag: &str, entity: &str, schema: &str, relname: &str) -> String {
crate::utils::fit_identifier(format!("trg_tview_{tag}_{entity}_on_{schema}_{relname}"))
}
struct InstalledTrigger {
table_oid: pg_sys::Oid,
table: String,
trigger: String,
function: String,
}
fn entity_triggers(
entity: &str,
table_oid: Option<pg_sys::Oid>,
) -> TViewResult<Vec<InstalledTrigger>> {
let query = format!(
"SELECT pg_catalog.quote_ident(n.nspname) || '.' || pg_catalog.quote_ident(c.relname), \
t.tgname::text, p.proname::text, t.tgrelid \
FROM pg_catalog.pg_trigger t \
JOIN pg_catalog.pg_proc p ON p.oid = t.tgfoid \
JOIN pg_catalog.pg_class c ON c.oid = t.tgrelid \
JOIN pg_catalog.pg_namespace n ON n.oid = c.relnamespace \
WHERE p.pronamespace = '{schema}'::pg_catalog.regnamespace \
AND p.proname IN ('{ROW_HANDLER}', '{FLUSH_HANDLER}') \
AND t.tgnargs = 1 \
AND t.tgargs = pg_catalog.convert_to($1, pg_catalog.getdatabaseencoding()) \
|| pg_catalog.decode('00', 'hex') \
AND ($2 IS NULL OR t.tgrelid = $2)",
schema = crate::utils::ext_schema(),
);
Spi::connect(|client| {
let args = [
unsafe { DatumWithOid::new(entity, PgOid::BuiltIn(PgBuiltInOids::TEXTOID).value()) },
unsafe { DatumWithOid::new(table_oid, PgOid::BuiltIn(PgBuiltInOids::OIDOID).value()) },
];
let mut found = Vec::new();
for row in client.select(&query, None, &args)? {
if let (Some(table), Some(trigger), Some(function), Some(table_oid)) = (
row.get::<String>(1)?,
row.get::<String>(2)?,
row.get::<String>(3)?,
row.get::<pg_sys::Oid>(4)?,
) {
found.push(InstalledTrigger {
table_oid,
table,
trigger,
function,
});
}
}
Ok::<_, spi::Error>(found)
})
.map_err(|e| TViewError::CatalogError {
operation: format!("Find triggers of TVIEW {entity}"),
pg_error: e.to_string(),
})
}
#[derive(Default)]
pub struct TriggerProblems {
pub orphaned: Vec<String>,
pub missing: Vec<String>,
pub untagged: Vec<String>,
}
pub fn trigger_problems() -> TViewResult<TriggerProblems> {
let query = format!(
"WITH ours AS ( \
SELECT t.tgname, t.tgrelid, p.proname, \
CASE WHEN t.tgnargs = 1 THEN pg_catalog.convert_from( \
pg_catalog.substring(t.tgargs, 1, pg_catalog.length(t.tgargs) - 1), \
pg_catalog.getdatabaseencoding()) END AS entity \
FROM pg_catalog.pg_trigger t \
JOIN pg_catalog.pg_proc p ON p.oid = t.tgfoid \
WHERE p.pronamespace = '{schema}'::pg_catalog.regnamespace \
AND p.proname IN ('{ROW_HANDLER}', '{FLUSH_HANDLER}') \
AND t.tgparentid = 0 \
), \
expected AS ( \
SELECT DISTINCT r.entity, r.relid \
FROM {schema}.pg_tview_reads r \
JOIN pg_catalog.pg_class c ON c.oid = r.relid AND c.relkind IN ('r', 'p') \
WHERE r.relid NOT IN (SELECT table_oid::oid FROM {meta}) \
) \
SELECT 'orphaned', pg_catalog.format('%I on %s', o.tgname, \
o.tgrelid::pg_catalog.regclass) \
FROM ours o \
WHERE o.entity IS NOT NULL \
AND NOT EXISTS (SELECT 1 FROM expected e \
WHERE e.entity = o.entity AND e.relid = o.tgrelid) \
UNION ALL \
SELECT 'missing', pg_catalog.format('%s (%s) on %s', e.entity, f.proname, \
e.relid::pg_catalog.regclass) \
FROM expected e \
CROSS JOIN (VALUES ('{ROW_HANDLER}'), ('{FLUSH_HANDLER}')) AS f(proname) \
WHERE NOT EXISTS (SELECT 1 FROM ours o \
WHERE o.entity = e.entity AND o.tgrelid = e.relid \
AND o.proname = f.proname) \
UNION ALL \
SELECT 'untagged', pg_catalog.format('%I on %s', o.tgname, \
o.tgrelid::pg_catalog.regclass) \
FROM ours o WHERE o.entity IS NULL \
ORDER BY 1, 2",
schema = crate::utils::ext_schema(),
meta = crate::utils::meta_table(),
);
Spi::connect(|client| {
let mut problems = TriggerProblems::default();
for row in client.select(&query, None, &[])? {
if let (Some(kind), Some(what)) = (row.get::<String>(1)?, row.get::<String>(2)?) {
match kind.as_str() {
"orphaned" => problems.orphaned.push(what),
"missing" => problems.missing.push(what),
_ => problems.untagged.push(what),
}
}
}
Ok::<_, spi::Error>(problems)
})
.map_err(|e| TViewError::CatalogError {
operation: "Check pg_tviews triggers".to_string(),
pg_error: e.to_string(),
})
}
pub fn install_triggers(table_oids: &[pg_sys::Oid], tview_entity: &str) -> TViewResult<()> {
let entity_arg = quote_identifier(tview_entity);
for &table_oid in table_oids {
let (schema, relname) = get_table_name(table_oid)?;
let qi_table = format!(
"{}.{}",
quote_identifier(&schema),
quote_identifier(&relname)
);
let installed = entity_triggers(tview_entity, Some(table_oid))?;
for (function, level, tag) in TRIGGERS {
if installed.iter().any(|t| t.function == function) {
continue;
}
let trigger_sql = format!(
"CREATE TRIGGER {}
AFTER INSERT OR UPDATE OR DELETE ON {qi_table}
FOR EACH {level}
EXECUTE FUNCTION {}.{function}({entity_arg})",
quote_identifier(&trigger_name(tag, tview_entity, &schema, &relname)),
crate::utils::ext_schema(),
);
crate::utils::spi_run_ddl(&trigger_sql).map_err(|e| TViewError::CatalogError {
operation: format!("Install {function} trigger on {qi_table}"),
pg_error: e,
})?;
}
}
Ok(())
}
pub fn sync_entity_triggers(table_oids: &[pg_sys::Oid], tview_entity: &str) -> TViewResult<()> {
for installed in entity_triggers(tview_entity, None)? {
if !table_oids.contains(&installed.table_oid) {
drop_trigger(installed.table_oid, &installed.table, &installed.trigger)?;
}
}
install_triggers(table_oids, tview_entity)
}
pub fn remove_entity_triggers(tview_entity: &str) -> TViewResult<()> {
for installed in entity_triggers(tview_entity, None)? {
drop_trigger(installed.table_oid, &installed.table, &installed.trigger)?;
}
Ok(())
}
fn drop_trigger(table_oid: pg_sys::Oid, table: &str, trigger: &str) -> TViewResult<()> {
let _owner = crate::owner::AsOwner::of_table(table_oid)?;
let drop_sql = format!(
"DROP TRIGGER IF EXISTS {} ON {table}",
quote_identifier(trigger)
);
crate::utils::spi_run_ddl(&drop_sql).map_err(|e| TViewError::CatalogError {
operation: format!("Drop trigger {trigger} from {table}"),
pg_error: e,
})
}
pub fn migrate_all_triggers_to_rust_handler() -> TViewResult<()> {
let pairs: Vec<(String, pg_sys::Oid)> = Spi::connect(|client| {
let rows = client.select(
&format!(
"SELECT m.entity, d.refobjid::oid AS table_oid \
FROM {} m \
JOIN pg_depend d ON d.objid = m.view_oid \
JOIN pg_class c ON c.oid = d.refobjid AND c.relkind = 'r' \
WHERE d.deptype = 'n'",
crate::utils::meta_table()
),
None,
&[],
)?;
let mut out = Vec::new();
for row in rows {
let entity: String = row["entity"].value()?.ok_or_else(|| {
spi::Error::from(TViewError::SpiError {
query: "migrate: SELECT entity".to_string(),
error: "entity column is NULL".to_string(),
})
})?;
let table_oid: pg_sys::Oid = row["table_oid"].value()?.ok_or_else(|| {
spi::Error::from(TViewError::SpiError {
query: "migrate: SELECT table_oid".to_string(),
error: "table_oid column is NULL".to_string(),
})
})?;
out.push((entity, table_oid));
}
Ok(out)
})
.map_err(|e: spi::Error| TViewError::CatalogError {
operation: "Migrate triggers: read pg_tview_meta".to_string(),
pg_error: format!("{e:?}"),
})?;
for (entity, table_oid) in pairs {
for (table, trigger) in legacy_triggers(table_oid)? {
drop_trigger(table_oid, &table, &trigger)?;
}
install_triggers(&[table_oid], &entity)?;
}
Ok(())
}
fn legacy_triggers(table_oid: pg_sys::Oid) -> TViewResult<Vec<(String, String)>> {
let query = format!(
"SELECT pg_catalog.quote_ident(n.nspname) || '.' || pg_catalog.quote_ident(c.relname), \
t.tgname::text \
FROM pg_catalog.pg_trigger t \
JOIN pg_catalog.pg_proc p ON p.oid = t.tgfoid \
JOIN pg_catalog.pg_class c ON c.oid = t.tgrelid \
JOIN pg_catalog.pg_namespace n ON n.oid = c.relnamespace \
WHERE t.tgrelid = $1 AND NOT t.tgisinternal \
AND (p.proname = 'tview_trigger_handler' \
OR (p.pronamespace = '{schema}'::pg_catalog.regnamespace \
AND p.proname IN ('{ROW_HANDLER}', '{FLUSH_HANDLER}') \
AND t.tgnargs = 0))",
schema = crate::utils::ext_schema(),
);
Spi::connect(|client| {
let args = [unsafe {
DatumWithOid::new(table_oid, PgOid::BuiltIn(PgBuiltInOids::OIDOID).value())
}];
let mut found = Vec::new();
for row in client.select(&query, None, &args)? {
if let (Some(table), Some(trigger)) = (row.get::<String>(1)?, row.get::<String>(2)?) {
found.push((table, trigger));
}
}
Ok::<_, spi::Error>(found)
})
.map_err(|e| TViewError::CatalogError {
operation: format!("Find legacy triggers on {table_oid:?}"),
pg_error: e.to_string(),
})
}
fn get_table_name(oid: pg_sys::Oid) -> TViewResult<(String, String)> {
let names = Spi::connect(|client| {
let args =
[unsafe { DatumWithOid::new(oid, PgOid::BuiltIn(PgBuiltInOids::OIDOID).value()) }];
client
.select(
"SELECT n.nspname::text, c.relname::text \
FROM pg_catalog.pg_class c \
JOIN pg_catalog.pg_namespace n ON n.oid = c.relnamespace \
WHERE c.oid = $1",
None,
&args,
)?
.first()
.get_two::<String, String>()
})
.map_err(|e| TViewError::CatalogError {
operation: format!("Get table name for OID {oid:?}"),
pg_error: e.to_string(),
})?;
match names {
(Some(schema), Some(relname)) => Ok((schema, relname)),
_ => Err(TViewError::DependencyResolutionFailed {
view_name: format!("OID {oid:?}"),
reason: "Table not found".to_string(),
}),
}
}
#[cfg(test)]
mod tests {
use super::trigger_name;
use crate::utils::MAX_IDENTIFIER_BYTES;
#[test]
fn test_trigger_name_short_is_verbatim() {
assert_eq!(
trigger_name("row", "post", "public", "tb_user"),
"trg_tview_row_post_on_public_tb_user"
);
assert_eq!(
trigger_name("flush", "post", "public", "tb_user"),
"trg_tview_flush_post_on_public_tb_user"
);
}
#[test]
fn test_trigger_name_row_and_flush_of_other_entities_differ() {
assert_ne!(
trigger_name("row", "flush_x", "public", "tb_t"),
trigger_name("flush", "x", "public", "tb_t")
);
}
#[test]
fn test_trigger_name_long_prefixes_stay_distinct() {
let common = "invoice_line_adjustment_with_a_deliberately_long_name_";
let a = trigger_name("row", &format!("{common}a"), "app", "tb_x");
let b = trigger_name("row", &format!("{common}b"), "app", "tb_x");
assert_eq!(a.len(), MAX_IDENTIFIER_BYTES);
assert_ne!(a, b);
}
#[test]
fn test_trigger_name_counts_bytes() {
let name = trigger_name("row", "note", &"é".repeat(40), "tb_note");
assert!(name.len() <= MAX_IDENTIFIER_BYTES);
}
}