use std::collections::{HashMap, HashSet};
use serde::{Deserialize, Serialize};
use crate::causal::{CausalNode, CausalStage, CausalStore, PredictedEffect};
use crate::state::NodeId;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub enum CounterfactualType {
WhatIf,
WhyNot,
WhatIfInstead,
RegretAnalysis,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum Intervention {
SetAction {
original: CausalNode,
replacement: CausalNode,
},
RemoveEvent {
node: CausalNode,
},
ForceActivation {
node: CausalNode,
strength: f64,
},
ChangeTime {
node: CausalNode,
time_delta_secs: f64,
},
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Observation {
pub actual_action: CausalNode,
pub actual_outcomes: Vec<(CausalNode, f64)>,
pub actual_utility: f64,
pub timestamp_ms: u64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CounterfactualQuery {
pub intervention: Intervention,
pub observation: Option<Observation>,
pub horizon_steps: usize,
pub query_type: CounterfactualType,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SimulatedStep {
pub step: usize,
pub trigger: CausalNode,
pub effect: CausalNode,
pub propagated_strength: f64,
pub confidence: f64,
pub edge_strength: f64,
pub edge_confidence: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct StateSnapshot {
pub node_activations: Vec<(CausalNode, f64)>,
pub steps_simulated: usize,
pub truncated: bool,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OutcomeDifference {
pub actual_utility: f64,
pub counterfactual_utility: f64,
pub changed_nodes: Vec<NodeDelta>,
pub narrative_impact: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NodeDelta {
pub node: CausalNode,
pub actual: f64,
pub counterfactual: f64,
pub direction: DeltaDirection,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum DeltaDirection {
Improved,
Worsened,
Neutral,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CounterfactualResult {
pub query: CounterfactualQuery,
pub trajectory: Vec<SimulatedStep>,
pub final_state: StateSnapshot,
pub divergence_point: usize,
pub outcome_difference: Option<OutcomeDifference>,
pub confidence: f64,
pub regret_score: f64,
pub explanation: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RegretReport {
pub top_regrets: Vec<CounterfactualResult>,
pub pattern: Option<String>,
pub actionable_insight: Option<String>,
pub decisions_analyzed: usize,
pub regret_rate: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DecisionRecord {
pub action_taken: CausalNode,
pub alternatives: Vec<CausalNode>,
pub observation: Observation,
pub context_tags: Vec<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SensitivityEntry {
pub factor: String,
pub sensitivity: f64,
pub direction: DeltaDirection,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CounterfactualConfig {
pub max_horizon: usize,
pub min_confidence: f64,
pub min_strength: f64,
pub hop_decay: f64,
pub confidence_decay: f64,
pub min_edge_stage: CausalStage,
pub include_hypothesized: bool,
}
impl Default for CounterfactualConfig {
fn default() -> Self {
Self {
max_horizon: 5,
min_confidence: 0.05,
min_strength: 0.03,
hop_decay: 0.80,
confidence_decay: 0.85,
min_edge_stage: CausalStage::Candidate,
include_hypothesized: false,
}
}
}
fn stage_meets_minimum(stage: CausalStage, minimum: CausalStage) -> bool {
let rank = |s: CausalStage| -> u8 {
match s {
CausalStage::Refuted => 0,
CausalStage::Weakening => 1,
CausalStage::Hypothesized => 2,
CausalStage::Candidate => 3,
CausalStage::Established => 4,
}
};
rank(stage) >= rank(minimum)
}
pub fn simulate_counterfactual(
query: &CounterfactualQuery,
store: &CausalStore,
config: &CounterfactualConfig,
) -> CounterfactualResult {
let horizon = query.horizon_steps.min(config.max_horizon);
let (intervention_node, initial_strength) = match &query.intervention {
Intervention::SetAction { replacement, .. } => (replacement.clone(), 1.0),
Intervention::RemoveEvent { node } => (node.clone(), -1.0), Intervention::ForceActivation { node, strength } => (node.clone(), *strength),
Intervention::ChangeTime { node, time_delta_secs } => {
let attenuation = 1.0 / (1.0 + time_delta_secs.abs() / 3600.0);
(node.clone(), attenuation)
}
};
let severed = match &query.intervention {
Intervention::RemoveEvent { node } => {
let mut s = HashSet::new();
s.insert(node.clone());
s
}
Intervention::SetAction { original, .. } => {
let mut s = HashSet::new();
s.insert(original.clone());
s
}
_ => HashSet::new(),
};
let mut trajectory = Vec::new();
let mut activations: HashMap<CausalNode, f64> = HashMap::new();
let mut visited: HashSet<CausalNode> = HashSet::new();
let mut truncated = false;
let mut frontier: Vec<(CausalNode, f64, f64, usize)> =
vec![(intervention_node.clone(), initial_strength, 1.0, 0)];
activations.insert(intervention_node.clone(), initial_strength);
while let Some((current, strength, confidence, step)) = frontier.pop() {
if step >= horizon {
continue;
}
if !visited.insert(current.clone()) {
continue; }
let effects = store.effects_of(¤t);
for edge in &effects {
if severed.contains(&edge.effect) {
continue;
}
if edge.stage == CausalStage::Refuted {
continue;
}
if !config.include_hypothesized
&& !stage_meets_minimum(edge.stage, config.min_edge_stage)
{
continue;
}
let prop_strength = strength * edge.strength * config.hop_decay;
let prop_confidence = confidence * edge.confidence * config.confidence_decay;
if prop_confidence < config.min_confidence
|| prop_strength.abs() < config.min_strength
{
if step == 0 {
truncated = true;
}
continue;
}
trajectory.push(SimulatedStep {
step,
trigger: current.clone(),
effect: edge.effect.clone(),
propagated_strength: prop_strength,
confidence: prop_confidence,
edge_strength: edge.strength,
edge_confidence: edge.confidence,
});
let entry = activations.entry(edge.effect.clone()).or_insert(0.0);
*entry += prop_strength;
frontier.push((
edge.effect.clone(),
prop_strength,
prop_confidence,
step + 1,
));
}
}
trajectory.sort_by_key(|s| s.step);
let node_activations: Vec<(CausalNode, f64)> = activations.into_iter().collect();
let steps_simulated = trajectory.last().map(|s| s.step + 1).unwrap_or(0);
let final_state = StateSnapshot {
node_activations: node_activations.clone(),
steps_simulated,
truncated,
};
let confidence = if trajectory.is_empty() {
0.0
} else {
let product: f64 = trajectory.iter().map(|s| s.confidence).product();
product.powf(1.0 / trajectory.len() as f64)
};
let (outcome_difference, regret_score) = if let Some(obs) = &query.observation {
let diff = compare_outcomes(obs, &final_state);
let regret = diff.counterfactual_utility - diff.actual_utility;
(Some(diff), regret)
} else {
(None, 0.0)
};
let explanation = explain_simulation(&query.intervention, &trajectory, &final_state);
CounterfactualResult {
query: query.clone(),
trajectory,
final_state,
divergence_point: 0, outcome_difference,
confidence,
regret_score,
explanation,
}
}
pub fn compare_outcomes(
observation: &Observation,
counterfactual: &StateSnapshot,
) -> OutcomeDifference {
let actual_map: HashMap<CausalNode, f64> = observation
.actual_outcomes
.iter()
.cloned()
.collect();
let cf_map: HashMap<CausalNode, f64> = counterfactual
.node_activations
.iter()
.cloned()
.collect();
let mut all_nodes: HashSet<CausalNode> = HashSet::new();
for (n, _) in &observation.actual_outcomes {
all_nodes.insert(n.clone());
}
for (n, _) in &counterfactual.node_activations {
all_nodes.insert(n.clone());
}
let mut changed_nodes = Vec::new();
for node in all_nodes {
let actual = actual_map.get(&node).copied().unwrap_or(0.0);
let cf = cf_map.get(&node).copied().unwrap_or(0.0);
let delta = cf - actual;
if delta.abs() > 0.05 {
let direction = if delta > 0.05 {
DeltaDirection::Improved
} else if delta < -0.05 {
DeltaDirection::Worsened
} else {
DeltaDirection::Neutral
};
changed_nodes.push(NodeDelta {
node,
actual,
counterfactual: cf,
direction,
});
}
}
changed_nodes.sort_by(|a, b| {
let mag_a = (a.counterfactual - a.actual).abs();
let mag_b = (b.counterfactual - b.actual).abs();
mag_b
.partial_cmp(&mag_a)
.unwrap_or(std::cmp::Ordering::Equal)
});
let cf_utility = estimate_utility_from_activations(&counterfactual.node_activations);
let narrative_impact = generate_comparison_narrative(
observation.actual_utility,
cf_utility,
&changed_nodes,
);
OutcomeDifference {
actual_utility: observation.actual_utility,
counterfactual_utility: cf_utility,
changed_nodes,
narrative_impact,
}
}
fn estimate_utility_from_activations(activations: &[(CausalNode, f64)]) -> f64 {
if activations.is_empty() {
return 0.0;
}
let sum: f64 = activations.iter().map(|(_, v)| *v).sum();
(sum / activations.len() as f64).clamp(-1.0, 1.0)
}
fn generate_comparison_narrative(
actual_utility: f64,
cf_utility: f64,
changed: &[NodeDelta],
) -> String {
let delta = cf_utility - actual_utility;
let mut parts = Vec::new();
if delta > 0.1 {
parts.push(format!(
"The alternative would likely have been better (utility {:.2} vs {:.2}).",
cf_utility, actual_utility
));
} else if delta < -0.1 {
parts.push(format!(
"The actual choice was likely better (utility {:.2} vs counterfactual {:.2}).",
actual_utility, cf_utility
));
} else {
parts.push("The alternative would likely have produced similar results.".to_string());
}
let improved: Vec<&NodeDelta> = changed
.iter()
.filter(|d| d.direction == DeltaDirection::Improved)
.collect();
let worsened: Vec<&NodeDelta> = changed
.iter()
.filter(|d| d.direction == DeltaDirection::Worsened)
.collect();
if !improved.is_empty() {
parts.push(format!(
"{} factor(s) would have improved.",
improved.len()
));
}
if !worsened.is_empty() {
parts.push(format!(
"{} factor(s) would have worsened.",
worsened.len()
));
}
parts.join(" ")
}
fn explain_simulation(
intervention: &Intervention,
trajectory: &[SimulatedStep],
final_state: &StateSnapshot,
) -> String {
let mut parts = Vec::new();
match intervention {
Intervention::SetAction { original, replacement } => {
parts.push(format!(
"If {:?} were replaced with {:?}:",
original, replacement
));
}
Intervention::RemoveEvent { node } => {
parts.push(format!("If {:?} had not occurred:", node));
}
Intervention::ForceActivation { node, strength } => {
parts.push(format!(
"If {:?} were forced to activation {:.2}:",
node, strength
));
}
Intervention::ChangeTime { node, time_delta_secs } => {
let direction = if *time_delta_secs > 0.0 {
"later"
} else {
"earlier"
};
parts.push(format!(
"If {:?} happened {:.0}s {}:",
node,
time_delta_secs.abs(),
direction
));
}
}
if trajectory.is_empty() {
parts.push("No downstream effects predicted.".to_string());
} else {
parts.push(format!(
"{} causal step(s) simulated across {} effect(s).",
final_state.steps_simulated,
trajectory.len()
));
let mut sorted = trajectory.to_vec();
sorted.sort_by(|a, b| {
b.propagated_strength
.abs()
.partial_cmp(&a.propagated_strength.abs())
.unwrap_or(std::cmp::Ordering::Equal)
});
for step in sorted.iter().take(3) {
let verb = if step.propagated_strength > 0.0 {
"promotes"
} else {
"inhibits"
};
parts.push(format!(
" {:?} {} {:?} (strength {:.2}, confidence {:.2})",
step.trigger, verb, step.effect, step.propagated_strength, step.confidence
));
}
}
if final_state.truncated {
parts.push("(Simulation truncated due to low confidence.)".to_string());
}
parts.join("\n")
}
pub fn why_not(
desired_outcome: &CausalNode,
actual_actions: &[CausalNode],
store: &CausalStore,
config: &CounterfactualConfig,
) -> Vec<CounterfactualResult> {
let mut results = Vec::new();
let potential_causes = store.causes_of(desired_outcome);
for edge in &potential_causes {
if edge.stage == CausalStage::Refuted {
continue;
}
if !config.include_hypothesized
&& !stage_meets_minimum(edge.stage, config.min_edge_stage)
{
continue;
}
let was_tried = actual_actions.contains(&edge.cause);
if was_tried {
continue; }
let query = CounterfactualQuery {
intervention: Intervention::ForceActivation {
node: edge.cause.clone(),
strength: 1.0,
},
observation: None,
horizon_steps: config.max_horizon.min(3),
query_type: CounterfactualType::WhyNot,
};
let result = simulate_counterfactual(&query, store, config);
if !result.trajectory.is_empty() {
results.push(result);
}
}
results.sort_by(|a, b| {
let score_a = a.confidence * a.trajectory.len() as f64;
let score_b = b.confidence * b.trajectory.len() as f64;
score_b
.partial_cmp(&score_a)
.unwrap_or(std::cmp::Ordering::Equal)
});
results
}
pub fn detect_regret_opportunities(
decisions: &[DecisionRecord],
store: &CausalStore,
config: &CounterfactualConfig,
) -> RegretReport {
let mut all_results = Vec::new();
let mut regret_count = 0usize;
let mut regret_tags: HashMap<String, usize> = HashMap::new();
for decision in decisions {
for alt in &decision.alternatives {
let query = CounterfactualQuery {
intervention: Intervention::SetAction {
original: decision.action_taken.clone(),
replacement: alt.clone(),
},
observation: Some(decision.observation.clone()),
horizon_steps: config.max_horizon.min(3),
query_type: CounterfactualType::RegretAnalysis,
};
let result = simulate_counterfactual(&query, store, config);
if result.regret_score < -0.1 {
regret_count += 1;
for tag in &decision.context_tags {
*regret_tags.entry(tag.clone()).or_insert(0) += 1;
}
}
all_results.push(result);
}
}
all_results.sort_by(|a, b| {
b.regret_score
.abs()
.partial_cmp(&a.regret_score.abs())
.unwrap_or(std::cmp::Ordering::Equal)
});
let pattern = detect_regret_pattern(®ret_tags, decisions.len());
let insight = generate_actionable_insight(®ret_tags, &pattern);
let total = decisions.len().max(1);
let regret_rate = regret_count as f64 / total as f64;
RegretReport {
top_regrets: all_results.into_iter().take(10).collect(),
pattern,
actionable_insight: insight,
decisions_analyzed: decisions.len(),
regret_rate,
}
}
fn detect_regret_pattern(
tag_counts: &HashMap<String, usize>,
total_decisions: usize,
) -> Option<String> {
if tag_counts.is_empty() || total_decisions == 0 {
return None;
}
let (top_tag, count) = tag_counts
.iter()
.max_by_key(|(_, c)| **c)?;
let rate = *count as f64 / total_decisions as f64;
if rate > 0.3 {
Some(format!(
"Regretted decisions often involve '{}' ({:.0}% of cases).",
top_tag,
rate * 100.0
))
} else {
None
}
}
fn generate_actionable_insight(
tag_counts: &HashMap<String, usize>,
pattern: &Option<String>,
) -> Option<String> {
if pattern.is_none() {
return None;
}
let (top_tag, _) = tag_counts.iter().max_by_key(|(_, c)| **c)?;
Some(format!(
"Consider giving extra attention to decisions involving '{}'. \
Historical data suggests alternatives in this area tend to produce better outcomes.",
top_tag
))
}
pub fn sensitivity_analysis(
base_query: &CounterfactualQuery,
store: &CausalStore,
config: &CounterfactualConfig,
) -> Vec<SensitivityEntry> {
let mut entries = Vec::new();
let intervention_node = match &base_query.intervention {
Intervention::SetAction { replacement, .. } => replacement.clone(),
Intervention::RemoveEvent { node } => node.clone(),
Intervention::ForceActivation { node, .. } => node.clone(),
Intervention::ChangeTime { node, .. } => node.clone(),
};
let baseline = simulate_counterfactual(base_query, store, config);
let baseline_utility = baseline
.outcome_difference
.as_ref()
.map(|d| d.counterfactual_utility)
.unwrap_or_else(|| {
estimate_utility_from_activations(&baseline.final_state.node_activations)
});
let edges = store.effects_of(&intervention_node);
for edge in &edges {
if edge.stage == CausalStage::Refuted {
continue;
}
let modified_query = CounterfactualQuery {
intervention: Intervention::RemoveEvent {
node: edge.effect.clone(),
},
observation: base_query.observation.clone(),
horizon_steps: base_query.horizon_steps,
query_type: CounterfactualType::WhatIf,
};
let modified = simulate_counterfactual(&modified_query, store, config);
let modified_utility = modified
.outcome_difference
.as_ref()
.map(|d| d.counterfactual_utility)
.unwrap_or_else(|| {
estimate_utility_from_activations(&modified.final_state.node_activations)
});
let sensitivity = (baseline_utility - modified_utility).abs();
let direction = if modified_utility < baseline_utility {
DeltaDirection::Improved } else if modified_utility > baseline_utility {
DeltaDirection::Worsened } else {
DeltaDirection::Neutral
};
entries.push(SensitivityEntry {
factor: format!("{:?}", edge.effect),
sensitivity,
direction,
});
}
entries.sort_by(|a, b| {
b.sensitivity
.partial_cmp(&a.sensitivity)
.unwrap_or(std::cmp::Ordering::Equal)
});
entries
}
pub fn strengthen_causal_model(
result: &CounterfactualResult,
actual_observation: &Observation,
store: &mut CausalStore,
) -> usize {
let actual_map: HashMap<CausalNode, f64> = actual_observation
.actual_outcomes
.iter()
.cloned()
.collect();
let mut updates = 0;
for step in &result.trajectory {
let predicted = step.propagated_strength;
let actual = actual_map.get(&step.effect).copied().unwrap_or(0.0);
let error = (predicted - actual).abs();
if let Some(edge) = store.find_edge(&step.trigger, &step.effect) {
let mut updated = edge.clone();
if error < 0.2 {
updated.confidence = (updated.confidence + 0.05).min(1.0);
} else if error > 0.5 {
updated.confidence = (updated.confidence - 0.1).max(0.0);
}
updated.updated_at = actual_observation.timestamp_ms as f64 / 1000.0;
store.upsert(updated);
updates += 1;
}
}
updates
}
pub fn net_impact(effects: &[PredictedEffect]) -> f64 {
effects
.iter()
.map(|e| e.expected_strength * e.confidence)
.sum()
}
pub fn compare_alternatives(
action_a: &CausalNode,
action_b: &CausalNode,
store: &CausalStore,
config: &CounterfactualConfig,
) -> f64 {
let query_a = CounterfactualQuery {
intervention: Intervention::ForceActivation {
node: action_a.clone(),
strength: 1.0,
},
observation: None,
horizon_steps: config.max_horizon.min(3),
query_type: CounterfactualType::WhatIfInstead,
};
let query_b = CounterfactualQuery {
intervention: Intervention::ForceActivation {
node: action_b.clone(),
strength: 1.0,
},
observation: None,
horizon_steps: config.max_horizon.min(3),
query_type: CounterfactualType::WhatIfInstead,
};
let result_a = simulate_counterfactual(&query_a, store, config);
let result_b = simulate_counterfactual(&query_b, store, config);
let utility_a =
estimate_utility_from_activations(&result_a.final_state.node_activations);
let utility_b =
estimate_utility_from_activations(&result_b.final_state.node_activations);
utility_a - utility_b
}
#[cfg(test)]
mod tests {
use super::*;
use crate::causal::{CausalEdge, CausalEvidence, CausalNode, CausalStage, CausalStore, CausalTrace, DiscoveryMethod};
use crate::observer::EventKind;
use crate::state::NodeId;
use crate::world_model::ActionKind;
fn make_edge(
cause: CausalNode,
effect: CausalNode,
strength: f64,
confidence: f64,
stage: CausalStage,
) -> CausalEdge {
CausalEdge {
cause,
effect,
strength,
confidence,
observation_count: 10,
intervention_count: 0,
non_occurrence_count: 2,
median_lag_secs: 60.0,
lag_iqr_secs: 30.0,
context_strengths: vec![],
trace: CausalTrace {
evidence: vec![CausalEvidence::TemporalPrecedence {
co_occurrences: 10,
avg_lag_secs: 60.0,
lag_stddev_secs: 15.0,
}],
primary_method: DiscoveryMethod::TemporalAssociation,
summary: "test edge".to_string(),
},
created_at: 1000.0,
updated_at: 2000.0,
stage,
}
}
fn node(id: u32) -> CausalNode {
CausalNode::GraphNode(NodeId::from_raw(id))
}
fn action_node(kind: ActionKind) -> CausalNode {
CausalNode::Action(kind)
}
fn event_node(kind: EventKind) -> CausalNode {
CausalNode::Event(kind)
}
fn signal(name: &str) -> CausalNode {
CausalNode::Signal(name.to_string())
}
fn make_store_with_chain() -> CausalStore {
let mut store = CausalStore::new();
store.upsert(make_edge(
signal("A"),
signal("B"),
0.8,
0.9,
CausalStage::Established,
));
store.upsert(make_edge(
signal("B"),
signal("C"),
0.7,
0.85,
CausalStage::Established,
));
store.upsert(make_edge(
signal("C"),
signal("D"),
0.6,
0.8,
CausalStage::Candidate,
));
store
}
fn default_config() -> CounterfactualConfig {
CounterfactualConfig::default()
}
#[test]
fn test_single_step_simulation() {
let mut store = CausalStore::new();
store.upsert(make_edge(
signal("cause"),
signal("effect"),
0.9,
0.95,
CausalStage::Established,
));
let query = CounterfactualQuery {
intervention: Intervention::ForceActivation {
node: signal("cause"),
strength: 1.0,
},
observation: None,
horizon_steps: 1,
query_type: CounterfactualType::WhatIf,
};
let result = simulate_counterfactual(&query, &store, &default_config());
assert_eq!(result.trajectory.len(), 1);
assert_eq!(result.trajectory[0].effect, signal("effect"));
assert!(result.trajectory[0].propagated_strength > 0.0);
assert!(result.confidence > 0.0);
}
#[test]
fn test_multi_step_chain() {
let store = make_store_with_chain();
let config = default_config();
let query = CounterfactualQuery {
intervention: Intervention::ForceActivation {
node: signal("A"),
strength: 1.0,
},
observation: None,
horizon_steps: 5,
query_type: CounterfactualType::WhatIf,
};
let result = simulate_counterfactual(&query, &store, &config);
assert!(result.trajectory.len() >= 2);
assert!(result.final_state.steps_simulated >= 2);
let first_conf = result.trajectory[0].confidence;
let last_conf = result.trajectory.last().unwrap().confidence;
assert!(last_conf < first_conf);
}
#[test]
fn test_remove_event_intervention() {
let mut store = CausalStore::new();
store.upsert(make_edge(
signal("A"),
signal("B"),
0.8,
0.9,
CausalStage::Established,
));
store.upsert(make_edge(
signal("A"),
signal("C"),
0.7,
0.85,
CausalStage::Established,
));
store.upsert(make_edge(
signal("B"),
signal("D"),
0.6,
0.8,
CausalStage::Established,
));
let query = CounterfactualQuery {
intervention: Intervention::RemoveEvent {
node: signal("B"),
},
observation: None,
horizon_steps: 3,
query_type: CounterfactualType::WhatIf,
};
let result = simulate_counterfactual(&query, &store, &default_config());
assert!(result.explanation.contains("had not occurred"));
}
#[test]
fn test_set_action_intervention() {
let mut store = CausalStore::new();
store.upsert(make_edge(
action_node(ActionKind::SurfaceSuggestion),
signal("good_outcome"),
0.8,
0.9,
CausalStage::Established,
));
store.upsert(make_edge(
action_node(ActionKind::SendNotification),
signal("bad_outcome"),
0.7,
0.85,
CausalStage::Established,
));
let query = CounterfactualQuery {
intervention: Intervention::SetAction {
original: action_node(ActionKind::SendNotification),
replacement: action_node(ActionKind::SurfaceSuggestion),
},
observation: None,
horizon_steps: 2,
query_type: CounterfactualType::WhatIfInstead,
};
let result = simulate_counterfactual(&query, &store, &default_config());
assert!(!result.explanation.is_empty());
}
#[test]
fn test_change_time_intervention() {
let store = make_store_with_chain();
let query = CounterfactualQuery {
intervention: Intervention::ChangeTime {
node: signal("A"),
time_delta_secs: 7200.0, },
observation: None,
horizon_steps: 3,
query_type: CounterfactualType::WhatIf,
};
let result = simulate_counterfactual(&query, &store, &default_config());
assert!(result.explanation.contains("later"));
}
#[test]
fn test_outcome_comparison() {
let observation = Observation {
actual_action: signal("action_A"),
actual_outcomes: vec![
(signal("goal_1"), 0.3),
(signal("goal_2"), 0.8),
],
actual_utility: 0.5,
timestamp_ms: 1000000,
};
let counterfactual = StateSnapshot {
node_activations: vec![
(signal("goal_1"), 0.9), (signal("goal_2"), 0.2), ],
steps_simulated: 2,
truncated: false,
};
let diff = compare_outcomes(&observation, &counterfactual);
assert_eq!(diff.actual_utility, 0.5);
assert!(diff.changed_nodes.len() >= 2);
let goal1_delta = diff
.changed_nodes
.iter()
.find(|d| d.node == signal("goal_1"))
.unwrap();
assert_eq!(goal1_delta.direction, DeltaDirection::Improved);
let goal2_delta = diff
.changed_nodes
.iter()
.find(|d| d.node == signal("goal_2"))
.unwrap();
assert_eq!(goal2_delta.direction, DeltaDirection::Worsened);
}
#[test]
fn test_regret_score_positive() {
let mut store = CausalStore::new();
store.upsert(make_edge(
signal("bad_alt"),
signal("bad_outcome"),
-0.5,
0.8,
CausalStage::Established,
));
let query = CounterfactualQuery {
intervention: Intervention::ForceActivation {
node: signal("bad_alt"),
strength: 1.0,
},
observation: Some(Observation {
actual_action: signal("good_action"),
actual_outcomes: vec![(signal("result"), 0.8)],
actual_utility: 0.8,
timestamp_ms: 1000000,
}),
horizon_steps: 2,
query_type: CounterfactualType::RegretAnalysis,
};
let result = simulate_counterfactual(&query, &store, &default_config());
assert!(result.regret_score <= 0.0 || result.trajectory.is_empty());
}
#[test]
fn test_why_not_finds_causes() {
let mut store = CausalStore::new();
store.upsert(make_edge(
signal("X"),
signal("desired"),
0.9,
0.9,
CausalStage::Established,
));
store.upsert(make_edge(
signal("Y"),
signal("desired"),
0.7,
0.8,
CausalStage::Candidate,
));
let actual = vec![signal("Z")];
let results = why_not(
&signal("desired"),
&actual,
&store,
&default_config(),
);
assert_eq!(results.len(), 2);
}
#[test]
fn test_why_not_excludes_tried_actions() {
let mut store = CausalStore::new();
store.upsert(make_edge(
signal("X"),
signal("desired"),
0.9,
0.9,
CausalStage::Established,
));
store.upsert(make_edge(
signal("Y"),
signal("desired"),
0.7,
0.8,
CausalStage::Candidate,
));
let actual = vec![signal("X")];
let results = why_not(
&signal("desired"),
&actual,
&store,
&default_config(),
);
assert_eq!(results.len(), 1);
}
#[test]
fn test_regret_analysis() {
let mut store = CausalStore::new();
store.upsert(make_edge(
signal("alt_good"),
signal("great_result"),
0.9,
0.9,
CausalStage::Established,
));
let decisions = vec![DecisionRecord {
action_taken: signal("mediocre"),
alternatives: vec![signal("alt_good")],
observation: Observation {
actual_action: signal("mediocre"),
actual_outcomes: vec![(signal("ok_result"), 0.3)],
actual_utility: 0.3,
timestamp_ms: 1000000,
},
context_tags: vec!["morning".to_string(), "rushed".to_string()],
}];
let report = detect_regret_opportunities(&decisions, &store, &default_config());
assert_eq!(report.decisions_analyzed, 1);
assert!(!report.top_regrets.is_empty());
}
#[test]
fn test_regret_pattern_detection() {
let mut store = CausalStore::new();
store.upsert(make_edge(
signal("alt"),
signal("better"),
0.9,
0.9,
CausalStage::Established,
));
let decisions: Vec<DecisionRecord> = (0..5)
.map(|i| DecisionRecord {
action_taken: signal("mediocre"),
alternatives: vec![signal("alt")],
observation: Observation {
actual_action: signal("mediocre"),
actual_outcomes: vec![(signal("meh"), 0.2)],
actual_utility: 0.2,
timestamp_ms: 1000000 + i * 1000,
},
context_tags: vec!["rushed".to_string()],
})
.collect();
let report = detect_regret_opportunities(&decisions, &store, &default_config());
if report.regret_rate > 0.3 {
assert!(report.pattern.is_some());
}
}
#[test]
fn test_sensitivity_analysis() {
let mut store = CausalStore::new();
store.upsert(make_edge(
signal("A"),
signal("B"),
0.9,
0.95,
CausalStage::Established,
));
store.upsert(make_edge(
signal("A"),
signal("C"),
0.2,
0.6,
CausalStage::Candidate,
));
let query = CounterfactualQuery {
intervention: Intervention::ForceActivation {
node: signal("A"),
strength: 1.0,
},
observation: None,
horizon_steps: 2,
query_type: CounterfactualType::WhatIf,
};
let entries = sensitivity_analysis(&query, &store, &default_config());
assert!(!entries.is_empty());
}
#[test]
fn test_strengthen_good_prediction() {
let mut store = CausalStore::new();
store.upsert(make_edge(
signal("A"),
signal("B"),
0.8,
0.7, CausalStage::Established,
));
let result = CounterfactualResult {
query: CounterfactualQuery {
intervention: Intervention::ForceActivation {
node: signal("A"),
strength: 1.0,
},
observation: None,
horizon_steps: 1,
query_type: CounterfactualType::WhatIf,
},
trajectory: vec![SimulatedStep {
step: 0,
trigger: signal("A"),
effect: signal("B"),
propagated_strength: 0.64, confidence: 0.6,
edge_strength: 0.8,
edge_confidence: 0.7,
}],
final_state: StateSnapshot {
node_activations: vec![(signal("B"), 0.64)],
steps_simulated: 1,
truncated: false,
},
divergence_point: 0,
outcome_difference: None,
confidence: 0.6,
regret_score: 0.0,
explanation: String::new(),
};
let observation = Observation {
actual_action: signal("A"),
actual_outcomes: vec![(signal("B"), 0.6)], actual_utility: 0.6,
timestamp_ms: 3000000,
};
let updates = strengthen_causal_model(&result, &observation, &mut store);
assert_eq!(updates, 1);
let edge = store.find_edge(&signal("A"), &signal("B")).unwrap();
assert!(edge.confidence > 0.7);
}
#[test]
fn test_weaken_bad_prediction() {
let mut store = CausalStore::new();
store.upsert(make_edge(
signal("A"),
signal("B"),
0.8,
0.7,
CausalStage::Established,
));
let result = CounterfactualResult {
query: CounterfactualQuery {
intervention: Intervention::ForceActivation {
node: signal("A"),
strength: 1.0,
},
observation: None,
horizon_steps: 1,
query_type: CounterfactualType::WhatIf,
},
trajectory: vec![SimulatedStep {
step: 0,
trigger: signal("A"),
effect: signal("B"),
propagated_strength: 0.64,
confidence: 0.6,
edge_strength: 0.8,
edge_confidence: 0.7,
}],
final_state: StateSnapshot {
node_activations: vec![(signal("B"), 0.64)],
steps_simulated: 1,
truncated: false,
},
divergence_point: 0,
outcome_difference: None,
confidence: 0.6,
regret_score: 0.0,
explanation: String::new(),
};
let observation = Observation {
actual_action: signal("A"),
actual_outcomes: vec![(signal("B"), -0.3)], actual_utility: -0.3,
timestamp_ms: 3000000,
};
let updates = strengthen_causal_model(&result, &observation, &mut store);
assert_eq!(updates, 1);
let edge = store.find_edge(&signal("A"), &signal("B")).unwrap();
assert!(edge.confidence < 0.7);
}
#[test]
fn test_compare_alternatives() {
let mut store = CausalStore::new();
store.upsert(make_edge(
signal("good"),
signal("positive"),
0.9,
0.9,
CausalStage::Established,
));
store.upsert(make_edge(
signal("bad"),
signal("negative"),
-0.7,
0.8,
CausalStage::Established,
));
let config = default_config();
let diff = compare_alternatives(
&signal("good"),
&signal("bad"),
&store,
&config,
);
assert!(diff > 0.0);
}
#[test]
fn test_empty_store_simulation() {
let store = CausalStore::new();
let config = default_config();
let query = CounterfactualQuery {
intervention: Intervention::ForceActivation {
node: signal("anything"),
strength: 1.0,
},
observation: None,
horizon_steps: 3,
query_type: CounterfactualType::WhatIf,
};
let result = simulate_counterfactual(&query, &store, &config);
assert!(result.trajectory.is_empty());
assert_eq!(result.confidence, 0.0);
}
#[test]
fn test_cycle_prevention() {
let mut store = CausalStore::new();
store.upsert(make_edge(
signal("A"),
signal("B"),
0.8,
0.9,
CausalStage::Established,
));
store.upsert(make_edge(
signal("B"),
signal("A"),
0.7,
0.85,
CausalStage::Established,
));
let query = CounterfactualQuery {
intervention: Intervention::ForceActivation {
node: signal("A"),
strength: 1.0,
},
observation: None,
horizon_steps: 10,
query_type: CounterfactualType::WhatIf,
};
let result = simulate_counterfactual(&query, &store, &default_config());
assert!(result.trajectory.len() <= 4);
}
#[test]
fn test_refuted_edges_skipped() {
let mut store = CausalStore::new();
store.upsert(make_edge(
signal("A"),
signal("B"),
0.8,
0.9,
CausalStage::Refuted, ));
store.upsert(make_edge(
signal("A"),
signal("C"),
0.7,
0.85,
CausalStage::Established,
));
let query = CounterfactualQuery {
intervention: Intervention::ForceActivation {
node: signal("A"),
strength: 1.0,
},
observation: None,
horizon_steps: 2,
query_type: CounterfactualType::WhatIf,
};
let result = simulate_counterfactual(&query, &store, &default_config());
assert!(result.trajectory.len() == 1);
assert_eq!(result.trajectory[0].effect, signal("C"));
}
#[test]
fn test_explanation_generation() {
let store = make_store_with_chain();
let config = default_config();
let query = CounterfactualQuery {
intervention: Intervention::ForceActivation {
node: signal("A"),
strength: 1.0,
},
observation: None,
horizon_steps: 5,
query_type: CounterfactualType::WhatIf,
};
let result = simulate_counterfactual(&query, &store, &config);
assert!(!result.explanation.is_empty());
assert!(result.explanation.contains("forced to activation"));
}
}