use pgrx::AllocatedByPostgres;
use pgrx::datum::DatumWithOid;
use pgrx::heap_tuple::PgHeapTuple;
use pgrx::pg_sys;
use pgrx::prelude::*;
use std::collections::HashMap;
use std::sync::{LazyLock, Mutex};
macro_rules! log_debug {
($($arg:tt)+) => {
if $crate::config::log_level().eq_ignore_ascii_case("debug") {
::pgrx::notice!($($arg)+);
} else {
::pgrx::debug1!($($arg)+);
}
};
}
pub(crate) use log_debug;
pub fn spi_run_ddl(sql: &str) -> Result<(), String> {
use std::ffi::CString;
log_debug!(
"spi_run_ddl() called with SQL ({} chars): {}",
sql.len(),
&sql[..sql.len().min(200)]
);
let c_sql = CString::new(sql).map_err(|e| format!("DDL SQL contains null byte: {e}"))?;
unsafe {
#[allow(clippy::cast_possible_wrap)] let connect_result = pg_sys::SPI_connect_ext(pg_sys::SPI_OPT_NONATOMIC as i32);
#[allow(clippy::cast_possible_wrap)]
if connect_result != pg_sys::SPI_OK_CONNECT as i32 {
error!(
"spi_run_ddl() FAILED: SPI_connect_ext returned error code: {}",
connect_result
);
#[allow(unreachable_code)]
return Err(format!(
"SPI_connect_ext failed (error! should diverge): {connect_result}"
));
}
let opts = pg_sys::SPIExecuteOptions {
read_only: false,
allow_nonatomic: true,
tcount: 0,
..pg_sys::SPIExecuteOptions::default()
};
let execute_result =
pg_sys::SPI_execute_extended(c_sql.as_ptr(), std::ptr::from_ref(&opts));
pg_sys::SPI_finish();
if execute_result < 0 {
error!(
"spi_run_ddl() FAILED: SPI_execute_extended error {} for DDL: {}",
execute_result, sql
);
#[allow(unreachable_code)]
return Err(format!(
"SPI_execute_extended failed (error! should diverge): {execute_result}"
));
}
}
log_debug!("spi_run_ddl() succeeded");
Ok(())
}
pub fn spi_get_string(query: &str) -> spi::Result<Option<String>> {
Spi::connect(|client| {
let mut rows = client.select(query, Some(1), &[])?;
match rows.next() {
Some(row) => Ok(row[1].value::<String>()?),
None => Ok(None),
}
})
}
use pgrx::pg_sys::Oid;
pub enum IntExtraction {
Value(i64),
Null,
Missing,
}
pub fn tuple_get_i64(tuple: &PgHeapTuple<'_, AllocatedByPostgres>, col: &str) -> IntExtraction {
match tuple.get_by_name::<i64>(col) {
Ok(Some(v)) => return IntExtraction::Value(v),
Ok(None) => return IntExtraction::Null,
Err(_) => {} }
match tuple.get_by_name::<i32>(col) {
Ok(Some(v)) => IntExtraction::Value(i64::from(v)),
Ok(None) => IntExtraction::Null,
Err(_) => IntExtraction::Missing,
}
}
static OID_QUALIFIED_RELNAME_CACHE: LazyLock<Mutex<HashMap<Oid, String>>> =
LazyLock::new(|| Mutex::new(HashMap::new()));
pub fn invalidate_oid_relname_cache() {
OID_QUALIFIED_RELNAME_CACHE
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.clear();
}
pub static VIEW_COLUMNS_CACHE: LazyLock<Mutex<HashMap<String, Vec<String>>>> =
LazyLock::new(|| Mutex::new(HashMap::new()));
pub fn invalidate_view_columns_cache() {
let mut cache = VIEW_COLUMNS_CACHE
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
cache.clear();
}
pub fn bound_cache<K, V>(cache: &mut HashMap<K, V>) {
if cache.len() >= crate::config::cache_size() {
cache.clear();
}
}
pub fn qualified_relname_from_oid(oid: Oid) -> spi::Result<String> {
{
let cache = OID_QUALIFIED_RELNAME_CACHE
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(name) = cache.get(&oid) {
return Ok(name.clone());
}
}
crate::metrics::metrics_api::record_catalog_lookup();
let qname: String = Spi::connect(|client| {
let args =
vec![unsafe { DatumWithOid::new(oid, PgOid::BuiltIn(PgBuiltInOids::OIDOID).value()) }];
let mut rows = client.select(
"SELECT quote_ident(n.nspname) || '.' || quote_ident(c.relname) AS qname \
FROM pg_class c \
JOIN pg_namespace n ON n.oid = c.relnamespace \
WHERE c.oid = $1",
None,
&args,
)?;
if let Some(row) = rows.next() {
row["qname"].value::<String>()?.ok_or_else(|| {
spi::Error::from(crate::TViewError::SpiError {
query: "qualified_relname_from_oid".to_string(),
error: "qname column is NULL".to_string(),
})
})
} else {
Err(spi::Error::from(crate::TViewError::SpiError {
query: "qualified_relname_from_oid".to_string(),
error: format!("No pg_class entry for oid: {oid:?}"),
}))
}
})?;
{
let mut cache = OID_QUALIFIED_RELNAME_CACHE
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
bound_cache(&mut cache);
cache.insert(oid, qname.clone());
}
Ok(qname)
}
#[must_use]
pub fn qualified_type_name(typid: Oid, typmod: i32) -> String {
#[allow(clippy::cast_possible_truncation)] const FLAGS: u16 =
(pg_sys::FORMAT_TYPE_TYPEMOD_GIVEN | pg_sys::FORMAT_TYPE_FORCE_QUALIFY) as u16;
unsafe {
let name = pg_sys::format_type_extended(typid, typmod, FLAGS);
let out = std::ffi::CStr::from_ptr(name)
.to_string_lossy()
.into_owned();
pg_sys::pfree(name.cast());
out
}
}
pub fn column_types(relid: Oid) -> crate::TViewResult<Vec<(String, String)>> {
Spi::connect(|client| {
let args =
[
unsafe {
pgrx::datum::DatumWithOid::new(relid, pgrx::PgBuiltInOids::OIDOID.value())
},
];
let mut out = Vec::new();
for row in client.select(
"SELECT attname::pg_catalog.text, atttypid, atttypmod FROM pg_catalog.pg_attribute \
WHERE attrelid = $1 AND attnum > 0 AND NOT attisdropped ORDER BY attnum",
None,
&args,
)? {
if let (Some(name), Some(typid), Some(typmod)) = (
row.get::<String>(1)?,
row.get::<Oid>(2)?,
row.get::<i32>(3)?,
) {
out.push((name, qualified_type_name(typid, typmod)));
}
}
Ok::<_, pgrx::spi::Error>(out)
})
.map_err(|e| crate::TViewError::CatalogError {
operation: format!("Read the column types of relation {relid:?}"),
pg_error: e.to_string(),
})
}
const EXT_SCHEMA: &str = "tviews";
thread_local! {
static LOGGED_ONCE: std::cell::RefCell<std::collections::HashSet<String>> =
std::cell::RefCell::new(std::collections::HashSet::new());
}
pub fn log_once(key: &str, message: &str) {
if first_time(key) {
log!("pg_tviews: {message}");
}
}
pub fn first_time(key: &str) -> bool {
LOGGED_ONCE.with(|seen| seen.borrow_mut().insert(key.to_string()))
}
pub fn forget_logged(key: &str) {
LOGGED_ONCE.with(|seen| seen.borrow_mut().remove(key));
}
pub fn meta_table() -> String {
format!("{EXT_SCHEMA}.pg_tview_meta")
}
pub const fn ext_schema() -> &'static str {
EXT_SCHEMA
}
pub fn get_view_columns(schema_name: &str, view_name: &str) -> spi::Result<Vec<String>> {
let cache_key = format!("{schema_name}.{view_name}");
{
let cache = VIEW_COLUMNS_CACHE
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(cols) = cache.get(&cache_key) {
return Ok(cols.clone());
}
}
crate::metrics::metrics_api::record_catalog_lookup();
let cols: Vec<String> = Spi::connect(|client| -> spi::Result<Vec<String>> {
let args = vec![
unsafe {
DatumWithOid::new(schema_name, PgOid::BuiltIn(PgBuiltInOids::TEXTOID).value())
},
unsafe { DatumWithOid::new(view_name, PgOid::BuiltIn(PgBuiltInOids::TEXTOID).value()) },
];
let rows = client.select(
"SELECT a.attname::text \
FROM pg_attribute a \
JOIN pg_class c ON c.oid = a.attrelid \
JOIN pg_namespace n ON n.oid = c.relnamespace \
WHERE n.nspname = $1 AND c.relname = $2 AND a.attnum > 0 AND NOT a.attisdropped \
ORDER BY a.attnum",
None,
&args,
)?;
let mut result = Vec::with_capacity(10);
for r in rows {
if let Some(name) = r["attname"].value::<String>()? {
result.push(name);
}
}
Ok(result)
})?;
{
let mut cache = VIEW_COLUMNS_CACHE
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
bound_cache(&mut cache);
cache.insert(cache_key, cols.clone());
}
Ok(cols)
}
pub fn get_view_columns_by_oid(rel_oid: Oid) -> spi::Result<Vec<String>> {
let oid_key = format!("oid:{}", rel_oid.to_u32());
{
let cache = VIEW_COLUMNS_CACHE
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(cols) = cache.get(&oid_key) {
return Ok(cols.clone());
}
}
crate::metrics::metrics_api::record_catalog_lookup();
let (schema_name, table_name): (String, String) = Spi::connect(|client| {
let args = vec![unsafe {
DatumWithOid::new(rel_oid, PgOid::BuiltIn(PgBuiltInOids::OIDOID).value())
}];
let mut rows = client.select(
"SELECT n.nspname::text, c.relname::text \
FROM pg_class c \
JOIN pg_namespace n ON n.oid = c.relnamespace \
WHERE c.oid = $1",
None,
&args,
)?;
if let Some(row) = rows.next() {
let schema = row["nspname"].value::<String>()?.ok_or_else(|| {
spi::Error::from(crate::TViewError::SpiError {
query: "get_view_columns_by_oid schema lookup".to_string(),
error: "nspname column is NULL".to_string(),
})
})?;
let table = row["relname"].value::<String>()?.ok_or_else(|| {
spi::Error::from(crate::TViewError::SpiError {
query: "get_view_columns_by_oid table lookup".to_string(),
error: "relname column is NULL".to_string(),
})
})?;
Ok((schema, table))
} else {
Err(spi::Error::from(crate::TViewError::SpiError {
query: "get_view_columns_by_oid".to_string(),
error: format!("No pg_class entry for oid: {rel_oid:?}"),
}))
}
})?;
let cols = get_view_columns(&schema_name, &table_name)?;
VIEW_COLUMNS_CACHE
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.insert(oid_key, cols.clone());
Ok(cols)
}
#[must_use]
pub fn quote_identifier(name: &str) -> String {
format!("\"{}\"", name.replace('"', "\"\""))
}
pub const MAX_IDENTIFIER_BYTES: usize = 63;
#[must_use]
pub fn fit_identifier(full: String) -> String {
if full.len() <= MAX_IDENTIFIER_BYTES {
return full;
}
let hash = full.bytes().fold(0x811c_9dc5_u32, |h, b| {
(h ^ u32::from(b)).wrapping_mul(0x0100_0193)
});
let tag = format!("_{hash:08x}");
let mut cut = MAX_IDENTIFIER_BYTES - tag.len();
while !full.is_char_boundary(cut) {
cut -= 1;
}
format!("{}{tag}", &full[..cut])
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_quote_identifier_normal() {
assert_eq!(quote_identifier("post"), "\"post\"");
}
#[test]
fn test_quote_identifier_uppercase() {
assert_eq!(quote_identifier("Post"), "\"Post\"");
}
#[test]
fn test_quote_identifier_with_underscore() {
assert_eq!(quote_identifier("pk_user"), "\"pk_user\"");
}
#[test]
fn test_quote_identifier_with_internal_quotes() {
assert_eq!(quote_identifier("test\"col"), "\"test\"\"col\"");
}
#[test]
fn test_oid_relname_cache_invalidation() {
use pg_sys::Oid;
invalidate_oid_relname_cache();
{
let mut cache = OID_QUALIFIED_RELNAME_CACHE.lock().unwrap();
cache.insert(Oid::from(123), "test_table".to_string());
}
{
let cache = OID_QUALIFIED_RELNAME_CACHE.lock().unwrap();
assert!(cache.get(&Oid::from(123)).is_some());
}
invalidate_oid_relname_cache();
{
let cache = OID_QUALIFIED_RELNAME_CACHE.lock().unwrap();
assert!(cache.is_empty());
}
}
#[test]
fn test_view_columns_cache_invalidation() {
invalidate_view_columns_cache();
{
let mut cache = VIEW_COLUMNS_CACHE.lock().unwrap();
cache.insert(
"public.v_user".to_string(),
vec!["id".to_string(), "name".to_string()],
);
cache.insert(
"public.v_post".to_string(),
vec!["id".to_string(), "title".to_string(), "user_id".to_string()],
);
cache.insert(
"app.v_user".to_string(),
vec!["id".to_string(), "name".to_string(), "org_id".to_string()],
);
}
{
let cache = VIEW_COLUMNS_CACHE.lock().unwrap();
assert_eq!(cache.len(), 3);
assert!(cache.contains_key("public.v_user"));
assert!(cache.contains_key("public.v_post"));
assert!(cache.contains_key("app.v_user"));
}
{
let cache = VIEW_COLUMNS_CACHE.lock().unwrap();
let public_user = cache.get("public.v_user").unwrap();
let app_user = cache.get("app.v_user").unwrap();
assert_eq!(public_user.len(), 2);
assert_eq!(app_user.len(), 3);
}
invalidate_view_columns_cache();
{
let cache = VIEW_COLUMNS_CACHE.lock().unwrap();
assert!(cache.is_empty());
}
}
}