use std::sync::Arc;
use std::sync::OnceLock;
use nmbrs_metrics::labels::Labels;
pub trait OpFieldModifier<T>: Send + Sync + 'static {
fn field_name(&self) -> &'static str;
fn apply(&self, target: &mut T);
fn diagnostic_value(&self) -> serde_json::Value;
}
pub struct ModifierChain<T> {
op_label: String,
active: Vec<Box<dyn OpFieldModifier<T>>>,
event_sink: Option<Arc<dyn ModifierTraceSink>>,
}
impl<T: 'static> ModifierChain<T> {
pub fn new(
op_label: impl Into<String>,
active: Vec<Box<dyn OpFieldModifier<T>>>,
event_sink: Option<Arc<dyn ModifierTraceSink>>,
) -> Self {
Self {
op_label: op_label.into(),
active,
event_sink,
}
}
pub fn is_empty(&self) -> bool {
self.active.is_empty()
}
pub fn len(&self) -> usize {
self.active.len()
}
pub fn op_label(&self) -> &str {
&self.op_label
}
pub fn apply(&self, target: &mut T) {
match &self.event_sink {
None => {
for m in &self.active {
m.apply(target);
}
}
Some(sink) => {
for m in &self.active {
m.apply(target);
sink.modifier_applied(&self.op_label, m.field_name(), &|| m.diagnostic_value());
}
}
}
}
}
impl<T: 'static> std::fmt::Debug for ModifierChain<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ModifierChain")
.field("op_label", &self.op_label)
.field("active_count", &self.active.len())
.field("event_sink", &self.event_sink.is_some())
.finish()
}
}
pub trait ModifierTraceSink: Send + Sync {
fn modifier_applied(
&self,
op: &str,
field: &'static str,
value_fn: &dyn Fn() -> serde_json::Value,
);
}
pub struct TraceRouterSink;
impl ModifierTraceSink for TraceRouterSink {
fn modifier_applied(
&self,
op: &str,
field: &'static str,
value_fn: &dyn Fn() -> serde_json::Value,
) {
if !crate::trace_router::enabled() {
return;
}
let value = value_fn();
let labels = Labels::of("component", "op_modifier")
.with("op", op.to_string())
.with("field", field);
let message = format!("{field}={value}");
crate::trace_router::log(&labels, &message);
}
}
static SESSION_MODIFIER_SINK: OnceLock<Arc<dyn ModifierTraceSink>> = OnceLock::new();
pub fn install_session_sink(sink: Arc<dyn ModifierTraceSink>) {
let _ = SESSION_MODIFIER_SINK.set(sink);
}
pub fn session_sink() -> Option<Arc<dyn ModifierTraceSink>> {
SESSION_MODIFIER_SINK.get().cloned()
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Mutex;
use std::sync::atomic::{AtomicUsize, Ordering};
#[derive(Default, Debug, PartialEq)]
struct FakeStmt {
timeout_ms: Option<u64>,
consistency: Option<String>,
page_size: Option<i32>,
}
struct TimeoutMod {
ms: u64,
}
impl OpFieldModifier<FakeStmt> for TimeoutMod {
fn field_name(&self) -> &'static str {
"request_timeout_ms"
}
fn apply(&self, t: &mut FakeStmt) {
t.timeout_ms = Some(self.ms);
}
fn diagnostic_value(&self) -> serde_json::Value {
serde_json::Value::from(self.ms)
}
}
struct ConsistencyMod {
cl: String,
}
impl OpFieldModifier<FakeStmt> for ConsistencyMod {
fn field_name(&self) -> &'static str {
"consistency"
}
fn apply(&self, t: &mut FakeStmt) {
t.consistency = Some(self.cl.clone());
}
fn diagnostic_value(&self) -> serde_json::Value {
serde_json::Value::String(self.cl.clone())
}
}
#[test]
fn empty_chain_is_noop() {
let chain: ModifierChain<FakeStmt> = ModifierChain::new("op1", vec![], None);
assert!(chain.is_empty());
let mut stmt = FakeStmt::default();
chain.apply(&mut stmt);
assert_eq!(stmt, FakeStmt::default());
}
#[test]
fn single_modifier_applies_captured_value() {
let chain: ModifierChain<FakeStmt> = ModifierChain::new(
"op_drop_index",
vec![Box::new(TimeoutMod { ms: 300_000 })],
None,
);
assert_eq!(chain.len(), 1);
let mut stmt = FakeStmt::default();
chain.apply(&mut stmt);
assert_eq!(stmt.timeout_ms, Some(300_000));
assert_eq!(stmt.consistency, None);
assert_eq!(stmt.page_size, None);
}
#[test]
fn multiple_modifiers_apply_in_order() {
let chain: ModifierChain<FakeStmt> = ModifierChain::new(
"op_select",
vec![
Box::new(TimeoutMod { ms: 5_000 }),
Box::new(ConsistencyMod {
cl: "LOCAL_QUORUM".to_string(),
}),
],
None,
);
let mut stmt = FakeStmt::default();
chain.apply(&mut stmt);
assert_eq!(stmt.timeout_ms, Some(5_000));
assert_eq!(stmt.consistency, Some("LOCAL_QUORUM".to_string()));
}
struct RecordingSink {
records: Mutex<Vec<(String, &'static str, serde_json::Value)>>,
invocation_count: AtomicUsize,
}
impl ModifierTraceSink for RecordingSink {
fn modifier_applied(
&self,
op: &str,
field: &'static str,
value_fn: &dyn Fn() -> serde_json::Value,
) {
self.invocation_count.fetch_add(1, Ordering::Relaxed);
let value = value_fn(); self.records
.lock()
.unwrap()
.push((op.to_string(), field, value));
}
}
struct GatedSink {
skipped: AtomicUsize,
}
impl ModifierTraceSink for GatedSink {
fn modifier_applied(
&self,
_op: &str,
_field: &'static str,
_value_fn: &dyn Fn() -> serde_json::Value,
) {
self.skipped.fetch_add(1, Ordering::Relaxed);
}
}
struct CountingMod {
diag_calls: Arc<AtomicUsize>,
}
impl OpFieldModifier<FakeStmt> for CountingMod {
fn field_name(&self) -> &'static str {
"request_timeout_ms"
}
fn apply(&self, t: &mut FakeStmt) {
t.timeout_ms = Some(42);
}
fn diagnostic_value(&self) -> serde_json::Value {
self.diag_calls.fetch_add(1, Ordering::Relaxed);
serde_json::Value::from(42u64)
}
}
#[test]
fn recording_sink_sees_all_fired_modifiers() {
let sink = Arc::new(RecordingSink {
records: Mutex::new(Vec::new()),
invocation_count: AtomicUsize::new(0),
});
let chain: ModifierChain<FakeStmt> = ModifierChain::new(
"op_drop_index",
vec![
Box::new(TimeoutMod { ms: 300_000 }),
Box::new(ConsistencyMod {
cl: "ONE".to_string(),
}),
],
Some(sink.clone()),
);
let mut stmt = FakeStmt::default();
chain.apply(&mut stmt);
assert_eq!(sink.invocation_count.load(Ordering::Relaxed), 2);
let records = sink.records.lock().unwrap();
assert_eq!(records.len(), 2);
assert_eq!(records[0].0, "op_drop_index");
assert_eq!(records[0].1, "request_timeout_ms");
assert_eq!(records[0].2, serde_json::json!(300_000));
assert_eq!(records[1].1, "consistency");
assert_eq!(records[1].2, serde_json::json!("ONE"));
}
#[test]
fn gated_sink_does_not_invoke_diagnostic_closure() {
let diag_calls = Arc::new(AtomicUsize::new(0));
let sink = Arc::new(GatedSink {
skipped: AtomicUsize::new(0),
});
let chain: ModifierChain<FakeStmt> = ModifierChain::new(
"op_select",
vec![Box::new(CountingMod {
diag_calls: diag_calls.clone(),
})],
Some(sink.clone()),
);
let mut stmt = FakeStmt::default();
for _ in 0..1000 {
chain.apply(&mut stmt);
}
assert_eq!(sink.skipped.load(Ordering::Relaxed), 1000);
assert_eq!(diag_calls.load(Ordering::Relaxed), 0);
}
#[test]
fn no_sink_hot_path_does_not_invoke_diagnostic_closure() {
let diag_calls = Arc::new(AtomicUsize::new(0));
let chain: ModifierChain<FakeStmt> = ModifierChain::new(
"op_select",
vec![Box::new(CountingMod {
diag_calls: diag_calls.clone(),
})],
None, );
let mut stmt = FakeStmt::default();
for _ in 0..1000 {
chain.apply(&mut stmt);
}
assert_eq!(diag_calls.load(Ordering::Relaxed), 0);
assert_eq!(stmt.timeout_ms, Some(42));
}
#[test]
fn diagnostic_value_returns_json() {
let m = TimeoutMod { ms: 300_000 };
assert_eq!(m.diagnostic_value(), serde_json::json!(300_000));
let m = ConsistencyMod {
cl: "LOCAL_QUORUM".to_string(),
};
assert_eq!(m.diagnostic_value(), serde_json::json!("LOCAL_QUORUM"));
}
#[test]
fn debug_impl_reports_active_count_without_revealing_state() {
let chain: ModifierChain<FakeStmt> =
ModifierChain::new("op1", vec![Box::new(TimeoutMod { ms: 100 })], None);
let s = format!("{:?}", chain);
assert!(s.contains("op_label"));
assert!(s.contains("active_count"));
assert!(s.contains("event_sink"));
}
}