use fsqlite_error::FrankenError;
use fsqlite_types::value::SqliteValue;
pub const SQLITE_MAX_TRIGGER_DEPTH: usize = fsqlite_types::limits::MAX_TRIGGER_DEPTH as usize;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RaiseResult {
Ignore,
Rollback(String),
Abort(String),
Fail(String),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PseudoTableMapping {
pub old_base: Option<i32>,
pub new_base: Option<i32>,
pub num_columns: i32,
}
#[derive(Debug, Clone)]
pub struct VdbeFrame {
pub saved_pc: i32,
pub registers: Vec<SqliteValue>,
pub n_cursor: i32,
pub subprogram_idx: i32,
pub trigger_name: String,
pub pseudo_tables: Option<PseudoTableMapping>,
pub raise_result: Option<RaiseResult>,
}
impl VdbeFrame {
#[allow(clippy::cast_possible_truncation)]
pub fn estimated_memory(&self) -> usize {
let base_overhead = std::mem::size_of::<Self>();
let reg_mem: usize = self
.registers
.iter()
.map(|v| {
std::mem::size_of::<SqliteValue>()
+ match v {
SqliteValue::Text(s) => s.len(),
SqliteValue::Blob(b) => b.len(),
_ => 0,
}
})
.sum();
let name_mem = self.trigger_name.len();
base_overhead + reg_mem + name_mem
}
}
#[derive(Debug)]
pub struct FrameStack {
frames: Vec<(VdbeFrame, usize)>,
max_depth: usize,
cx_memory_budget: usize,
current_memory: usize,
recursive_triggers: bool,
}
impl FrameStack {
pub fn new(max_depth: usize, cx_memory_budget: usize) -> Self {
Self {
frames: Vec::new(),
max_depth,
cx_memory_budget,
current_memory: 0,
recursive_triggers: false,
}
}
pub fn with_defaults() -> Self {
Self::new(SQLITE_MAX_TRIGGER_DEPTH, 64 * 1024 * 1024)
}
pub fn depth(&self) -> usize {
self.frames.len()
}
pub fn is_empty(&self) -> bool {
self.frames.is_empty()
}
pub fn set_recursive_triggers(&mut self, enabled: bool) {
self.recursive_triggers = enabled;
}
pub fn recursive_triggers(&self) -> bool {
self.recursive_triggers
}
pub fn current_memory(&self) -> usize {
self.current_memory
}
pub fn push_frame(&mut self, frame: VdbeFrame) -> Result<(), FrankenError> {
if self.frames.len() >= self.max_depth {
return Err(FrankenError::Internal(format!(
"trigger depth limit exceeded (max {})",
self.max_depth
)));
}
if !self.recursive_triggers {
let is_self_recursive = self
.frames
.iter()
.any(|(f, _)| f.trigger_name == frame.trigger_name);
if is_self_recursive {
return Err(FrankenError::Internal(
"recursive triggers are disabled (PRAGMA recursive_triggers = OFF)".to_owned(),
));
}
}
let frame_mem = frame.estimated_memory();
let new_total = self
.current_memory
.checked_add(frame_mem)
.ok_or(FrankenError::OutOfMemory)?;
if new_total > self.cx_memory_budget {
return Err(FrankenError::OutOfMemory);
}
self.current_memory = new_total;
self.frames.push((frame, frame_mem));
Ok(())
}
pub fn pop_frame(&mut self) -> Option<VdbeFrame> {
let (frame, pushed_mem) = self.frames.pop()?;
self.current_memory = self.current_memory.saturating_sub(pushed_mem);
Some(frame)
}
pub fn top(&self) -> Option<&VdbeFrame> {
self.frames.last().map(|(f, _)| f)
}
pub fn update_top(
&mut self,
update: impl FnOnce(&mut VdbeFrame),
) -> Result<bool, FrankenError> {
let Some((top, recorded_mem)) = self.frames.last_mut() else {
return Ok(false);
};
let mut candidate = top.clone();
update(&mut candidate);
let new_mem = candidate.estimated_memory();
let base_memory = self.current_memory.saturating_sub(*recorded_mem);
let new_total = base_memory.saturating_add(new_mem);
if new_total > self.cx_memory_budget {
return Err(FrankenError::OutOfMemory);
}
*top = candidate;
*recorded_mem = new_mem;
self.current_memory = new_total;
Ok(true)
}
pub fn unwind_all(&mut self) -> Vec<VdbeFrame> {
let mut unwound = Vec::with_capacity(self.frames.len());
while let Some(frame) = self.pop_frame() {
unwound.push(frame);
}
unwound
}
pub fn unwind_to_index(&mut self, target_index: usize) -> Result<Vec<VdbeFrame>, FrankenError> {
if target_index >= self.frames.len() {
return Err(FrankenError::Internal(format!(
"no frame at index {target_index}"
)));
}
let mut unwound = Vec::with_capacity(self.frames.len() - target_index);
while self.frames.len() > target_index {
if let Some(frame) = self.pop_frame() {
unwound.push(frame);
}
}
Ok(unwound)
}
}
pub fn make_frame(
saved_pc: i32,
num_registers: usize,
n_cursor: i32,
subprogram_idx: i32,
trigger_name: impl Into<String>,
) -> VdbeFrame {
VdbeFrame {
saved_pc,
registers: vec![SqliteValue::Null; num_registers],
n_cursor,
subprogram_idx,
trigger_name: trigger_name.into(),
pseudo_tables: None,
raise_result: None,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_trigger_depth_limit_1000() {
let mut stack = FrameStack::new(SQLITE_MAX_TRIGGER_DEPTH, usize::MAX);
stack.set_recursive_triggers(true);
for i in 0..SQLITE_MAX_TRIGGER_DEPTH {
let frame = make_frame(
i32::try_from(i).unwrap(),
4, 0,
1,
"trg_recursive",
);
stack
.push_frame(frame)
.unwrap_or_else(|e| unreachable!("push at depth {i} should succeed: {e}"));
}
assert_eq!(stack.depth(), SQLITE_MAX_TRIGGER_DEPTH);
let overflow_frame = make_frame(1000, 4, 0, 1, "trg_recursive");
let err = stack
.push_frame(overflow_frame)
.expect_err("push at depth 1001 must fail");
assert!(
err.to_string().contains("trigger depth limit exceeded"),
"expected depth limit error, got: {err}"
);
}
#[test]
fn test_trigger_no_stack_overflow_at_max_depth() {
let mut stack = FrameStack::new(SQLITE_MAX_TRIGGER_DEPTH, usize::MAX);
stack.set_recursive_triggers(true);
for i in 0..SQLITE_MAX_TRIGGER_DEPTH {
let frame = make_frame(
i32::try_from(i).unwrap(),
8, 2,
1,
"trg_deep",
);
stack.push_frame(frame).unwrap();
}
assert_eq!(stack.depth(), 1000);
let unwound = stack.unwind_all();
assert_eq!(unwound.len(), 1000);
assert!(stack.is_empty());
assert_eq!(stack.current_memory(), 0);
}
#[test]
fn test_trigger_cx_memory_budget_enforced() {
let frame_size_estimate = {
let sample = make_frame(0, 1000, 0, 1, "trg_mem");
sample.estimated_memory()
};
let budget = frame_size_estimate * 10;
let mut stack = FrameStack::new(SQLITE_MAX_TRIGGER_DEPTH, budget);
stack.set_recursive_triggers(true);
let mut pushed = 0;
for i in 0..SQLITE_MAX_TRIGGER_DEPTH {
let frame = make_frame(i32::try_from(i).unwrap(), 1000, 0, 1, "trg_mem");
match stack.push_frame(frame) {
Ok(()) => pushed += 1,
Err(e) => {
assert!(
matches!(e, FrankenError::OutOfMemory),
"expected OutOfMemory, got: {e}"
);
break;
}
}
}
assert!(
pushed < SQLITE_MAX_TRIGGER_DEPTH,
"expected budget to stop nesting, but pushed all {pushed}"
);
assert!(
pushed <= 10,
"expected ~10 frames allowed, but got {pushed}"
);
assert!(
pushed >= 10,
"expected at least 10 frames, but got {pushed}"
);
}
#[test]
fn test_trigger_recursive_off_prevents_self_fire() {
let mut stack = FrameStack::with_defaults();
assert!(!stack.recursive_triggers());
let frame1 = make_frame(0, 4, 0, 1, "trg_self");
stack.push_frame(frame1).expect("first push should succeed");
assert_eq!(stack.depth(), 1);
let frame2 = make_frame(1, 4, 0, 1, "trg_self");
let err = stack
.push_frame(frame2)
.expect_err("self-recursive push must fail when recursive_triggers=OFF");
assert!(
err.to_string().contains("recursive triggers are disabled"),
"expected recursion-disabled error, got: {err}"
);
let frame3 = make_frame(2, 4, 0, 2, "trg_other");
stack
.push_frame(frame3)
.expect("different trigger should succeed");
assert_eq!(stack.depth(), 2);
}
#[test]
fn test_trigger_frame_stack_cleanup_on_error() {
let mut stack = FrameStack::with_defaults();
stack.set_recursive_triggers(true);
for i in 0..5 {
let frame = make_frame(i, 16, 1, i, format!("trg_chain_{i}"));
stack.push_frame(frame).unwrap();
}
assert_eq!(stack.depth(), 5);
stack
.update_top(|frame| {
frame.raise_result = Some(RaiseResult::Abort("constraint failed".to_owned()));
})
.expect("raise_result update should succeed");
let top = stack.top().unwrap();
assert!(matches!(
&top.raise_result,
Some(RaiseResult::Abort(msg)) if msg == "constraint failed"
));
let unwound = stack.unwind_all();
assert_eq!(unwound.len(), 5);
assert!(stack.is_empty());
assert_eq!(stack.depth(), 0);
assert_eq!(stack.current_memory(), 0);
assert_eq!(unwound[0].trigger_name, "trg_chain_4");
assert_eq!(unwound[4].trigger_name, "trg_chain_0");
assert!(unwound[0].raise_result.is_some());
}
#[test]
fn test_trigger_old_new_pseudo_tables() {
let mut stack = FrameStack::with_defaults();
let mut parent_frame = make_frame(100, 7, 1, 0, "trg_update");
parent_frame.registers[1] = SqliteValue::Integer(10); parent_frame.registers[2] = SqliteValue::Text("hello".into()); parent_frame.registers[3] = SqliteValue::Float(std::f64::consts::PI); parent_frame.registers[4] = SqliteValue::Integer(20); parent_frame.registers[5] = SqliteValue::Text("world".into()); parent_frame.registers[6] = SqliteValue::Float(2.72); parent_frame.pseudo_tables = Some(PseudoTableMapping {
old_base: Some(1),
new_base: Some(4),
num_columns: 3,
});
stack.push_frame(parent_frame).unwrap();
let trigger_frame = make_frame(0, 4, 0, 1, "trg_before_update");
stack.push_frame(trigger_frame).unwrap();
let parent = &stack.frames[0].0;
let mapping = parent.pseudo_tables.as_ref().unwrap();
let old_base = usize::try_from(mapping.old_base.expect("old_base must exist"))
.expect("old_base must be non-negative");
assert!(matches!(
parent.registers[old_base],
SqliteValue::Integer(10)
));
let new_base = usize::try_from(mapping.new_base.expect("new_base must exist"))
.expect("new_base must be non-negative");
assert!(matches!(&parent.registers[new_base + 1], SqliteValue::Text(s) if &**s == "world"));
let parent_mut = &mut stack.frames[0].0;
let new_base_mut = usize::try_from(
parent_mut
.pseudo_tables
.as_ref()
.unwrap()
.new_base
.expect("new_base must exist"),
)
.expect("new_base must be non-negative");
parent_mut.registers[new_base_mut] = SqliteValue::Integer(999);
let parent = &stack.frames[0].0;
let new_base = usize::try_from(
parent
.pseudo_tables
.as_ref()
.unwrap()
.new_base
.expect("new_base must exist"),
)
.expect("new_base must be non-negative");
assert!(matches!(
parent.registers[new_base],
SqliteValue::Integer(999)
));
let old_base = usize::try_from(
parent
.pseudo_tables
.as_ref()
.unwrap()
.old_base
.expect("old_base must exist"),
)
.expect("old_base must be non-negative");
assert!(matches!(
parent.registers[old_base],
SqliteValue::Integer(10)
));
stack.unwind_all();
assert!(stack.is_empty());
}
#[test]
fn test_update_top_rejects_budget_bypass() {
let base_frame = make_frame(0, 1, 0, 0, "trg_budget");
let budget = base_frame.estimated_memory() + 16;
let mut stack = FrameStack::new(SQLITE_MAX_TRIGGER_DEPTH, budget);
stack.push_frame(base_frame.clone()).unwrap();
let before_memory = stack.current_memory();
let before_frame = stack.top().unwrap().clone();
let err = stack
.update_top(|frame| {
frame.registers[0] = SqliteValue::Text("x".repeat(512).into());
})
.expect_err("budget-busting top mutation must fail");
assert!(matches!(err, FrankenError::OutOfMemory));
assert_eq!(stack.current_memory(), before_memory);
assert_eq!(
stack.top().unwrap().estimated_memory(),
before_frame.estimated_memory()
);
assert!(matches!(
stack.top().unwrap().registers[0],
SqliteValue::Null
));
}
#[test]
fn test_push_frame_rejects_accounting_overflow() {
let frame = make_frame(0, 1, 0, 0, "trg_overflow");
let frame_mem = frame.estimated_memory();
assert!(frame_mem > 0);
let preloaded_memory = usize::MAX - (frame_mem - 1);
let mut stack = FrameStack::new(SQLITE_MAX_TRIGGER_DEPTH, usize::MAX);
stack.current_memory = preloaded_memory;
let err = stack
.push_frame(frame)
.expect_err("overflowed memory accounting must fail cleanly");
assert!(matches!(err, FrankenError::OutOfMemory));
assert_eq!(stack.depth(), 0);
assert_eq!(stack.current_memory(), preloaded_memory);
}
#[test]
fn test_update_top_refreshes_memory_accounting() {
let base_frame = make_frame(0, 1, 0, 0, "trg_budget_ok");
let mut stack = FrameStack::new(SQLITE_MAX_TRIGGER_DEPTH, usize::MAX);
stack.push_frame(base_frame).unwrap();
stack
.update_top(|frame| {
frame.trigger_name.push_str("_expanded");
frame.registers[0] = SqliteValue::Blob(vec![7; 128].into());
})
.expect("top mutation should succeed");
let top = stack.top().unwrap();
assert_eq!(stack.current_memory(), top.estimated_memory());
assert_eq!(top.trigger_name, "trg_budget_ok_expanded");
assert!(matches!(&top.registers[0], SqliteValue::Blob(bytes) if bytes.len() == 128));
}
#[test]
fn test_raise_result_variants() {
let ignore = RaiseResult::Ignore;
let rollback = RaiseResult::Rollback("oops".to_owned());
let abort = RaiseResult::Abort("error".to_owned());
let fail = RaiseResult::Fail("bad".to_owned());
assert_eq!(ignore, RaiseResult::Ignore);
assert_ne!(rollback, abort);
assert!(matches!(fail, RaiseResult::Fail(msg) if msg == "bad"));
}
#[test]
fn test_pseudo_table_mapping_insert_trigger() {
let mapping = PseudoTableMapping {
old_base: None,
new_base: Some(1),
num_columns: 3,
};
assert!(mapping.old_base.is_none());
assert_eq!(mapping.new_base, Some(1));
}
#[test]
fn test_pseudo_table_mapping_delete_trigger() {
let mapping = PseudoTableMapping {
old_base: Some(1),
new_base: None,
num_columns: 3,
};
assert_eq!(mapping.old_base, Some(1));
assert!(mapping.new_base.is_none());
}
#[test]
fn test_unwind_to_index() {
let mut stack = FrameStack::with_defaults();
stack.set_recursive_triggers(true);
for i in 0..5 {
let frame = make_frame(i, 4, 0, i, format!("trg_{i}"));
stack.push_frame(frame).unwrap();
}
let unwound = stack.unwind_to_index(2).unwrap();
assert_eq!(unwound.len(), 3);
assert_eq!(unwound[0].trigger_name, "trg_4");
assert_eq!(unwound[1].trigger_name, "trg_3");
assert_eq!(unwound[2].trigger_name, "trg_2");
assert_eq!(stack.depth(), 2); }
#[test]
fn test_unwind_to_index_missing_target_errors() {
let mut stack = FrameStack::with_defaults();
stack.push_frame(make_frame(0, 4, 0, 0, "trg_0")).unwrap();
stack.push_frame(make_frame(1, 4, 0, 1, "trg_1")).unwrap();
let err = stack
.unwind_to_index(2)
.expect_err("missing target index must not unwind the full stack");
assert!(err.to_string().contains("no frame at index 2"));
assert_eq!(stack.depth(), 2);
}
#[test]
fn test_pop_empty_stack() {
let mut stack = FrameStack::with_defaults();
assert!(stack.pop_frame().is_none());
assert!(stack.top().is_none());
}
#[test]
fn test_frame_estimated_memory_with_text_blob() {
let mut frame = make_frame(0, 3, 0, 0, "trg_mem_test");
frame.registers[0] = SqliteValue::Text("a".repeat(1024).into());
frame.registers[1] = SqliteValue::Blob(vec![0u8; 2048].into());
frame.registers[2] = SqliteValue::Integer(42);
let mem = frame.estimated_memory();
assert!(mem > 3000, "memory estimate too low: {mem}");
}
}