use std::collections::HashMap;
use std::sync::Arc;
use awaken_contract::contract::message::Message;
use awaken_contract::contract::storage::{RunRecord, StorageError, ThreadRunStore};
use futures::StreamExt;
use super::{config::NatsBufferedThreadConfig, entry, hot_meta};
struct ThreadBatch {
latest_messages: Vec<Message>,
latest_thread_seq: u64,
runs_by_id: HashMap<String, (RunRecord, u64)>,
msgs_to_ack: Vec<async_nats::jetstream::Message>,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum WalAckAction {
Ack,
Nak,
}
pub fn spawn_flusher<T: ThreadRunStore + Send + Sync + 'static>(
inner: Arc<T>,
consumer: async_nats::jetstream::consumer::PullConsumer,
kv_hot: async_nats::jetstream::kv::Store,
config: NatsBufferedThreadConfig,
flush_notify: Arc<tokio::sync::Notify>,
mut shutdown_rx: tokio::sync::watch::Receiver<bool>,
) -> tokio::task::JoinHandle<()> {
tokio::spawn(async move {
let mut messages = match consumer.messages().await {
Ok(m) => m,
Err(e) => {
tracing::warn!(error = %e, "flusher failed to open message stream");
return;
}
};
loop {
let mut by_thread: HashMap<String, ThreadBatch> = HashMap::new();
let window = tokio::time::sleep(config.flush_interval);
tokio::pin!(window);
let mut shutdown_requested = false;
loop {
tokio::select! {
_ = shutdown_rx.changed() => {
if *shutdown_rx.borrow() {
shutdown_requested = true;
break;
}
}
_ = flush_notify.notified() => break,
_ = &mut window => break,
next = messages.next() => match next {
None => return,
Some(Err(e)) => {
tracing::warn!(error = %e, "flusher message stream error");
}
Some(Ok(msg)) => {
if let Some(decoded) = decode_or_ack(&msg).await {
merge_entry(&mut by_thread, decoded, msg);
if by_thread.len() >= config.flush_batch_size {
break;
}
}
}
}
}
}
if !by_thread.is_empty() {
flush_batch(&inner, &kv_hot, by_thread).await;
}
if shutdown_requested {
return;
}
}
})
}
async fn decode_or_ack(msg: &async_nats::jetstream::Message) -> Option<entry::CheckpointEntry> {
match entry::decode(&msg.payload) {
Ok(e) => Some(e),
Err(err) => {
tracing::warn!(error = %err, "poison entry; acking to skip");
let _ = msg.ack().await;
None
}
}
}
fn wal_ack_action(checkpoints_ok: bool, watermark_ok: bool) -> WalAckAction {
if checkpoints_ok && watermark_ok {
WalAckAction::Ack
} else {
WalAckAction::Nak
}
}
async fn finish_wal_messages(messages: Vec<async_nats::jetstream::Message>, action: WalAckAction) {
for msg in messages {
let result = match action {
WalAckAction::Ack => msg.ack().await,
WalAckAction::Nak => {
msg.ack_with(async_nats::jetstream::AckKind::Nak(None))
.await
}
};
if let Err(error) = result {
tracing::warn!(%error, ?action, "failed to finish WAL message");
}
}
}
fn merge_entry(
by_thread: &mut HashMap<String, ThreadBatch>,
decoded: entry::CheckpointEntry,
msg: async_nats::jetstream::Message,
) {
let thread_id = decoded.thread_id.clone();
let batch = by_thread.entry(thread_id).or_insert_with(|| ThreadBatch {
latest_messages: Vec::new(),
latest_thread_seq: 0,
runs_by_id: HashMap::new(),
msgs_to_ack: Vec::new(),
});
if decoded.thread_seq >= batch.latest_thread_seq {
batch.latest_messages = decoded.messages;
batch.latest_thread_seq = decoded.thread_seq;
}
let run_id = decoded.run.run_id.clone();
match batch.runs_by_id.get(&run_id) {
Some((_, existing_seq)) if *existing_seq >= decoded.thread_seq => {}
_ => {
batch
.runs_by_id
.insert(run_id, (decoded.run, decoded.thread_seq));
}
}
batch.msgs_to_ack.push(msg);
}
fn order_runs_for_flush(runs_by_id: HashMap<String, (RunRecord, u64)>) -> Vec<(RunRecord, u64)> {
let mut ordered: Vec<(RunRecord, u64)> = runs_by_id.into_values().collect();
ordered.sort_by_key(|(_, seq)| *seq);
ordered
}
async fn apply_thread_batch_ordered<T: ThreadRunStore + Send + Sync>(
inner: &T,
thread_id: &str,
messages: &[Message],
ordered_runs: &[(RunRecord, u64)],
) -> bool {
for (run, _) in ordered_runs {
if let Err(e) = inner.checkpoint(thread_id, messages, run).await {
tracing::warn!(thread_id, run_id = %run.run_id, error = %e, "inner checkpoint failed");
return false;
}
}
true
}
async fn persist_stale_runs_without_projection<T: ThreadRunStore + Send + Sync>(
inner: &T,
thread_id: &str,
stale_runs: &[(RunRecord, u64)],
) -> bool {
for (run, _) in stale_runs {
match inner.load_run(&run.run_id).await {
Ok(Some(_)) => {}
Ok(None) => match inner.create_run(run).await {
Ok(()) | Err(StorageError::AlreadyExists(_)) => {}
Err(e) => {
tracing::warn!(
thread_id,
run_id = %run.run_id,
error = %e,
"inner create_run for stale WAL entry failed"
);
return false;
}
},
Err(e) => {
tracing::warn!(
thread_id,
run_id = %run.run_id,
error = %e,
"inner load_run for stale WAL entry failed"
);
return false;
}
}
}
true
}
async fn flush_batch<T: ThreadRunStore + Send + Sync + 'static>(
inner: &Arc<T>,
kv_hot: &async_nats::jetstream::kv::Store,
by_thread: HashMap<String, ThreadBatch>,
) {
for (thread_id, batch) in by_thread {
let current_flushed = match hot_meta::read_flushed_seq(kv_hot, &thread_id).await {
Ok(seq) => seq,
Err(e) => {
tracing::warn!(thread_id, error = %e, "read flushed_seq failed");
finish_wal_messages(batch.msgs_to_ack, WalAckAction::Nak).await;
continue;
}
};
let ordered = order_runs_for_flush(batch.runs_by_id);
let (stale_runs, fresh_runs): (Vec<_>, Vec<_>) = ordered
.into_iter()
.partition(|(_, seq)| *seq <= current_flushed);
if batch.latest_thread_seq <= current_flushed {
let stale_ok =
persist_stale_runs_without_projection(inner.as_ref(), &thread_id, &stale_runs)
.await;
let action = wal_ack_action(stale_ok, true);
finish_wal_messages(batch.msgs_to_ack, action).await;
continue;
}
let stale_ok =
persist_stale_runs_without_projection(inner.as_ref(), &thread_id, &stale_runs).await;
let fresh_ok = if fresh_runs.is_empty() {
true
} else {
apply_thread_batch_ordered(
inner.as_ref(),
&thread_id,
&batch.latest_messages,
&fresh_runs,
)
.await
};
let all_ok = stale_ok && fresh_ok;
let watermark_ok = if all_ok {
match hot_meta::write_flushed_seq(kv_hot, &thread_id, batch.latest_thread_seq).await {
Ok(()) => true,
Err(e) => {
tracing::warn!(
thread_id,
error = %e,
"write flushed_seq failed; nacking WAL batch for redelivery"
);
false
}
}
} else {
false
};
let action = wal_ack_action(all_ok, watermark_ok);
finish_wal_messages(batch.msgs_to_ack, action).await;
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::InMemoryStore;
use awaken_contract::contract::lifecycle::RunStatus;
use awaken_contract::contract::storage::ThreadStore;
fn mk_run(run_id: &str, thread_id: &str) -> RunRecord {
RunRecord {
run_id: run_id.into(),
thread_id: thread_id.into(),
agent_id: "a".into(),
parent_run_id: None,
request: None,
input: None,
output: None,
status: RunStatus::Created,
termination_reason: None,
final_output: None,
error_payload: None,
dispatch_id: None,
session_id: None,
transport_request_id: None,
waiting: None,
outcome: None,
created_at: 1,
started_at: None,
finished_at: None,
updated_at: 1,
steps: 0,
input_tokens: 0,
output_tokens: 0,
state: None,
}
}
#[test]
fn order_runs_for_flush_sorts_by_thread_seq() {
let mut map: HashMap<String, (RunRecord, u64)> = HashMap::new();
map.insert("a".into(), (mk_run("a", "t"), 42));
map.insert("b".into(), (mk_run("b", "t"), 10));
map.insert("c".into(), (mk_run("c", "t"), 99));
let ordered = order_runs_for_flush(map);
let seqs: Vec<u64> = ordered.iter().map(|(_, s)| *s).collect();
assert_eq!(seqs, vec![10, 42, 99]);
}
#[test]
fn watermark_write_failure_naks_wal_batch() {
assert_eq!(wal_ack_action(true, true), WalAckAction::Ack);
assert_eq!(wal_ack_action(true, false), WalAckAction::Nak);
assert_eq!(wal_ack_action(false, true), WalAckAction::Nak);
assert_eq!(wal_ack_action(false, false), WalAckAction::Nak);
}
#[tokio::test]
async fn flush_ordering_preserves_highest_seq_projection() {
let inner = InMemoryStore::new();
let older_run = mk_run("run-old", "t");
let newer_run = mk_run("run-new", "t");
let mut map: HashMap<String, (RunRecord, u64)> = HashMap::new();
map.insert("run-new".into(), (newer_run, 20));
map.insert("run-old".into(), (older_run, 10));
let ordered = order_runs_for_flush(map);
let ok = apply_thread_batch_ordered(&inner, "t", &[], &ordered).await;
assert!(ok);
let thread = ThreadStore::load_thread(&inner as &InMemoryStore, "t")
.await
.unwrap()
.unwrap();
assert_eq!(
thread.latest_run_id.as_deref(),
Some("run-new"),
"thread projection must point at highest-seq run"
);
let open = thread.open_run_id.as_deref();
assert!(
open == Some("run-new") || open.is_none(),
"open_run_id must match highest-seq projection; got {open:?}"
);
}
}