use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use agent_client_protocol::schema::v1::{ContentBlock, SessionNotification, SessionUpdate};
use serde_json::{Value, json};
use sqlx::SqlitePool;
use tokio::sync::mpsc;
use uuid::Uuid;
use crate::acp::chat_persistence;
const DEBOUNCE: Duration = Duration::from_millis(250);
const MAX_LATENCY: Duration = Duration::from_millis(1000);
const WRITER_CHANNEL_CAPACITY: usize = 256;
const RAW_FRAMES_VERSION: u64 = 1;
#[derive(Default)]
struct TurnState {
active: bool,
row_id: Option<String>,
frames: Vec<Value>,
text: String,
seq: u64,
dirty: bool,
}
struct Sink {
cmd_tx: mpsc::Sender<WriterCmd>,
}
enum WriterCmd {
Flush,
Finalize,
}
pub struct TurnAccumulator {
inner: Mutex<TurnState>,
sink: Mutex<Option<Sink>>,
}
struct FlushSnapshot {
row_id: String,
text: String,
blocks: String,
last_seq: i64,
}
impl TurnAccumulator {
pub fn new() -> Self {
Self { inner: Mutex::new(TurnState::default()), sink: Mutex::new(None) }
}
pub fn attach_persistence(self: &Arc<Self>, db: SqlitePool, db_session_id: String) {
let (cmd_tx, cmd_rx) = mpsc::channel(WRITER_CHANNEL_CAPACITY);
if let Ok(mut guard) = self.sink.lock() {
*guard = Some(Sink { cmd_tx });
}
let acc = self.clone();
tokio::spawn(writer_loop(acc, db, db_session_id, cmd_rx));
}
pub fn begin_turn(&self) {
if let Ok(mut st) = self.inner.lock() {
st.active = true;
st.row_id = None;
st.frames.clear();
st.text.clear();
st.dirty = false;
}
}
pub fn fold(&self, notification: &SessionNotification) -> Option<u64> {
let seq = {
let Ok(mut st) = self.inner.lock() else { return None };
if !st.active {
return None;
}
st.seq += 1;
if st.row_id.is_none() {
st.row_id = Some(Uuid::new_v4().to_string());
}
if let Some(text) = agent_message_text(¬ification.update) {
st.text.push_str(text);
}
match serde_json::to_value(¬ification.update) {
Ok(v) => st.frames.push(v),
Err(e) => {
tracing::warn!("failed to serialize session update for persistence: {}", e)
}
}
st.dirty = true;
st.seq
};
self.send_cmd(WriterCmd::Flush);
Some(seq)
}
pub fn finalize_turn(&self) {
let should_finalize = {
let Ok(mut st) = self.inner.lock() else { return };
if !st.active {
return;
}
st.active = false;
st.row_id.is_some()
};
if should_finalize {
self.send_cmd(WriterCmd::Finalize);
}
}
fn send_cmd(&self, cmd: WriterCmd) {
if let Ok(guard) = self.sink.lock()
&& let Some(sink) = guard.as_ref()
{
let _ = sink.cmd_tx.try_send(cmd);
}
}
fn snapshot(&self) -> Option<FlushSnapshot> {
let st = self.inner.lock().ok()?;
let row_id = st.row_id.clone()?;
let blocks = json!({ "v": RAW_FRAMES_VERSION, "frames": st.frames }).to_string();
Some(FlushSnapshot { row_id, text: st.text.clone(), blocks, last_seq: st.seq as i64 })
}
pub fn turn_snapshot(&self) -> TurnSnapshot {
let Ok(st) = self.inner.lock() else {
return TurnSnapshot::default();
};
let blocks = json!({ "v": RAW_FRAMES_VERSION, "frames": st.frames }).to_string();
TurnSnapshot {
active: st.active,
row_id: st.row_id.clone(),
text: st.text.clone(),
blocks,
seq: st.seq,
}
}
}
#[derive(Default)]
pub struct TurnSnapshot {
pub active: bool,
pub row_id: Option<String>,
pub text: String,
pub blocks: String,
pub seq: u64,
}
impl Default for TurnAccumulator {
fn default() -> Self {
Self::new()
}
}
fn agent_message_text(update: &SessionUpdate) -> Option<&str> {
if let SessionUpdate::AgentMessageChunk(chunk) = update
&& let ContentBlock::Text(t) = &chunk.content
{
return Some(&t.text);
}
None
}
async fn writer_loop(
acc: Arc<TurnAccumulator>,
db: SqlitePool,
session_id: String,
mut cmd_rx: mpsc::Receiver<WriterCmd>,
) {
let mut pending_since: Option<Instant> = None;
loop {
let cmd = if let Some(since) = pending_since {
let elapsed = since.elapsed();
let wait = if elapsed >= MAX_LATENCY {
Duration::ZERO
} else {
DEBOUNCE.min(MAX_LATENCY - elapsed)
};
match tokio::time::timeout(wait, cmd_rx.recv()).await {
Err(_elapsed) => {
flush_once(&acc, &db, &session_id).await;
pending_since = None;
continue;
}
Ok(None) => break, Ok(Some(c)) => c,
}
} else {
match cmd_rx.recv().await {
Some(c) => c,
None => break,
}
};
match cmd {
WriterCmd::Flush => {
if pending_since.is_none() {
pending_since = Some(Instant::now());
}
}
WriterCmd::Finalize => {
flush_once(&acc, &db, &session_id).await;
pending_since = None;
if let Some(snap) = acc.snapshot()
&& let Err(e) = chat_persistence::finalize_message(&db, &snap.row_id).await
{
tracing::warn!("failed to finalize streaming chat message: {}", e);
}
}
}
}
}
async fn flush_once(acc: &Arc<TurnAccumulator>, db: &SqlitePool, session_id: &str) {
let Some(snap) = acc.snapshot() else { return };
if let Err(e) = chat_persistence::upsert_streaming_message(
db,
&snap.row_id,
session_id,
&snap.text,
Some(&snap.blocks),
snap.last_seq,
)
.await
{
tracing::warn!("failed to upsert streaming chat message: {}", e);
}
}