use pgrx::datum::DatumWithOid;
use pgrx::pg_sys;
use pgrx::prelude::*;
use std::ffi::CStr;
use crate::TViewError;
use crate::ddl::drop_tview;
use crate::error::TViewResult;
static mut PREV_PROCESS_UTILITY_HOOK: pg_sys::ProcessUtility_hook_type = None;
static mut HOOK_IN_PROGRESS: bool = false;
static mut HOOK_GUARD_LEVEL: i32 = 0;
pub unsafe fn install_hook() {
unsafe {
PREV_PROCESS_UTILITY_HOOK = pg_sys::ProcessUtility_hook;
pg_sys::ProcessUtility_hook = Some(tview_process_utility_hook);
}
}
pub unsafe fn ensure_hook_installed() {
unsafe {
static mut HOOK_INSTALLED: bool = false;
if !HOOK_INSTALLED {
install_hook();
HOOK_INSTALLED = true;
}
}
}
#[pg_guard]
#[allow(clippy::too_many_arguments)] unsafe extern "C-unwind" fn tview_process_utility_hook(
pstmt: *mut pg_sys::PlannedStmt,
query_string: *const ::std::os::raw::c_char,
read_only_tree: bool,
context: pg_sys::ProcessUtilityContext::Type,
params: pg_sys::ParamListInfo,
query_env: *mut pg_sys::QueryEnvironment,
dest: *mut pg_sys::DestReceiver,
qc: *mut pg_sys::QueryCompletion,
) {
if !pstmt.is_null()
&& unsafe { !(*pstmt).utilityStmt.is_null() }
&& unsafe { (*(*pstmt).utilityStmt).type_ } == pg_sys::NodeTag::T_TruncateStmt
{
crate::delta::begin_truncate();
}
if unsafe { HOOK_IN_PROGRESS } {
unsafe {
call_prev_hook_or_standard(
pstmt,
query_string,
read_only_tree,
context,
params,
query_env,
dest,
qc,
);
};
return;
}
if !pstmt.is_null() && unsafe { !(*pstmt).utilityStmt.is_null() } {
let tag = unsafe { (*(*pstmt).utilityStmt).type_ };
if tag == pg_sys::NodeTag::T_DoStmt || tag == pg_sys::NodeTag::T_CallStmt {
unsafe {
call_prev_hook_or_standard(
pstmt,
query_string,
read_only_tree,
context,
params,
query_env,
dest,
qc,
);
};
return;
}
}
unsafe {
HOOK_IN_PROGRESS = true;
HOOK_GUARD_LEVEL = pg_sys::GetCurrentTransactionNestLevel();
};
if !pstmt.is_null()
&& unsafe { !(*pstmt).utilityStmt.is_null() }
&& context == pg_sys::ProcessUtilityContext::PROCESS_UTILITY_TOPLEVEL
{
let utility_stmt = unsafe { (*pstmt).utilityStmt };
if unsafe { (*utility_stmt).type_ } == pg_sys::NodeTag::T_TransactionStmt {
#[allow(clippy::cast_ptr_alignment)] let xact_stmt = utility_stmt.cast::<pg_sys::TransactionStmt>();
if !xact_stmt.is_null() {
let kind = unsafe { (*xact_stmt).kind };
let ending = if kind == pg_sys::TransactionStmtKind::TRANS_STMT_COMMIT {
Some("COMMIT")
} else if kind == pg_sys::TransactionStmtKind::TRANS_STMT_PREPARE {
Some("PREPARE TRANSACTION")
} else {
None
};
if let Some(stmt) = ending {
if crate::suspend::is_suspended() {
crate::suspend::force_resume();
if let Err(e) = crate::suspend::catch_up() {
unsafe { HOOK_IN_PROGRESS = false };
error!("TVIEW catch-up after suspension failed before {stmt}: {e:?}");
}
}
if let Err(e) = crate::queue::flush_refresh_queue() {
unsafe { HOOK_IN_PROGRESS = false };
error!("TVIEW refresh failed before {stmt}: {e:?}");
}
if let Err(e) = crate::audit::flush_audit_buffer() {
unsafe { HOOK_IN_PROGRESS = false };
error!("Audit flush failed before {stmt}: {e:?}");
}
}
}
}
}
let (column_rename, partition_ddl) = if extension_installed() {
unsafe { (column_rename_of(pstmt), partition_ddl_of(pstmt)) }
} else {
(None, None)
};
let result = std::panic::catch_unwind(|| -> Result<Intercept, TViewError> {
if pstmt.is_null() {
return Ok(Intercept::PassThrough);
}
let pstmt_ref = unsafe { &*pstmt };
if pstmt_ref.utilityStmt.is_null() {
return Ok(Intercept::PassThrough);
}
let utility_stmt = pstmt_ref.utilityStmt;
let node_tag = unsafe { (*utility_stmt).type_ };
if let Some(extensions) = unsafe { extension_statement_names(utility_stmt) } {
if extensions.iter().any(|e| e == "jsonb_delta") {
crate::lifecycle::invalidate_jsonb_delta_cache();
}
if extensions.iter().any(|e| e == "pg_tviews") {
crate::revision::reset();
}
return Ok(Intercept::PassThrough);
}
if !extension_installed() {
return Ok(Intercept::PassThrough);
}
if node_tag == pg_sys::NodeTag::T_CreateTableAsStmt {
#[allow(clippy::cast_ptr_alignment)]
let ctas = utility_stmt.cast::<pg_sys::CreateTableAsStmt>();
return unsafe { inspect_create_table_as(ctas, pstmt, query_string) };
}
if node_tag == pg_sys::NodeTag::T_ExplainStmt {
#[allow(clippy::cast_ptr_alignment)] let query = unsafe { (*utility_stmt.cast::<pg_sys::ExplainStmt>()).query };
if let Some(table) = unsafe { tview_ctas_target(utility_of(query)) } {
return Ok(Intercept::Refuse(
format!("EXPLAIN of CREATE TABLE {table} AS … cannot create a TVIEW"),
None,
));
}
}
if node_tag == pg_sys::NodeTag::T_DropStmt {
#[allow(clippy::cast_ptr_alignment)] return Ok(Intercept::DropTable(
utility_stmt.cast::<pg_sys::DropStmt>(),
));
}
if node_tag == pg_sys::NodeTag::T_AlterTableStmt {
#[allow(clippy::cast_ptr_alignment)] let alter_stmt = utility_stmt.cast::<pg_sys::AlterTableStmt>();
match unsafe { handle_alter_table(alter_stmt, query_string) } {
Ok(true) => return Ok(Intercept::Handled),
Ok(false) => {}
Err(e) => return Err(e),
}
}
Ok(Intercept::PassThrough)
});
let should_pass_through = match result {
Ok(Ok(Intercept::PassThrough)) => true,
Ok(Ok(Intercept::Handled)) => false,
Ok(Ok(
Intercept::Refuse(_, Some(target)) | Intercept::CreateTview(Ctas { target, .. }),
)) if skipped_or_raise(&target) => true,
Ok(Ok(Intercept::Refuse(reason, _))) => {
unsafe { HOOK_IN_PROGRESS = false };
pg_sys::panic::ErrorReport::new(
PgSqlErrorCode::ERRCODE_FEATURE_NOT_SUPPORTED,
format!("pg_tviews: {reason}"),
function_name!(),
)
.set_hint(CTAS_HINT)
.report(PgLogLevel::ERROR);
unreachable!("ERROR does not return")
}
Ok(Ok(Intercept::CreateTview(ctas))) => {
unsafe { create_tview_from_ctas(&ctas, qc) };
false
}
Ok(Ok(Intercept::DropTable(drop_stmt))) => {
match unsafe { handle_drop_table(drop_stmt, query_string) } {
Ok(handled) => !handled,
Err(e) => {
unsafe { HOOK_IN_PROGRESS = false };
error!("{e}");
}
}
}
Ok(Err(handler_err)) => {
unsafe { HOOK_IN_PROGRESS = false };
error!("{handler_err}");
#[allow(unreachable_code)] {
true
}
}
Err(panic_info) => {
unsafe { HOOK_IN_PROGRESS = false };
let panic_info = match panic_info.downcast::<pg_sys::panic::CaughtError>() {
Ok(caught) => caught.rethrow(),
Err(panic_info) => panic_info,
};
let panic_info = match panic_info.downcast::<pg_sys::panic::ErrorReportWithLevel>() {
Ok(report) => pg_sys::panic::CaughtError::ErrorReport(*report).rethrow(),
Err(panic_info) => panic_info,
};
let panic_msg = panic_info
.downcast_ref::<&str>()
.map(|s| (*s).to_string())
.or_else(|| panic_info.downcast_ref::<String>().cloned())
.unwrap_or_else(|| format!("{panic_info:?}"));
error!(
"PANIC in ProcessUtility hook: {panic_msg} - This is a bug in pg_tviews - please report it!"
);
#[allow(unreachable_code)]
{
true
}
}
};
if should_pass_through {
unsafe {
call_prev_hook_or_standard(
pstmt,
query_string,
read_only_tree,
context,
params,
query_env,
dest,
qc,
);
}
if let Some((relid, old_name, new_name)) = column_rename.and_then(ColumnRename::resolve)
&& let Err(e) = crate::ddl::rename::handle_column_rename(relid, &old_name, &new_name)
{
unsafe { HOOK_IN_PROGRESS = false };
error!("pg_tviews: could not follow the column rename: {e}");
}
if let Some(ddl) = partition_ddl
&& let Err(e) = unsafe { ddl.apply() }
{
unsafe { HOOK_IN_PROGRESS = false };
error!("pg_tviews: could not update the triggers of a partition: {e}");
}
}
unsafe { HOOK_IN_PROGRESS = false };
}
enum Intercept {
PassThrough,
Handled,
DropTable(*mut pg_sys::DropStmt),
CreateTview(Ctas),
Refuse(String, Option<CtasTarget>),
}
struct Ctas {
target: CtasTarget,
query: String,
logged: Option<bool>,
fillfactor: Option<i32>,
}
struct CtasTarget {
schema: Option<String>,
table: String,
if_not_exists: bool,
}
impl CtasTarget {
fn name(&self) -> String {
match &self.schema {
Some(schema) => format!("\"{}\".{}", schema.replace('"', "\"\""), self.table),
None => self.table.clone(),
}
}
fn skipped(&self) -> TViewResult<bool> {
if !self.if_not_exists {
return Ok(false);
}
let args = [
unsafe {
DatumWithOid::new(
self.schema.as_deref(),
PgOid::BuiltIn(PgBuiltInOids::TEXTOID).value(),
)
},
unsafe {
DatumWithOid::new(
self.table.as_str(),
PgOid::BuiltIn(PgBuiltInOids::TEXTOID).value(),
)
},
];
Spi::connect(|client| {
client
.select(
"SELECT EXISTS (SELECT 1 FROM pg_catalog.pg_class c \
JOIN pg_catalog.pg_namespace n ON n.oid = c.relnamespace \
WHERE c.relname = $2 AND n.nspname = COALESCE($1, current_schema()))",
None,
&args,
)?
.first()
.get_one::<bool>()
})
.map(|exists| exists == Some(true))
.map_err(|e| TViewError::CatalogError {
operation: format!("Look up {}", self.name()),
pg_error: e.to_string(),
})
}
}
const CTAS_HINT: &str = "Create or change the TVIEW with \
SELECT tviews.pg_tviews_create_or_replace('tv_<entity>', $$<query>$$, \
options => '{\"logged\": …, \"fillfactor\": …}').";
struct ColumnRename {
relation: *const pg_sys::RangeVar,
old_name: String,
new_name: String,
}
impl ColumnRename {
fn resolve(self) -> Option<(pg_sys::Oid, String, String)> {
let relid = unsafe {
pg_sys::RangeVarGetRelidExtended(
self.relation,
pg_sys::NoLock.cast_signed(),
pg_sys::RVROption::RVR_MISSING_OK,
None,
std::ptr::null_mut(),
)
};
(relid != pg_sys::InvalidOid).then_some((relid, self.old_name, self.new_name))
}
}
fn extension_installed() -> bool {
unsafe {
pg_sys::IsTransactionState()
&& pg_sys::get_extension_oid(c"pg_tviews".as_ptr(), true) != pg_sys::InvalidOid
}
}
unsafe fn column_rename_of(pstmt: *const pg_sys::PlannedStmt) -> Option<ColumnRename> {
unsafe {
if pstmt.is_null() || (*pstmt).utilityStmt.is_null() {
return None;
}
let node = (*pstmt).utilityStmt;
if (*node).type_ != pg_sys::NodeTag::T_RenameStmt {
return None;
}
#[allow(clippy::cast_ptr_alignment)] let stmt = &*node.cast::<pg_sys::RenameStmt>();
if stmt.renameType != pg_sys::ObjectType::OBJECT_COLUMN
|| stmt.relation.is_null()
|| stmt.subname.is_null()
|| stmt.newname.is_null()
{
return None;
}
Some(ColumnRename {
relation: stmt.relation,
old_name: CStr::from_ptr(stmt.subname).to_string_lossy().into_owned(),
new_name: CStr::from_ptr(stmt.newname).to_string_lossy().into_owned(),
})
}
}
struct PartitionDdl {
tables: Vec<*mut pg_sys::RangeVar>,
changed_rows_of: Option<*mut pg_sys::RangeVar>,
}
impl PartitionDdl {
unsafe fn apply(self) -> TViewResult<()> {
crate::delta::clear_caches();
for rv in self.tables {
let oid = unsafe { resolve_relation_oid(rv) };
if oid != pg_sys::InvalidOid {
crate::dependency::triggers::ensure_partition_triggers(oid)?;
}
}
let parent = self
.changed_rows_of
.map_or(pg_sys::InvalidOid, |rv| unsafe { resolve_relation_oid(rv) });
if parent != pg_sys::InvalidOid {
crate::delta::refresh_tviews_over(parent)?;
}
Ok(())
}
}
unsafe fn partition_ddl_of(pstmt: *const pg_sys::PlannedStmt) -> Option<PartitionDdl> {
unsafe {
if pstmt.is_null() || (*pstmt).utilityStmt.is_null() {
return None;
}
let node = (*pstmt).utilityStmt;
let mut tables = Vec::new();
let mut changed_rows_of = None;
match (*node).type_ {
pg_sys::NodeTag::T_CreateStmt => {
#[allow(clippy::cast_ptr_alignment)] let stmt = &*node.cast::<pg_sys::CreateStmt>();
if !stmt.partbound.is_null() && !stmt.relation.is_null() {
tables.push(stmt.relation);
}
}
pg_sys::NodeTag::T_AlterTableStmt => {
#[allow(clippy::cast_ptr_alignment)]
let stmt = &*node.cast::<pg_sys::AlterTableStmt>();
for i in 0..pg_sys::list_length(stmt.cmds) {
let cmd = pg_sys::list_nth(stmt.cmds, i).cast::<pg_sys::AlterTableCmd>();
if cmd.is_null()
|| !matches!(
(*cmd).subtype,
pg_sys::AlterTableType::AT_AttachPartition
| pg_sys::AlterTableType::AT_DetachPartition
| pg_sys::AlterTableType::AT_DetachPartitionFinalize
)
|| (*cmd).def.is_null()
{
continue;
}
#[allow(clippy::cast_ptr_alignment)]
let partition = &*(*cmd).def.cast::<pg_sys::PartitionCmd>();
if !partition.name.is_null() {
tables.push(partition.name);
changed_rows_of = Some(stmt.relation);
}
}
}
_ => {}
}
(!tables.is_empty()).then_some(PartitionDdl {
tables,
changed_rows_of,
})
}
}
unsafe fn inspect_create_table_as(
ctas: *mut pg_sys::CreateTableAsStmt,
pstmt: *const pg_sys::PlannedStmt,
query_string: *const ::std::os::raw::c_char,
) -> Result<Intercept, TViewError> {
unsafe {
let Some(table_name) = tview_ctas_target(ctas.cast()) else {
return Ok(Intercept::PassThrough);
};
if crate::config::test_skip_ctas_intercept() {
return Ok(Intercept::PassThrough);
}
let ctas_ref = &*ctas;
let into = &*ctas_ref.into;
let rel = &*into.rel;
let schema = (!rel.schemaname.is_null()).then(|| {
CStr::from_ptr(rel.schemaname)
.to_string_lossy()
.into_owned()
});
let target = || CtasTarget {
schema: schema.clone(),
table: table_name.clone(),
if_not_exists: ctas_ref.if_not_exists,
};
let refuse = |reason: &str| Ok(Intercept::Refuse(reason.to_string(), Some(target())));
if ctas_ref.is_select_into {
return Ok(Intercept::Refuse(
format!("SELECT … INTO {table_name} cannot create a TVIEW"),
None,
));
}
let persistence = rel.relpersistence.cast_unsigned();
if persistence == pg_sys::RELPERSISTENCE_TEMP
|| schema
.as_deref()
.is_some_and(|schema| schema == "pg_temp" || schema.starts_with("pg_temp_"))
{
return refuse(&format!("{table_name} cannot be a temporary TVIEW"));
}
if !into.colNames.is_null() {
return refuse(&format!(
"{table_name} takes its column names from its query, not from a column list"
));
}
if !into.tableSpaceName.is_null() {
return refuse(&format!(
"TABLESPACE is not supported for TVIEW {table_name}"
));
}
if !into.accessMethod.is_null() {
return refuse(&format!("USING is not supported for TVIEW {table_name}"));
}
if into.skipData {
return refuse(&format!(
"WITH NO DATA is not supported: TVIEW {table_name} is always populated"
));
}
if is_execute(ctas_ref.query) {
return refuse(&format!(
"CREATE TABLE {table_name} AS EXECUTE cannot create a TVIEW"
));
}
if !ctas_ref.query.is_null() && contains_param(ctas_ref.query, std::ptr::null_mut()) {
return refuse(&format!(
"a query with parameters (such as PL/pgSQL variables) cannot define TVIEW \
{table_name}"
));
}
let mut fillfactor = None;
for i in 0..pg_sys::list_length(into.options) {
let option = pg_sys::list_nth(into.options, i).cast::<pg_sys::DefElem>();
if option.is_null() || (*option).defname.is_null() {
continue;
}
let name = CStr::from_ptr((*option).defname).to_string_lossy();
if name != "fillfactor" || !(*option).defnamespace.is_null() {
return refuse(&format!(
"storage parameter {name} is not supported for TVIEW {table_name}; only \
fillfactor is"
));
}
match option_integer(option) {
Some(value) if (10..=100).contains(&value) => fillfactor = Some(value),
_ => {
return refuse(&format!(
"fillfactor for TVIEW {table_name} must be an integer from 10 to 100"
));
}
}
}
let sql = if query_string.is_null() {
""
} else {
CStr::from_ptr(query_string).to_str().unwrap_or("")
};
let stmt_sql = statement_text(sql, pstmt);
let query = extract_ctas_select(stmt_sql, &table_name).ok_or_else(|| {
TViewError::InvalidSelectStatement {
sql: stmt_sql.to_string(),
reason: format!("Could not find 'CREATE TABLE {table_name} AS' in query"),
}
})?;
Ok(Intercept::CreateTview(Ctas {
target: target(),
query,
logged: (persistence == pg_sys::RELPERSISTENCE_UNLOGGED).then_some(false),
fillfactor,
}))
}
}
unsafe fn tview_ctas_target(node: *mut pg_sys::Node) -> Option<String> {
unsafe {
if node.is_null() || (*node).type_ != pg_sys::NodeTag::T_CreateTableAsStmt {
return None;
}
#[allow(clippy::cast_ptr_alignment)] let ctas = &*node.cast::<pg_sys::CreateTableAsStmt>();
if ctas.objtype != pg_sys::ObjectType::OBJECT_TABLE
|| ctas.into.is_null()
|| (*ctas.into).rel.is_null()
|| (*(*ctas.into).rel).relname.is_null()
{
return None;
}
let table = CStr::from_ptr((*(*ctas.into).rel).relname).to_str().ok()?;
(table.starts_with("tv_") && table.len() > 3).then(|| table.to_string())
}
}
unsafe fn utility_of(node: *mut pg_sys::Node) -> *mut pg_sys::Node {
unsafe {
if !node.is_null() && (*node).type_ == pg_sys::NodeTag::T_Query {
#[allow(clippy::cast_ptr_alignment)] let utility = (*node.cast::<pg_sys::Query>()).utilityStmt;
if !utility.is_null() {
return utility;
}
}
node
}
}
unsafe fn is_execute(query: *mut pg_sys::Node) -> bool {
unsafe {
let node = utility_of(query);
!node.is_null() && (*node).type_ == pg_sys::NodeTag::T_ExecuteStmt
}
}
unsafe extern "C-unwind" fn contains_param(
node: *mut pg_sys::Node,
context: *mut std::ffi::c_void,
) -> bool {
if node.is_null() {
return false;
}
unsafe {
match (*node).type_ {
pg_sys::NodeTag::T_Param => true,
#[allow(clippy::cast_ptr_alignment)] pg_sys::NodeTag::T_Query => pg_sys::query_tree_walker(
node.cast::<pg_sys::Query>(),
Some(contains_param),
context,
0,
),
_ => pg_sys::expression_tree_walker(node, Some(contains_param), context),
}
}
}
unsafe fn option_integer(option: *mut pg_sys::DefElem) -> Option<i32> {
unsafe {
if option.is_null() || (*option).arg.is_null() {
return None;
}
let arg = (*option).arg;
match (*arg).type_ {
#[allow(clippy::cast_ptr_alignment)] pg_sys::NodeTag::T_Integer => Some((*arg.cast::<pg_sys::Integer>()).ival),
#[allow(clippy::cast_ptr_alignment)] pg_sys::NodeTag::T_String => {
let value = (*arg.cast::<pg_sys::String>()).sval;
(!value.is_null())
.then(|| CStr::from_ptr(value).to_str().ok()?.parse().ok())
.flatten()
}
_ => None,
}
}
}
fn skipped_or_raise(target: &CtasTarget) -> bool {
target.skipped().unwrap_or_else(|e| {
unsafe { HOOK_IN_PROGRESS = false };
error!("{e}")
})
}
unsafe fn create_tview_from_ctas(ctas: &Ctas, qc: *mut pg_sys::QueryCompletion) {
if !crate::revision::is_current() {
unsafe { HOOK_IN_PROGRESS = false };
crate::revision::check();
}
let created = crate::ddl::replace::create_only(
&ctas.target.name(),
&ctas.query,
crate::ddl::replace::Options::storage(ctas.logged, ctas.fillfactor),
ctas.target.if_not_exists,
);
match created {
Ok(crate::ddl::replace::Created::Rows(rows)) => {
if !qc.is_null() {
unsafe {
(*qc).commandTag = pg_sys::CommandTag::CMDTAG_SELECT;
(*qc).nprocessed = rows;
}
}
}
Ok(crate::ddl::replace::Created::Skipped) => {}
Ok(crate::ddl::replace::Created::Exists(name)) => {
unsafe { HOOK_IN_PROGRESS = false };
pg_sys::panic::ErrorReport::new(
PgSqlErrorCode::ERRCODE_DUPLICATE_TABLE,
format!("TVIEW {name} already exists"),
function_name!(),
)
.set_hint(CTAS_HINT)
.report(PgLogLevel::ERROR);
}
Err(e) => {
unsafe { HOOK_IN_PROGRESS = false };
error!("{e}");
}
}
}
unsafe fn statement_text(query_string: &str, pstmt: *const pg_sys::PlannedStmt) -> &str {
if pstmt.is_null() {
return query_string;
}
let (location, len) = unsafe { ((*pstmt).stmt_location, (*pstmt).stmt_len) };
let Ok(start) = usize::try_from(location) else {
return query_string;
};
let end = match usize::try_from(len) {
Ok(len) if len > 0 => start.saturating_add(len),
_ => query_string.len(),
};
query_string
.get(start..end.min(query_string.len()))
.unwrap_or(query_string)
}
fn extract_ctas_select(stmt_sql: &str, table_name: &str) -> Option<String> {
let re = regex::Regex::new(&format!(
r#"(?is)^\s*create\s+(?:[a-z]+\s+){{0,2}}?table\s+(?:if\s+not\s+exists\s+)?(?:"?[^\s."]+"?\s*\.\s*)?"?{}"?\s+(?:with\s*\([^)]*\)\s*)?as\s+"#,
regex::escape(table_name)
))
.ok()?;
let m = re.find(stmt_sql)?;
let select = stmt_sql[m.end()..].trim().trim_end_matches(';').trim();
(!select.is_empty()).then(|| select.to_string())
}
unsafe fn extension_statement_names(node: *mut pg_sys::Node) -> Option<Vec<String>> {
unsafe {
let tag = (*node).type_;
if tag == pg_sys::NodeTag::T_CreateExtensionStmt {
#[allow(clippy::cast_ptr_alignment)]
let stmt = node.cast::<pg_sys::CreateExtensionStmt>();
let name = if (*stmt).extname.is_null() {
String::new()
} else {
CStr::from_ptr((*stmt).extname)
.to_string_lossy()
.into_owned()
};
return Some(vec![name]);
}
if tag == pg_sys::NodeTag::T_DropStmt {
#[allow(clippy::cast_ptr_alignment)] let stmt = node.cast::<pg_sys::DropStmt>();
if (*stmt).removeType != pg_sys::ObjectType::OBJECT_EXTENSION {
return None;
}
let mut names = Vec::new();
let objects = (*stmt).objects;
for i in 0..pg_sys::list_length(objects) {
let item = pg_sys::list_nth(objects, i).cast::<pg_sys::String>();
if !item.is_null() && !(*item).sval.is_null() {
names.push(CStr::from_ptr((*item).sval).to_string_lossy().into_owned());
}
}
return Some(names);
}
None
}
}
unsafe fn resolve_relation_oid(rv: *const pg_sys::RangeVar) -> pg_sys::Oid {
if rv.is_null() {
return pg_sys::InvalidOid;
}
unsafe {
pg_sys::RangeVarGetRelidExtended(
rv,
pg_sys::NoLock.cast_signed(),
pg_sys::RVROption::RVR_MISSING_OK,
None,
std::ptr::null_mut(),
)
}
}
pub fn release_hook_guard_on_abort(whole_xact: bool) {
unsafe {
if HOOK_IN_PROGRESS
&& (whole_xact || pg_sys::GetCurrentTransactionNestLevel() <= HOOK_GUARD_LEVEL)
{
HOOK_IN_PROGRESS = false;
}
}
}
unsafe fn handle_drop_table(
drop_stmt: *mut pg_sys::DropStmt,
_query_string: *const ::std::os::raw::c_char,
) -> Result<bool, TViewError> {
unsafe {
if drop_stmt.is_null() {
return Ok(false);
}
let drop_ref = &*drop_stmt;
if drop_ref.removeType != pg_sys::ObjectType::OBJECT_TABLE {
return Ok(false);
}
let objects = drop_ref.objects;
if objects.is_null() {
return Ok(false);
}
let if_exists = drop_ref.missing_ok;
let cascade = drop_ref.behavior == pg_sys::DropBehavior::DROP_CASCADE;
let num_tables = pg_sys::list_length(objects);
let mut tv_entries: Vec<(i32, String)> = Vec::new(); let mut has_non_tv = false;
for i in 0..num_tables {
let name_list = pg_sys::list_nth(objects, i).cast::<pg_sys::List>();
if name_list.is_null() || pg_sys::list_length(name_list) == 0 {
has_non_tv = true;
continue;
}
let rv = pg_sys::makeRangeVarFromNameList(name_list);
let relid = resolve_relation_oid(rv);
if relid == pg_sys::InvalidOid {
has_non_tv = true;
continue;
}
match crate::catalog::TviewMeta::load_for_tview(relid) {
Ok(Some(meta)) => tv_entries.push((i, format!("tv_{}", meta.entity_name))),
_ => has_non_tv = true,
}
}
if tv_entries.is_empty() {
return Ok(false);
}
for (_, name) in &tv_entries {
drop_tview(name, if_exists, cascade)?;
}
if has_non_tv {
for (idx, _) in tv_entries.iter().rev() {
pg_sys::list_delete_nth_cell(objects, *idx);
}
return Ok(false);
}
Ok(true)
} }
unsafe fn handle_alter_table(
alter_stmt: *mut pg_sys::AlterTableStmt,
_query_string: *const ::std::os::raw::c_char,
) -> Result<bool, TViewError> {
unsafe {
if alter_stmt.is_null() {
return Ok(false);
}
let alter_ref = &*alter_stmt;
let relation = alter_ref.relation;
if relation.is_null() {
return Ok(false);
}
let rel_ref = &*relation;
let table_name_cstr = rel_ref.relname;
if table_name_cstr.is_null() {
return Ok(false);
}
let table_name = CStr::from_ptr(table_name_cstr).to_str().unwrap_or("");
if !table_name.starts_with("tv_") {
return Ok(false);
}
let cmds = alter_ref.cmds;
if cmds.is_null() {
return Ok(false);
}
let num_cmds = pg_sys::list_length(cmds);
for i in 0..num_cmds {
let cmd_node = pg_sys::list_nth(cmds, i);
if cmd_node.is_null() {
continue;
}
let cmd = cmd_node.cast::<pg_sys::AlterTableCmd>();
if cmd.is_null() {
continue;
}
let cmd_ref = &*cmd;
if cmd_ref.subtype == pg_sys::AlterTableType::AT_SetUnLogged {
return Ok(false); } else if cmd_ref.subtype == pg_sys::AlterTableType::AT_SetLogged {
return Ok(false); }
}
Ok(false)
}
}
#[allow(clippy::too_many_arguments)] unsafe fn call_prev_hook_or_standard(
pstmt: *mut pg_sys::PlannedStmt,
query_string: *const ::std::os::raw::c_char,
read_only_tree: bool,
context: pg_sys::ProcessUtilityContext::Type,
params: pg_sys::ParamListInfo,
query_env: *mut pg_sys::QueryEnvironment,
dest: *mut pg_sys::DestReceiver,
qc: *mut pg_sys::QueryCompletion,
) {
unsafe {
match PREV_PROCESS_UTILITY_HOOK {
Some(prev_hook) => {
prev_hook(
pstmt,
query_string,
read_only_tree,
context,
params,
query_env,
dest,
qc,
);
}
None => {
pg_sys::standard_ProcessUtility(
pstmt,
query_string,
read_only_tree,
context,
params,
query_env,
dest,
qc,
);
}
}
}
}
#[cfg(test)]
mod ctas_extraction_tests {
use super::extract_ctas_select;
#[test]
fn extracts_select_and_strips_semicolon() {
let sql = "CREATE TABLE tv_post AS SELECT pk_post, id, data FROM tb_post;";
assert_eq!(
extract_ctas_select(sql, "tv_post").as_deref(),
Some("SELECT pk_post, id, data FROM tb_post")
);
}
#[test]
fn handles_schema_if_not_exists_and_case() {
let sql = "create table if not exists public.tv_post as\n select 1";
assert_eq!(
extract_ctas_select(sql, "tv_post").as_deref(),
Some("select 1")
);
}
#[test]
fn ignores_earlier_occurrences_of_the_name() {
let sql = "/* tv_post as */ CREATE TABLE tv_post AS SELECT 'tv_post as x' FROM t";
assert_eq!(extract_ctas_select(sql, "tv_post"), None);
let sql = "CREATE TABLE tv_post AS SELECT 'tv_post as x' FROM t";
assert_eq!(
extract_ctas_select(sql, "tv_post").as_deref(),
Some("SELECT 'tv_post as x' FROM t")
);
}
#[test]
fn does_not_match_a_different_table() {
let sql = "CREATE TABLE tv_post_extra AS SELECT 1";
assert_eq!(extract_ctas_select(sql, "tv_post"), None);
}
#[test]
fn handles_quoted_names() {
let sql = "CREATE TABLE \"s\".\"tv_post\" AS SELECT 1";
assert_eq!(
extract_ctas_select(sql, "tv_post").as_deref(),
Some("SELECT 1")
);
}
}