use pgrx::pg_sys;
use pgrx::prelude::*;
use std::cell::{Cell, RefCell};
use std::ffi::CStr;
static mut PREV_EXECUTOR_RUN: pg_sys::ExecutorRun_hook_type = None;
static mut PREV_EXECUTOR_FINISH: pg_sys::ExecutorFinish_hook_type = None;
thread_local! {
static FRAMES: RefCell<Vec<bool>> = const { RefCell::new(Vec::new()) };
static FLUSH_FUNCTION: Cell<pg_sys::Oid> = const { Cell::new(pg_sys::InvalidOid) };
}
const FLUSH_TRIGGER_PREFIX: &[u8] = b"trg_tview_flush_";
const FLUSH_FUNCTION_NAME: &CStr = c"pg_tview_flush_trigger";
pub unsafe fn install_hooks() {
unsafe {
PREV_EXECUTOR_RUN = pg_sys::ExecutorRun_hook;
pg_sys::ExecutorRun_hook = Some(executor_run);
PREV_EXECUTOR_FINISH = pg_sys::ExecutorFinish_hook;
pg_sys::ExecutorFinish_hook = Some(executor_finish);
}
}
pub fn flush_deferred() -> bool {
FRAMES.with(|f| {
let frames = f.borrow();
frames
.split_last()
.is_some_and(|(_, enclosing)| enclosing.iter().any(|&writing| writing))
})
}
pub fn reset() {
FRAMES.with(|f| f.borrow_mut().clear());
}
struct Frame;
impl Frame {
fn push(writing: bool) -> Self {
FRAMES.with(|f| f.borrow_mut().push(writing));
Self
}
}
impl Drop for Frame {
fn drop(&mut self) {
FRAMES.with(|f| {
f.borrow_mut().pop();
});
}
}
#[cfg(any(feature = "pg16", feature = "pg17"))]
#[pg_guard]
unsafe extern "C-unwind" fn executor_run(
query_desc: *mut pg_sys::QueryDesc,
direction: pg_sys::ScanDirection::Type,
count: u64,
execute_once: bool,
) {
let _frame = Frame::push(unsafe { writes_tracked_table(query_desc) });
unsafe {
match PREV_EXECUTOR_RUN {
Some(prev) => prev(query_desc, direction, count, execute_once),
None => pg_sys::standard_ExecutorRun(query_desc, direction, count, execute_once),
}
}
}
#[cfg(feature = "pg18")]
#[pg_guard]
unsafe extern "C-unwind" fn executor_run(
query_desc: *mut pg_sys::QueryDesc,
direction: pg_sys::ScanDirection::Type,
count: u64,
) {
let _frame = Frame::push(unsafe { writes_tracked_table(query_desc) });
unsafe {
match PREV_EXECUTOR_RUN {
Some(prev) => prev(query_desc, direction, count),
None => pg_sys::standard_ExecutorRun(query_desc, direction, count),
}
}
}
#[pg_guard]
unsafe extern "C-unwind" fn executor_finish(query_desc: *mut pg_sys::QueryDesc) {
let writing = unsafe { writes_tracked_table(query_desc) };
{
let _frame = Frame::push(writing);
unsafe {
match PREV_EXECUTOR_FINISH {
Some(prev) => prev(query_desc),
None => pg_sys::standard_ExecutorFinish(query_desc),
}
}
}
if writing && !FRAMES.with(|f| f.borrow().iter().any(|&w| w)) {
crate::trigger::flush_after_statement();
}
}
unsafe fn writes_tracked_table(query_desc: *mut pg_sys::QueryDesc) -> bool {
unsafe {
if query_desc.is_null() || (*query_desc).estate.is_null() {
return false;
}
let opened = (*(*query_desc).estate).es_opened_result_relations;
for i in 0..pg_sys::list_length(opened) {
let rel = pg_sys::list_nth(opened, i).cast::<pg_sys::ResultRelInfo>();
if !rel.is_null() && has_flush_trigger((*rel).ri_TrigDesc) {
return true;
}
}
false
}
}
unsafe fn has_flush_trigger(desc: *const pg_sys::TriggerDesc) -> bool {
unsafe {
if desc.is_null() {
return false;
}
let count = usize::try_from((*desc).numtriggers).unwrap_or(0);
(0..count).any(|i| {
let trigger = (*desc).triggers.add(i);
fires((*trigger).tgenabled.cast_unsigned())
&& !(*trigger).tgname.is_null()
&& CStr::from_ptr((*trigger).tgname)
.to_bytes()
.starts_with(FLUSH_TRIGGER_PREFIX)
&& is_flush_function((*trigger).tgfoid)
})
}
}
fn fires(enabled: u8) -> bool {
let replica = unsafe { pg_sys::SessionReplicationRole }
== pg_sys::SESSION_REPLICATION_ROLE_REPLICA.cast_signed();
match enabled {
pg_sys::TRIGGER_DISABLED => false,
pg_sys::TRIGGER_FIRES_ON_ORIGIN => !replica,
pg_sys::TRIGGER_FIRES_ON_REPLICA => replica,
_ => true,
}
}
fn is_flush_function(function: pg_sys::Oid) -> bool {
if FLUSH_FUNCTION.with(Cell::get) == function {
return true;
}
let name = unsafe { pg_sys::get_func_name(function) };
if name.is_null() {
return false;
}
let matches = unsafe {
let matches = CStr::from_ptr(name) == FLUSH_FUNCTION_NAME;
pg_sys::pfree(name.cast());
matches
};
if matches {
FLUSH_FUNCTION.with(|f| f.set(function));
}
matches
}