use std::collections::HashMap;
use serde::{Deserialize, Serialize};
use crate::state::NodeId;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum SamplingStrategy {
Uniform,
PrioritizedByTDError,
PrioritizedByRecency,
PrioritizedBySurprise,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ActionRecord {
pub description: String,
pub domain: String,
pub involved_nodes: Vec<NodeId>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OutcomeData {
pub utility: f64,
pub expected: bool,
pub domains: Vec<String>,
pub affected_nodes: Vec<NodeId>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ReplayEntry {
pub episode_id: NodeId,
pub expected_utility: f64,
pub action: ActionRecord,
pub outcome: OutcomeData,
pub td_error: f64,
pub replay_count: u32,
pub last_replayed_at: u64,
pub created_at: u64,
pub priority: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BeliefDelta {
pub belief_id: NodeId,
pub confidence_delta: f64,
pub reason: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CausalDelta {
pub cause: NodeId,
pub effect: NodeId,
pub strength_delta: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ReplayOutcome {
pub episode_id: NodeId,
pub new_td_error: f64,
pub belief_updates: Vec<BeliefDelta>,
pub causal_updates: Vec<CausalDelta>,
pub associations: Vec<(NodeId, NodeId, f64)>,
pub insights: Vec<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DreamReport {
pub replays_executed: usize,
pub beliefs_updated: usize,
pub causal_updates: usize,
pub new_associations: usize,
pub insights: Vec<String>,
pub duration_ms: u64,
pub avg_td_error: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ReplayBudget {
pub max_replays_per_cycle: usize,
pub min_idle_ms: u64,
pub max_replay_age_ms: u64,
pub priority_exponent: f64,
pub min_td_error: f64,
pub max_replays_per_episode: u32,
}
impl Default for ReplayBudget {
fn default() -> Self {
Self {
max_replays_per_cycle: 10,
min_idle_ms: 30_000,
max_replay_age_ms: 30 * 24 * 3600 * 1000, priority_exponent: 0.6,
min_td_error: 0.05,
max_replays_per_episode: 5,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ReplayBuffer {
pub entries: Vec<ReplayEntry>,
pub capacity: usize,
pub strategy: SamplingStrategy,
}
impl ReplayBuffer {
pub fn new(capacity: usize, strategy: SamplingStrategy) -> Self {
Self {
entries: Vec::new(),
capacity,
strategy,
}
}
pub fn len(&self) -> usize {
self.entries.len()
}
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct ReplayStats {
pub total_cycles: u64,
pub total_replays: u64,
pub total_belief_updates: u64,
pub total_causal_updates: u64,
pub total_associations: u64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ReplayEngine {
pub buffer: ReplayBuffer,
pub budget: ReplayBudget,
pub last_replay_at: u64,
pub stats: ReplayStats,
}
impl ReplayEngine {
pub fn new() -> Self {
Self {
buffer: ReplayBuffer::new(500, SamplingStrategy::PrioritizedByTDError),
budget: ReplayBudget::default(),
last_replay_at: 0,
stats: ReplayStats::default(),
}
}
pub fn with_config(capacity: usize, strategy: SamplingStrategy, budget: ReplayBudget) -> Self {
Self {
buffer: ReplayBuffer::new(capacity, strategy),
budget,
last_replay_at: 0,
stats: ReplayStats::default(),
}
}
}
pub fn compute_td_error(expected_utility: f64, actual_utility: f64) -> f64 {
(actual_utility - expected_utility).abs()
}
pub fn add_to_buffer(
engine: &mut ReplayEngine,
episode_id: NodeId,
expected_utility: f64,
action: ActionRecord,
outcome: OutcomeData,
now_ms: u64,
) {
let td_error = compute_td_error(expected_utility, outcome.utility);
if td_error < engine.budget.min_td_error {
return;
}
let priority = compute_priority(
td_error,
now_ms,
now_ms, 0, engine.budget.priority_exponent,
);
let entry = ReplayEntry {
episode_id,
expected_utility,
action,
outcome,
td_error,
replay_count: 0,
last_replayed_at: 0,
created_at: now_ms,
priority,
};
if engine.buffer.entries.len() >= engine.buffer.capacity {
if let Some(min_idx) = engine
.buffer
.entries
.iter()
.enumerate()
.min_by(|(_, a), (_, b)| a.priority.total_cmp(&b.priority))
.map(|(i, _)| i)
{
if engine.buffer.entries[min_idx].priority < priority {
engine.buffer.entries.swap_remove(min_idx);
} else {
return; }
}
}
engine.buffer.entries.push(entry);
}
fn compute_priority(
td_error: f64,
now_ms: u64,
created_at: u64,
replay_count: u32,
exponent: f64,
) -> f64 {
let td_component = td_error.abs().powf(exponent);
let age_hours = (now_ms.saturating_sub(created_at)) as f64 / 3_600_000.0;
let recency = 1.0 / (1.0 + age_hours / 24.0);
let novelty = 1.0 / (1.0 + replay_count as f64);
td_component * recency * novelty
}
pub fn should_replay(engine: &ReplayEngine, now_ms: u64) -> bool {
if engine.buffer.is_empty() {
return false;
}
let time_since_last = now_ms.saturating_sub(engine.last_replay_at);
time_since_last >= engine.budget.min_idle_ms
}
pub fn reprioritize_buffer(engine: &mut ReplayEngine, now_ms: u64) {
for entry in &mut engine.buffer.entries {
entry.priority = compute_priority(
entry.td_error,
now_ms,
entry.created_at,
entry.replay_count,
engine.budget.priority_exponent,
);
}
engine
.buffer
.entries
.sort_by(|a, b| b.priority.total_cmp(&a.priority));
}
fn sample_episodes(engine: &ReplayEngine, count: usize) -> Vec<usize> {
let n = engine.buffer.entries.len().min(count);
let budget = &engine.budget;
match engine.buffer.strategy {
SamplingStrategy::Uniform => {
(0..engine.buffer.entries.len())
.filter(|&i| {
let e = &engine.buffer.entries[i];
e.replay_count < budget.max_replays_per_episode
})
.take(n)
.collect()
}
SamplingStrategy::PrioritizedByTDError | SamplingStrategy::PrioritizedBySurprise => {
let mut indices = Vec::new();
for (i, entry) in engine.buffer.entries.iter().enumerate() {
if entry.replay_count < budget.max_replays_per_episode {
indices.push(i);
if indices.len() >= n {
break;
}
}
}
indices
}
SamplingStrategy::PrioritizedByRecency => {
let mut sorted_indices: Vec<usize> = (0..engine.buffer.entries.len())
.filter(|&i| engine.buffer.entries[i].replay_count < budget.max_replays_per_episode)
.collect();
sorted_indices.sort_by(|&a, &b| {
engine.buffer.entries[b]
.created_at
.cmp(&engine.buffer.entries[a].created_at)
});
sorted_indices.into_iter().take(n).collect()
}
}
}
pub fn re_evaluate_episode(
entry: &ReplayEntry,
current_belief_confidences: &[(NodeId, f64)],
) -> ReplayOutcome {
let mut belief_updates = Vec::new();
let mut insights = Vec::new();
let belief_map: HashMap<NodeId, f64> = current_belief_confidences.iter().cloned().collect();
for &affected in &entry.outcome.affected_nodes {
if let Some(¤t_conf) = belief_map.get(&affected) {
let should_increase = entry.outcome.utility > entry.expected_utility;
let delta = if should_increase { 0.1 } else { -0.1 };
if (current_conf < 0.8 && should_increase) || (current_conf > 0.2 && !should_increase) {
belief_updates.push(BeliefDelta {
belief_id: affected,
confidence_delta: delta * entry.td_error,
reason: format!(
"Replay of episode {:?}: outcome was {} than expected (TD error {:.2})",
entry.episode_id,
if should_increase { "better" } else { "worse" },
entry.td_error
),
});
}
}
}
let mut causal_updates = Vec::new();
for &involved in &entry.action.involved_nodes {
for &affected in &entry.outcome.affected_nodes {
if involved != affected {
let strength_delta = if entry.outcome.expected {
0.05 } else {
-0.05 };
causal_updates.push(CausalDelta {
cause: involved,
effect: affected,
strength_delta,
});
}
}
}
if entry.td_error > 0.3 {
let direction = if entry.outcome.utility > entry.expected_utility {
"positively"
} else {
"negatively"
};
insights.push(format!(
"Episode {:?} {} surprised us: expected utility {:.2}, got {:.2}. Domain: {}.",
entry.episode_id,
direction,
entry.expected_utility,
entry.outcome.utility,
entry.action.domain
));
}
let new_td_error = entry.td_error * 0.9;
ReplayOutcome {
episode_id: entry.episode_id,
new_td_error,
belief_updates,
causal_updates,
associations: Vec::new(), insights,
}
}
pub fn discover_cross_associations(
outcomes: &[ReplayOutcome],
entries: &[&ReplayEntry],
) -> Vec<(NodeId, NodeId, f64)> {
let mut associations = Vec::new();
for i in 0..entries.len() {
for j in (i + 1)..entries.len() {
let a = &entries[i];
let b = &entries[j];
if a.action.domain == b.action.domain {
continue;
}
let shared_nodes: Vec<NodeId> = a
.outcome
.affected_nodes
.iter()
.filter(|n| b.outcome.affected_nodes.contains(n))
.cloned()
.collect();
if !shared_nodes.is_empty() {
let td_similarity = 1.0 - (a.td_error - b.td_error).abs().min(1.0);
let node_overlap = shared_nodes.len() as f64
/ (a.outcome.affected_nodes.len() + b.outcome.affected_nodes.len()) as f64;
let similarity = 0.5 * td_similarity + 0.5 * node_overlap;
if similarity > 0.2 {
for node in &shared_nodes {
associations.push((a.episode_id, b.episode_id, similarity));
}
}
}
}
}
associations.sort_by(|a, b| b.2.total_cmp(&a.2));
associations.dedup_by(|a, b| a.0 == b.0 && a.1 == b.1);
associations
}
pub fn run_replay_cycle(
engine: &mut ReplayEngine,
current_beliefs: &[(NodeId, f64)],
now_ms: u64,
) -> DreamReport {
let start_ms = now_ms;
reprioritize_buffer(engine, now_ms);
let sample_indices = sample_episodes(engine, engine.budget.max_replays_per_cycle);
if sample_indices.is_empty() {
return DreamReport {
replays_executed: 0,
beliefs_updated: 0,
causal_updates: 0,
new_associations: 0,
insights: Vec::new(),
duration_ms: 0,
avg_td_error: 0.0,
};
}
let entries: Vec<ReplayEntry> = sample_indices
.iter()
.map(|&i| engine.buffer.entries[i].clone())
.collect();
let mut outcomes = Vec::new();
let mut total_td_error = 0.0;
for entry in &entries {
let outcome = re_evaluate_episode(entry, current_beliefs);
total_td_error += outcome.new_td_error;
outcomes.push(outcome);
}
let entry_refs: Vec<&ReplayEntry> = entries.iter().collect();
let associations = discover_cross_associations(&outcomes, &entry_refs);
let total_belief_updates: usize = outcomes.iter().map(|o| o.belief_updates.len()).sum();
let total_causal_updates: usize = outcomes.iter().map(|o| o.causal_updates.len()).sum();
let all_insights: Vec<String> = outcomes.iter().flat_map(|o| o.insights.clone()).collect();
for (idx_pos, &buf_idx) in sample_indices.iter().enumerate() {
if buf_idx < engine.buffer.entries.len() && idx_pos < outcomes.len() {
engine.buffer.entries[buf_idx].td_error = outcomes[idx_pos].new_td_error;
engine.buffer.entries[buf_idx].replay_count += 1;
engine.buffer.entries[buf_idx].last_replayed_at = now_ms;
}
}
engine.last_replay_at = now_ms;
engine.stats.total_cycles += 1;
engine.stats.total_replays += entries.len() as u64;
engine.stats.total_belief_updates += total_belief_updates as u64;
engine.stats.total_causal_updates += total_causal_updates as u64;
engine.stats.total_associations += associations.len() as u64;
let replays_executed = entries.len();
let avg_td_error = if replays_executed > 0 {
total_td_error / replays_executed as f64
} else {
0.0
};
DreamReport {
replays_executed,
beliefs_updated: total_belief_updates,
causal_updates: total_causal_updates,
new_associations: associations.len(),
insights: all_insights,
duration_ms: now_ms.saturating_sub(start_ms),
avg_td_error,
}
}
pub fn buffer_maintenance(engine: &mut ReplayEngine, now_ms: u64) -> usize {
let max_age = engine.budget.max_replay_age_ms;
let max_replays = engine.budget.max_replays_per_episode;
let initial_len = engine.buffer.entries.len();
engine.buffer.entries.retain(|entry| {
let age = now_ms.saturating_sub(entry.created_at);
let not_expired = age < max_age;
let not_exhausted = entry.replay_count < max_replays;
let still_surprising = entry.td_error > 0.01;
not_expired && not_exhausted && still_surprising
});
reprioritize_buffer(engine, now_ms);
if engine.buffer.entries.len() > engine.buffer.capacity {
engine.buffer.entries.truncate(engine.buffer.capacity);
}
initial_len - engine.buffer.entries.len()
}
pub fn replay_summary(engine: &ReplayEngine) -> ReplaySummary {
let avg_td = if engine.buffer.entries.is_empty() {
0.0
} else {
engine
.buffer
.entries
.iter()
.map(|e| e.td_error)
.sum::<f64>()
/ engine.buffer.entries.len() as f64
};
let max_td = engine
.buffer
.entries
.iter()
.map(|e| e.td_error)
.fold(0.0f64, f64::max);
let mut domain_counts: HashMap<String, usize> = HashMap::new();
for entry in &engine.buffer.entries {
*domain_counts
.entry(entry.action.domain.clone())
.or_insert(0) += 1;
}
ReplaySummary {
buffer_size: engine.buffer.entries.len(),
buffer_capacity: engine.buffer.capacity,
avg_td_error: avg_td,
max_td_error: max_td,
total_cycles: engine.stats.total_cycles,
total_replays: engine.stats.total_replays,
domain_distribution: domain_counts,
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ReplaySummary {
pub buffer_size: usize,
pub buffer_capacity: usize,
pub avg_td_error: f64,
pub max_td_error: f64,
pub total_cycles: u64,
pub total_replays: u64,
pub domain_distribution: HashMap<String, usize>,
}
#[cfg(test)]
mod tests {
use super::*;
use crate::state::{NodeId, NodeKind};
fn episode(seq: u32) -> NodeId {
NodeId::new(NodeKind::Episode, seq)
}
fn entity(seq: u32) -> NodeId {
NodeId::new(NodeKind::Entity, seq)
}
fn belief(seq: u32) -> NodeId {
NodeId::new(NodeKind::Belief, seq)
}
fn make_action(domain: &str, nodes: &[NodeId]) -> ActionRecord {
ActionRecord {
description: format!("test action in {}", domain),
domain: domain.to_string(),
involved_nodes: nodes.to_vec(),
}
}
fn make_outcome(
utility: f64,
expected: bool,
domains: &[&str],
nodes: &[NodeId],
) -> OutcomeData {
OutcomeData {
utility,
expected,
domains: domains.iter().map(|s| s.to_string()).collect(),
affected_nodes: nodes.to_vec(),
}
}
#[test]
fn test_td_error_computation() {
assert!((compute_td_error(0.5, 0.8) - 0.3).abs() < 0.01);
assert!((compute_td_error(0.8, 0.2) - 0.6).abs() < 0.01);
assert!((compute_td_error(0.5, 0.5) - 0.0).abs() < 0.01);
}
#[test]
fn test_add_to_buffer() {
let mut engine = ReplayEngine::new();
add_to_buffer(
&mut engine,
episode(1),
0.5,
make_action("work", &[entity(1)]),
make_outcome(0.9, false, &["work"], &[belief(1)]),
1000,
);
assert_eq!(engine.buffer.len(), 1);
assert!(engine.buffer.entries[0].td_error > 0.0);
}
#[test]
fn test_add_below_threshold_ignored() {
let mut engine = ReplayEngine::new();
add_to_buffer(
&mut engine,
episode(1),
0.50,
make_action("work", &[entity(1)]),
make_outcome(0.51, true, &["work"], &[belief(1)]),
1000,
);
assert!(engine.buffer.is_empty());
}
#[test]
fn test_buffer_capacity_eviction() {
let mut engine = ReplayEngine::with_config(
3, SamplingStrategy::PrioritizedByTDError,
ReplayBudget::default(),
);
for i in 0..4 {
let td = 0.1 + i as f64 * 0.2; add_to_buffer(
&mut engine,
episode(i),
0.5,
make_action("work", &[entity(i)]),
make_outcome(0.5 + td, false, &["work"], &[belief(i)]),
1000 + i as u64 * 1000,
);
}
assert_eq!(engine.buffer.len(), 3);
}
#[test]
fn test_priority_td_error_dominates() {
let high_td = compute_priority(0.8, 1000, 900, 0, 0.6);
let low_td = compute_priority(0.1, 1000, 900, 0, 0.6);
assert!(high_td > low_td);
}
#[test]
fn test_priority_recency_matters() {
let recent = compute_priority(0.5, 10000, 9000, 0, 0.6);
let old = compute_priority(0.5, 10000, 1000, 0, 0.6);
assert!(recent > old);
}
#[test]
fn test_priority_novelty_matters() {
let fresh = compute_priority(0.5, 1000, 900, 0, 0.6);
let stale = compute_priority(0.5, 1000, 900, 5, 0.6);
assert!(fresh > stale);
}
#[test]
fn test_should_replay_empty_buffer() {
let engine = ReplayEngine::new();
assert!(!should_replay(&engine, 100000));
}
#[test]
fn test_should_replay_after_idle() {
let mut engine = ReplayEngine::new();
add_to_buffer(
&mut engine,
episode(1),
0.5,
make_action("work", &[entity(1)]),
make_outcome(0.9, false, &["work"], &[belief(1)]),
1000,
);
engine.last_replay_at = 1000;
assert!(!should_replay(&engine, 2000));
assert!(should_replay(&engine, 50000));
}
#[test]
fn test_re_evaluate_surprising_episode() {
let entry = ReplayEntry {
episode_id: episode(1),
expected_utility: 0.3,
action: make_action("work", &[entity(1)]),
outcome: make_outcome(0.8, false, &["work"], &[belief(1), belief(2)]),
td_error: 0.5,
replay_count: 0,
last_replayed_at: 0,
created_at: 1000,
priority: 1.0,
};
let beliefs = vec![(belief(1), 0.5), (belief(2), 0.6)];
let outcome = re_evaluate_episode(&entry, &beliefs);
assert!(!outcome.belief_updates.is_empty());
assert!(!outcome.insights.is_empty()); assert!(outcome.new_td_error < entry.td_error); }
#[test]
fn test_re_evaluate_expected_episode() {
let entry = ReplayEntry {
episode_id: episode(1),
expected_utility: 0.5,
action: make_action("work", &[entity(1)]),
outcome: make_outcome(0.55, true, &["work"], &[belief(1)]),
td_error: 0.05,
replay_count: 0,
last_replayed_at: 0,
created_at: 1000,
priority: 0.5,
};
let beliefs = vec![(belief(1), 0.7)];
let outcome = re_evaluate_episode(&entry, &beliefs);
assert!(outcome.insights.is_empty());
}
#[test]
fn test_cross_domain_association_discovery() {
let shared_node = belief(99);
let entries = vec![
ReplayEntry {
episode_id: episode(1),
expected_utility: 0.5,
action: make_action("work", &[entity(1)]),
outcome: make_outcome(0.8, false, &["work"], &[shared_node, belief(1)]),
td_error: 0.4,
replay_count: 0,
last_replayed_at: 0,
created_at: 1000,
priority: 1.0,
},
ReplayEntry {
episode_id: episode(2),
expected_utility: 0.6,
action: make_action("health", &[entity(2)]),
outcome: make_outcome(0.9, false, &["health"], &[shared_node, belief(2)]),
td_error: 0.35,
replay_count: 0,
last_replayed_at: 0,
created_at: 2000,
priority: 0.9,
},
];
let outcomes: Vec<ReplayOutcome> = entries
.iter()
.map(|e| re_evaluate_episode(e, &[]))
.collect();
let entry_refs: Vec<&ReplayEntry> = entries.iter().collect();
let assocs = discover_cross_associations(&outcomes, &entry_refs);
assert!(!assocs.is_empty());
}
#[test]
fn test_full_replay_cycle() {
let mut engine = ReplayEngine::new();
for i in 0..5 {
add_to_buffer(
&mut engine,
episode(i),
0.5,
make_action("work", &[entity(i)]),
make_outcome(0.9 - i as f64 * 0.1, false, &["work"], &[belief(i)]),
1000 + i as u64 * 1000,
);
}
let beliefs: Vec<(NodeId, f64)> = (0..5).map(|i| (belief(i), 0.5)).collect();
let report = run_replay_cycle(&mut engine, &beliefs, 100000);
assert!(report.replays_executed > 0);
assert_eq!(engine.stats.total_cycles, 1);
}
#[test]
fn test_maintenance_removes_expired() {
let mut engine = ReplayEngine::new();
add_to_buffer(
&mut engine,
episode(1),
0.5,
make_action("work", &[entity(1)]),
make_outcome(0.9, false, &["work"], &[belief(1)]),
1000,
);
let far_future = 1000 + engine.budget.max_replay_age_ms + 1;
let removed = buffer_maintenance(&mut engine, far_future);
assert_eq!(removed, 1);
assert!(engine.buffer.is_empty());
}
#[test]
fn test_maintenance_keeps_recent() {
let mut engine = ReplayEngine::new();
add_to_buffer(
&mut engine,
episode(1),
0.5,
make_action("work", &[entity(1)]),
make_outcome(0.9, false, &["work"], &[belief(1)]),
1000,
);
let removed = buffer_maintenance(&mut engine, 2000);
assert_eq!(removed, 0);
assert_eq!(engine.buffer.len(), 1);
}
#[test]
fn test_replay_summary() {
let mut engine = ReplayEngine::new();
add_to_buffer(
&mut engine,
episode(1),
0.3,
make_action("work", &[entity(1)]),
make_outcome(0.9, false, &["work"], &[belief(1)]),
1000,
);
add_to_buffer(
&mut engine,
episode(2),
0.4,
make_action("health", &[entity(2)]),
make_outcome(0.8, false, &["health"], &[belief(2)]),
2000,
);
let summary = replay_summary(&engine);
assert_eq!(summary.buffer_size, 2);
assert!(summary.avg_td_error > 0.0);
assert_eq!(summary.domain_distribution.len(), 2);
}
}