use super::key::{KeyValue, RefreshKey};
use super::state::{TX_CRASH_RECOVERY_CHECKED, TX_REFRESH_QUEUE};
use std::collections::HashSet;
fn check_queue_backpressure(limit: usize) -> Result<(), String> {
let current_size = TX_REFRESH_QUEUE.with(|q| q.borrow().len());
if current_size >= limit {
return Err(format!(
"refresh queue backpressure: queue size ({current_size}) would exceed max_queue_size ({limit})"
));
}
Ok(())
}
pub fn enqueue_refresh_with_limit(entity: &str, key: KeyValue, limit: usize) -> Result<(), String> {
check_queue_backpressure(limit)?;
TX_REFRESH_QUEUE.with(|q| {
q.borrow_mut().insert(RefreshKey::new(entity, key));
});
Ok(())
}
pub fn enqueue_refresh(entity: &str, key: KeyValue) {
if let Err(msg) =
enqueue_refresh_with_limit(entity, key.clone(), crate::config::max_queue_size())
{
pgrx::error!("{}", msg);
}
super::patch::poison(RefreshKey::new(entity, key));
}
pub fn enqueue_refresh_patched(
entity: &str,
pk: i64,
fields: serde_json::Map<String, serde_json::Value>,
) {
if let Err(msg) =
enqueue_refresh_with_limit(entity, KeyValue::Int(pk), crate::config::max_queue_size())
{
pgrx::error!("{}", msg);
}
if super::patch::record(RefreshKey::pk(entity, pk), Vec::new(), fields) {
crate::metrics::metrics_api::record_direct_patch_captured();
}
}
pub fn enqueue_refresh_all(entity: &str) {
TX_REFRESH_QUEUE.with(|q| {
let key = RefreshKey::all(entity);
if q.borrow().contains(&key) {
return;
}
check_queue_backpressure(crate::config::max_queue_size()).unwrap_or_else(|msg| {
pgrx::error!("{}", msg);
});
q.borrow_mut().insert(key);
});
}
pub fn enqueue_refresh_bulk(entity: &str, keys: Vec<KeyValue>) {
TX_REFRESH_QUEUE.with(|q| {
let limit = crate::config::max_queue_size();
check_queue_backpressure(limit).unwrap_or_else(|msg| {
pgrx::error!("{}", msg);
});
let mut queue = q.borrow_mut();
for key in &keys {
queue.insert(RefreshKey::new(entity, key.clone()));
}
});
for key in keys {
super::patch::poison(RefreshKey::new(entity, key));
}
}
pub fn take_queue_snapshot() -> HashSet<RefreshKey> {
TX_REFRESH_QUEUE.with(|q| {
let mut queue = q.borrow_mut();
std::mem::take(&mut *queue)
})
}
pub fn clear_queue() {
TX_REFRESH_QUEUE.with(|q| {
q.borrow_mut().clear();
});
}
pub fn is_crash_recovery_checked(entity: &str) -> bool {
TX_CRASH_RECOVERY_CHECKED.with(|checked| checked.borrow().contains(entity))
}
pub fn mark_crash_recovery_checked(entity: &str) {
TX_CRASH_RECOVERY_CHECKED.with(|checked| {
checked.borrow_mut().insert(entity.to_string());
});
}
pub fn clear_crash_recovery_cache() {
TX_CRASH_RECOVERY_CHECKED.with(|checked| {
checked.borrow_mut().clear();
});
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_enqueue_and_snapshot() {
clear_queue();
enqueue_refresh_with_limit("user", KeyValue::Int(1), usize::MAX).unwrap();
enqueue_refresh_with_limit("post", KeyValue::Int(2), usize::MAX).unwrap();
enqueue_refresh_with_limit("user", KeyValue::Int(1), usize::MAX).unwrap();
let snapshot = take_queue_snapshot();
assert_eq!(snapshot.len(), 2);
let empty_snapshot = take_queue_snapshot();
assert_eq!(empty_snapshot.len(), 0);
}
#[test]
fn test_clear_queue() {
clear_queue();
enqueue_refresh_with_limit("user", KeyValue::Int(1), usize::MAX).unwrap();
enqueue_refresh_with_limit("post", KeyValue::Int(2), usize::MAX).unwrap();
clear_queue();
let snapshot = take_queue_snapshot();
assert_eq!(snapshot.len(), 0);
}
#[test]
fn test_enqueue_respects_max_queue_size() {
clear_queue();
let limit = 2;
enqueue_refresh_with_limit("user", KeyValue::Int(1), limit)
.expect("first insert should succeed");
enqueue_refresh_with_limit("post", KeyValue::Int(2), limit)
.expect("second insert should succeed");
assert!(
enqueue_refresh_with_limit("user", KeyValue::Int(3), limit).is_err(),
"third insert should fail"
);
let snapshot = take_queue_snapshot();
assert_eq!(snapshot.len(), 2);
}
}