use crate::context::DjogiContext;
use crate::error::{DbError, DjogiError};
use crate::live_migrate::backfill::SIDE_EFFECT_SUPPRESSION_TXN_LOCAL;
use crate::live_migrate::plan::{LivePlan, StepKind, StepParameters};
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct DualReadHook {
pub table: String,
pub legacy_column: String,
pub shadow_column: String,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct DualWriteHook {
pub table: String,
pub legacy_column: String,
pub shadow_column: String,
pub codec_transform: Option<String>,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct ActiveHooks {
pub dual_read: Vec<DualReadHook>,
pub dual_write: Vec<DualWriteHook>,
pub side_effects_suppressed: bool,
}
const FIELD_SEPARATOR: &str = "::";
const DUAL_READ_TAG: &str = "dual_read";
const DUAL_WRITE_TAG: &str = "dual_write";
const CODEC_MARKER: &str = "codec";
const SHADOW_SUFFIX: &str = "_new";
#[derive(Debug, Clone, PartialEq, Eq)]
enum ParsedHook {
Read(DualReadHook),
Write(DualWriteHook),
}
fn parse_hook_id(id: &str) -> Result<ParsedHook, HookError> {
let tokens: Vec<&str> = id.split(FIELD_SEPARATOR).collect();
if tokens.iter().any(|t| t.is_empty()) {
return Err(HookError::MalformedHookId(id.to_owned()));
}
match tokens.as_slice() {
[DUAL_READ_TAG, table, column] => Ok(ParsedHook::Read(DualReadHook {
table: (*table).to_owned(),
legacy_column: (*column).to_owned(),
shadow_column: format!("{column}{SHADOW_SUFFIX}"),
})),
[DUAL_READ_TAG, CODEC_MARKER, table, column, _from_codec] => {
Ok(ParsedHook::Read(DualReadHook {
table: (*table).to_owned(),
legacy_column: (*column).to_owned(),
shadow_column: format!("{column}{SHADOW_SUFFIX}"),
}))
}
[DUAL_WRITE_TAG, table, column] => Ok(ParsedHook::Write(DualWriteHook {
table: (*table).to_owned(),
legacy_column: (*column).to_owned(),
shadow_column: format!("{column}{SHADOW_SUFFIX}"),
codec_transform: None,
})),
[DUAL_WRITE_TAG, CODEC_MARKER, table, column, codec_transform] => {
Ok(ParsedHook::Write(DualWriteHook {
table: (*table).to_owned(),
legacy_column: (*column).to_owned(),
shadow_column: format!("{column}{SHADOW_SUFFIX}"),
codec_transform: Some((*codec_transform).to_owned()),
}))
}
_ => Err(HookError::MalformedHookId(id.to_owned())),
}
}
pub fn active_hooks_at_step(plan: &LivePlan, step_ordinal: u32) -> Result<ActiveHooks, HookError> {
let mut snapshot = ActiveHooks::default();
for step in &plan.steps {
if step.ordinal > step_ordinal {
break;
}
snapshot.side_effects_suppressed = false;
match (&step.kind, &step.parameters) {
(
StepKind::BeginCompatibilityWindow,
StepParameters::BeginCompatibilityWindow { hooks },
) => {
for hook_id in hooks {
match parse_hook_id(hook_id)? {
ParsedHook::Read(read_hook) => snapshot.dual_read.push(read_hook),
ParsedHook::Write(write_hook) => snapshot.dual_write.push(write_hook),
}
}
}
(StepKind::CutoverReads, _) => {
snapshot.dual_read.clear();
}
(StepKind::CutoverWrites, _) => {
snapshot.dual_write.clear();
}
(StepKind::BackfillChunked, _) if step.ordinal == step_ordinal => {
snapshot.side_effects_suppressed = true;
}
_ => {}
}
}
Ok(snapshot)
}
pub async fn side_effects_suppressed(ctx: &mut DjogiContext) -> Result<bool, HookError> {
let sql = format!(
"SELECT current_setting('{name}', true)",
name = SIDE_EFFECT_SUPPRESSION_TXN_LOCAL,
);
let row_opt = ctx
.__query_opt_for_macros(&sql, &[])
.await
.map_err(HookError::from)?;
let row = match row_opt {
Some(row) => row,
None => return Ok(false),
};
let value: Option<String> = row
.try_get(0)
.map_err(|e| HookError::from(DjogiError::Db(DbError::other(e.to_string()))))?;
Ok(matches!(
value.as_deref(),
Some("1") | Some("true") | Some("on")
))
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum HookError {
#[error(transparent)]
Database(#[from] DjogiError),
#[error("malformed hook id: {0}")]
MalformedHookId(String),
}
#[cfg(test)]
mod tests {
use super::*;
use crate::live_migrate::plan::{LivePlan, PlanClassification, PlanHeader, Step};
use crate::types::HeerId;
#[test]
fn parse_hook_id_accepts_dual_read_basic_form() {
let parsed = parse_hook_id("dual_read::ledger_entry::amount").unwrap();
let ParsedHook::Read(hook) = parsed else {
panic!("expected ParsedHook::Read");
};
assert_eq!(hook.table, "ledger_entry");
assert_eq!(hook.legacy_column, "amount");
assert_eq!(hook.shadow_column, "amount_new");
}
#[test]
fn parse_hook_id_accepts_dual_write_basic_form() {
let parsed = parse_hook_id("dual_write::ledger_entry::amount").unwrap();
let ParsedHook::Write(hook) = parsed else {
panic!("expected ParsedHook::Write");
};
assert_eq!(hook.table, "ledger_entry");
assert_eq!(hook.legacy_column, "amount");
assert_eq!(hook.shadow_column, "amount_new");
assert!(hook.codec_transform.is_none());
}
#[test]
fn parse_hook_id_accepts_dual_read_codec_form() {
let parsed = parse_hook_id("dual_read::codec::secret::ciphertext::aes_gcm_v1").unwrap();
let ParsedHook::Read(hook) = parsed else {
panic!("expected ParsedHook::Read");
};
assert_eq!(hook.table, "secret");
assert_eq!(hook.legacy_column, "ciphertext");
assert_eq!(hook.shadow_column, "ciphertext_new");
}
#[test]
fn parse_hook_id_accepts_dual_write_codec_form_with_transform() {
let parsed =
parse_hook_id("dual_write::codec::secret::ciphertext::aes_gcm_v1->aes_gcm_v2").unwrap();
let ParsedHook::Write(hook) = parsed else {
panic!("expected ParsedHook::Write");
};
assert_eq!(hook.table, "secret");
assert_eq!(hook.legacy_column, "ciphertext");
assert_eq!(hook.shadow_column, "ciphertext_new");
assert_eq!(
hook.codec_transform.as_deref(),
Some("aes_gcm_v1->aes_gcm_v2"),
);
}
#[test]
fn parse_hook_id_rejects_unknown_family_tag() {
let err = parse_hook_id("dual_replicate::ledger::amount").unwrap_err();
assert!(matches!(err, HookError::MalformedHookId(_)));
}
#[test]
fn parse_hook_id_rejects_empty_token() {
let err = parse_hook_id("dual_read::ledger::").unwrap_err();
assert!(matches!(err, HookError::MalformedHookId(_)));
let err = parse_hook_id("::ledger::amount").unwrap_err();
assert!(matches!(err, HookError::MalformedHookId(_)));
let err = parse_hook_id("dual_read::ledger::::amount").unwrap_err();
assert!(matches!(err, HookError::MalformedHookId(_)));
}
#[test]
fn parse_hook_id_rejects_too_few_tokens() {
let err = parse_hook_id("dual_read").unwrap_err();
assert!(matches!(err, HookError::MalformedHookId(_)));
let err = parse_hook_id("dual_read::ledger").unwrap_err();
assert!(matches!(err, HookError::MalformedHookId(_)));
}
#[test]
fn parse_hook_id_rejects_too_many_tokens_for_basic_form() {
let err = parse_hook_id("dual_read::ledger::amount::extra::tail").unwrap_err();
assert!(matches!(err, HookError::MalformedHookId(_)));
}
#[test]
fn dual_read_hook_round_trips_through_hash_set() {
use std::collections::HashSet;
let mut set: HashSet<DualReadHook> = HashSet::new();
let hook = DualReadHook {
table: "t".to_owned(),
legacy_column: "c".to_owned(),
shadow_column: "c_new".to_owned(),
};
assert!(set.insert(hook.clone()));
assert!(!set.insert(hook.clone()));
assert!(set.contains(&hook));
}
#[test]
fn dual_write_hook_round_trips_through_hash_set() {
use std::collections::HashSet;
let mut set: HashSet<DualWriteHook> = HashSet::new();
let hook = DualWriteHook {
table: "t".to_owned(),
legacy_column: "c".to_owned(),
shadow_column: "c_new".to_owned(),
codec_transform: Some("v1->v2".to_owned()),
};
assert!(set.insert(hook.clone()));
assert!(!set.insert(hook.clone()));
assert!(set.contains(&hook));
}
fn step(kind: StepKind, ordinal: u32, params: StepParameters) -> Step {
Step {
kind,
ordinal,
parameters: params,
}
}
fn full_plan() -> LivePlan {
LivePlan {
header: PlanHeader {
plan_id: HeerId::ZERO,
slug: "demo".to_owned(),
classification: PlanClassification::ExpandContract,
originating_migration: "V20260428000000__demo".to_owned(),
target_database: "main".to_owned(),
app_label: "".to_owned(),
},
steps: vec![
step(
StepKind::ExpandSchema,
0,
StepParameters::ExpandSchema {
sql_segments: vec!["ALTER TABLE t ADD COLUMN c_new INT".to_owned()],
},
),
step(
StepKind::BeginCompatibilityWindow,
1,
StepParameters::BeginCompatibilityWindow {
hooks: vec!["dual_read::t::c".to_owned(), "dual_write::t::c".to_owned()],
},
),
step(
StepKind::BackfillChunked,
2,
StepParameters::BackfillChunked {
table: "t".to_owned(),
predicate_template: "WHERE c_new IS NULL LIMIT $1 RETURNING id".to_owned(),
chunk_size: 1000,
},
),
step(
StepKind::ValidateBackfill,
3,
StepParameters::ValidateBackfill {
gate_query: "SELECT count(*) FROM t WHERE c_new IS NULL".to_owned(),
},
),
step(
StepKind::CutoverReads,
4,
StepParameters::CutoverReads {
description: "flip reads".to_owned(),
},
),
step(
StepKind::CutoverWrites,
5,
StepParameters::CutoverWrites {
description: "flip writes".to_owned(),
},
),
step(
StepKind::CleanupLegacyState,
6,
StepParameters::CleanupLegacyState {
sql_segments: vec!["ALTER TABLE t DROP COLUMN c".to_owned()],
},
),
],
}
}
#[test]
fn active_hooks_empty_before_compatibility_window_opens() {
let plan = full_plan();
let snap = active_hooks_at_step(&plan, 0).unwrap();
assert!(snap.dual_read.is_empty());
assert!(snap.dual_write.is_empty());
assert!(!snap.side_effects_suppressed);
}
#[test]
fn active_hooks_populated_after_begin_compatibility_window() {
let plan = full_plan();
let snap = active_hooks_at_step(&plan, 1).unwrap();
assert_eq!(snap.dual_read.len(), 1);
assert_eq!(snap.dual_read[0].legacy_column, "c");
assert_eq!(snap.dual_read[0].shadow_column, "c_new");
assert_eq!(snap.dual_write.len(), 1);
assert_eq!(snap.dual_write[0].legacy_column, "c");
assert_eq!(snap.dual_write[0].shadow_column, "c_new");
assert!(!snap.side_effects_suppressed);
}
#[test]
fn active_hooks_side_effects_suppressed_during_backfill_step() {
let plan = full_plan();
let snap = active_hooks_at_step(&plan, 2).unwrap();
assert!(
snap.side_effects_suppressed,
"side_effects_suppressed must be set while a BackfillChunked step is active",
);
assert_eq!(snap.dual_read.len(), 1);
assert_eq!(snap.dual_write.len(), 1);
}
#[test]
fn active_hooks_side_effects_suppressed_resets_after_backfill_step() {
let plan = full_plan();
let snap = active_hooks_at_step(&plan, 3).unwrap();
assert!(
!snap.side_effects_suppressed,
"side_effects_suppressed must reset at the next step boundary",
);
assert_eq!(snap.dual_read.len(), 1);
assert_eq!(snap.dual_write.len(), 1);
}
#[test]
fn active_hooks_dual_read_dropped_after_cutover_reads() {
let plan = full_plan();
let snap = active_hooks_at_step(&plan, 4).unwrap();
assert!(
snap.dual_read.is_empty(),
"CutoverReads must drop dual_read entries",
);
assert_eq!(
snap.dual_write.len(),
1,
"CutoverReads must NOT drop dual_write entries",
);
}
#[test]
fn active_hooks_dual_write_dropped_after_cutover_writes() {
let plan = full_plan();
let snap = active_hooks_at_step(&plan, 5).unwrap();
assert!(snap.dual_read.is_empty());
assert!(
snap.dual_write.is_empty(),
"CutoverWrites must drop dual_write entries",
);
}
#[test]
fn active_hooks_terminal_state_clean() {
let plan = full_plan();
let snap = active_hooks_at_step(&plan, 6).unwrap();
assert!(snap.dual_read.is_empty());
assert!(snap.dual_write.is_empty());
assert!(!snap.side_effects_suppressed);
}
#[test]
fn active_hooks_ordinal_beyond_plan_saturates_to_terminal_state() {
let plan = full_plan();
let last_ordinal: u32 = plan.steps.last().unwrap().ordinal;
let terminal = active_hooks_at_step(&plan, last_ordinal).unwrap();
let beyond = active_hooks_at_step(&plan, last_ordinal + 100).unwrap();
assert_eq!(terminal, beyond);
}
#[test]
fn active_hooks_walker_propagates_malformed_hook_id() {
let mut plan = full_plan();
plan.steps[1].parameters = StepParameters::BeginCompatibilityWindow {
hooks: vec!["this_is_not_a_valid_hook".to_owned()],
};
let err = active_hooks_at_step(&plan, 1).unwrap_err();
assert!(matches!(err, HookError::MalformedHookId(_)));
}
#[test]
fn active_hooks_codec_form_round_trips_through_walker() {
let mut plan = full_plan();
plan.steps[1].parameters = StepParameters::BeginCompatibilityWindow {
hooks: vec![
"dual_read::codec::secret::ciphertext::aes_gcm_v1".to_owned(),
"dual_write::codec::secret::ciphertext::aes_gcm_v1->aes_gcm_v2".to_owned(),
],
};
let snap = active_hooks_at_step(&plan, 1).unwrap();
assert_eq!(snap.dual_read.len(), 1);
assert_eq!(snap.dual_read[0].table, "secret");
assert_eq!(snap.dual_read[0].legacy_column, "ciphertext");
assert_eq!(snap.dual_read[0].shadow_column, "ciphertext_new");
assert_eq!(snap.dual_write.len(), 1);
assert_eq!(
snap.dual_write[0].codec_transform.as_deref(),
Some("aes_gcm_v1->aes_gcm_v2"),
);
}
#[test]
fn active_hooks_walker_aggregates_multiple_compat_windows() {
let plan = LivePlan {
header: PlanHeader {
plan_id: HeerId::ZERO,
slug: "stacked".to_owned(),
classification: PlanClassification::ExpandContract,
originating_migration: "V20260428000000__stacked".to_owned(),
target_database: "main".to_owned(),
app_label: "".to_owned(),
},
steps: vec![
step(
StepKind::BeginCompatibilityWindow,
0,
StepParameters::BeginCompatibilityWindow {
hooks: vec!["dual_read::t::a".to_owned()],
},
),
step(
StepKind::BeginCompatibilityWindow,
1,
StepParameters::BeginCompatibilityWindow {
hooks: vec!["dual_write::t::b".to_owned()],
},
),
],
};
let snap = active_hooks_at_step(&plan, 1).unwrap();
assert_eq!(snap.dual_read.len(), 1);
assert_eq!(snap.dual_read[0].legacy_column, "a");
assert_eq!(snap.dual_write.len(), 1);
assert_eq!(snap.dual_write[0].legacy_column, "b");
}
#[test]
fn side_effect_suppression_const_shared_with_backfill_module() {
assert!(
SIDE_EFFECT_SUPPRESSION_TXN_LOCAL.starts_with("djogi."),
"GUC name must live under djogi.* namespace: {SIDE_EFFECT_SUPPRESSION_TXN_LOCAL:?}",
);
}
#[test]
fn hook_error_malformed_id_message_includes_offending_value() {
let err = HookError::MalformedHookId("garbage".to_owned());
let msg = format!("{err}");
assert!(
msg.contains("garbage"),
"expected offending id in message: {msg}",
);
assert!(
msg.contains("malformed"),
"expected `malformed` hint: {msg}",
);
}
}