use crate::error::{TViewError, TViewResult};
use crate::queue::affected::{self, Change, NetChange};
use crate::utils::quote_identifier;
use pgrx::JsonB;
use pgrx::datum::DatumWithOid;
use pgrx::prelude::*;
use serde_json::{Map, Value, json};
use std::collections::{BTreeSet, HashMap};
#[pg_extern]
fn pg_tviews_flush_and_report(
max_entities: default!(i32, 500),
include_data: default!(bool, true),
reset: default!(bool, true),
) -> Result<JsonB, TViewError> {
crate::revision::check();
crate::queue::flush_refresh_queue()?;
let (changes, overflow) = affected::summarize(reset);
let limit = usize::try_from(max_entities).unwrap_or(0);
Ok(JsonB(build_report(
&changes,
&overflow,
limit,
include_data,
)?))
}
#[pg_extern]
fn pg_tviews_set_typename(entity: &str, typename: Option<&str>) -> Result<(), TViewError> {
crate::revision::check();
if let Some(name) = typename {
let valid = name
.chars()
.next()
.is_some_and(|c| c.is_ascii_alphabetic() || c == '_')
&& name.chars().all(|c| c.is_ascii_alphanumeric() || c == '_');
if !valid {
return Err(TViewError::InvalidInput {
parameter: "typename".to_string(),
reason: format!("'{name}' is not a GraphQL name ([_A-Za-z][_0-9A-Za-z]*)"),
});
}
}
let meta = crate::catalog::TviewMeta::load_by_entity(entity)?.ok_or_else(|| {
TViewError::MetadataNotFound {
entity: entity.to_string(),
}
})?;
crate::owner::require_owner(meta.tview_oid, &format!("tv_{entity}"))?;
let _owner = crate::owner::AsOwner::of_extension()?;
let updated = Spi::connect_mut(|client| {
let args = [
unsafe { DatumWithOid::new(entity, PgOid::BuiltIn(PgBuiltInOids::TEXTOID).value()) },
unsafe { DatumWithOid::new(typename, PgOid::BuiltIn(PgBuiltInOids::TEXTOID).value()) },
];
client
.update(
&format!(
"UPDATE {} SET graphql_typename = $2 WHERE entity = $1",
crate::utils::meta_table()
),
None,
&args,
)
.map(|t| t.len())
})
.map_err(|e| TViewError::CatalogError {
operation: "Set GraphQL type name".to_string(),
pg_error: e.to_string(),
})?;
if updated == 0 {
return Err(TViewError::MetadataNotFound {
entity: entity.to_string(),
});
}
Ok(())
}
fn pascal_case(entity: &str) -> String {
entity
.split('_')
.filter(|part| !part.is_empty())
.map(|part| {
let mut chars = part.chars();
chars.next().map_or_else(String::new, |first| {
first.to_uppercase().collect::<String>() + chars.as_str()
})
})
.collect()
}
struct EntityInfo {
typename: String,
table: String,
}
fn entity_info(entities: &BTreeSet<&str>) -> TViewResult<HashMap<String, EntityInfo>> {
let names: Vec<String> = entities.iter().map(|e| (*e).to_string()).collect();
Spi::connect(|client| {
let args = [unsafe {
DatumWithOid::new(names, PgOid::BuiltIn(PgBuiltInOids::TEXTARRAYOID).value())
}];
let mut out = HashMap::new();
for row in client.select(
&format!(
"SELECT m.entity, m.graphql_typename, \
quote_ident(n.nspname) || '.' || quote_ident(c.relname) AS tbl \
FROM {} m \
JOIN pg_class c ON c.oid = m.table_oid \
JOIN pg_namespace n ON n.oid = c.relnamespace \
WHERE m.entity = ANY($1)",
crate::utils::meta_table()
),
None,
&args,
)? {
let entity: String = row["entity"].value()?.unwrap_or_default();
let typename: Option<String> = row["graphql_typename"].value()?;
let table: String = row["tbl"].value()?.unwrap_or_default();
out.insert(
entity.clone(),
EntityInfo {
typename: typename.unwrap_or_else(|| pascal_case(&entity)),
table,
},
);
}
Ok::<_, spi::Error>(out)
})
.map_err(|e| TViewError::CatalogError {
operation: "Load TVIEWs for the report".to_string(),
pg_error: e.to_string(),
})
}
fn current_rows(
entity: &str,
table: &str,
pks: Vec<String>,
) -> TViewResult<HashMap<String, (Value, Value)>> {
let qi_pk = quote_identifier(&format!("pk_{entity}"));
let sql = format!(
"SELECT t.{qi_pk}::text AS k, to_jsonb(t.*) AS r FROM {table} t \
WHERE t.{qi_pk}::text = ANY($1)"
);
Spi::connect(|client| {
let args = [unsafe {
DatumWithOid::new(pks, PgOid::BuiltIn(PgBuiltInOids::TEXTARRAYOID).value())
}];
let mut out = HashMap::new();
for row in client.select(&sql, None, &args)? {
let key: Option<String> = row["k"].value()?;
let rec: Option<JsonB> = row["r"].value()?;
if let (Some(key), Some(JsonB(Value::Object(mut rec)))) = (key, rec) {
let id = rec.remove("id").unwrap_or(Value::Null);
let data = rec.remove("data").unwrap_or(Value::Null);
out.insert(key, (id, data));
}
}
Ok::<_, spi::Error>(out)
})
.map_err(|e| TViewError::SpiError {
query: sql,
error: e.to_string(),
})
}
fn build_report(
changes: &[NetChange],
overflow: &BTreeSet<String>,
max_entities: usize,
include_data: bool,
) -> TViewResult<Value> {
let entities: BTreeSet<&str> = changes
.iter()
.map(|c| c.entity.as_str())
.chain(overflow.iter().map(String::as_str))
.collect();
let info = entity_info(&entities)?;
let typename = |entity: &str| {
info.get(entity)
.map_or_else(|| pascal_case(entity), |i| i.typename.clone())
};
let mut changes = changes.to_vec();
changes.sort_by(|a, b| {
a.entity
.cmp(&b.entity)
.then_with(|| match (a.pk.parse::<i64>(), b.pk.parse::<i64>()) {
(Ok(x), Ok(y)) => x.cmp(&y),
_ => a.pk.cmp(&b.pk),
})
});
let (reported, left_out) = changes.split_at(changes.len().min(max_entities));
let mut pks_by_entity: HashMap<&str, Vec<String>> = HashMap::new();
for c in reported {
if !matches!(c.change, Change::Deleted(_)) {
pks_by_entity
.entry(&c.entity)
.or_default()
.push(c.pk.clone());
}
}
let mut rows: HashMap<(String, String), (Value, Value)> = HashMap::new();
for (entity, pks) in pks_by_entity {
if let Some(i) = info.get(entity) {
for (pk, row) in current_rows(entity, &i.table, pks)? {
rows.insert((entity.to_string(), pk), row);
}
}
}
let mut updated = Vec::new();
let mut deleted = Vec::new();
for c in reported {
let mut item = Map::new();
item.insert("__typename".into(), Value::String(typename(&c.entity)));
match &c.change {
Change::Deleted(id) => {
item.insert("id".into(), id.clone().map_or(Value::Null, Value::String));
deleted.push(Value::Object(item));
}
change => {
let Some((id, data)) = rows.remove(&(c.entity.clone(), c.pk.clone())) else {
continue;
};
let op = if *change == Change::Inserted {
"CREATED"
} else {
"UPDATED"
};
item.insert("id".into(), id);
item.insert("operation".into(), Value::String(op.into()));
if include_data {
item.insert("data".into(), data);
}
updated.push(Value::Object(item));
}
}
}
let invalidated: BTreeSet<String> = left_out
.iter()
.map(|c| c.entity.as_str())
.chain(overflow.iter().map(String::as_str))
.map(typename)
.collect();
Ok(json!({
"updated": updated,
"deleted": deleted,
"truncated": !invalidated.is_empty(),
"invalidated_types": invalidated.into_iter().collect::<Vec<_>>(),
}))
}
#[cfg(test)]
mod tests {
use super::pascal_case;
#[test]
fn pascal_case_joins_snake_parts() {
assert_eq!(pascal_case("post"), "Post");
assert_eq!(pascal_case("blog_post"), "BlogPost");
assert_eq!(pascal_case("order_line_item"), "OrderLineItem");
}
}