use std::sync::Arc; use std::{fmt::Display, future::Future}; use tokio::sync::Mutex; use super::session::{MessagePersistSnapshot, Session}; use crate::bus::ChatMessage; use crate::storage::StorageError; use crate::{providers::Usage, session::TurnController}; async fn persist_added_messages( snapshots: Vec>, ) -> Result<(), StorageError> { let mut storage = None; let mut session_id = None; let mut messages = Vec::new(); let mut final_meta = None; for snapshot in snapshots.into_iter().flatten() { let (snapshot_storage, snapshot_session_id, message, meta) = snapshot; if let Some(ref expected) = session_id && expected != &snapshot_session_id { return Err(StorageError::Serialization( "attempted to persist messages from different sessions in one turn".to_string(), )); } storage = Some(snapshot_storage); session_id = Some(snapshot_session_id); messages.push(message); final_meta = Some(meta); } let (Some(storage), Some(session_id), Some(final_meta)) = (storage, session_id, final_meta) else { return Ok(()); }; storage .persist_message_batch_with_retry(&session_id, &messages, &final_meta) .await } pub(super) async fn append_persisted_messages( session: &Arc>, messages: Vec, ) -> Result<(), StorageError> { append_persisted_messages_with_meta(session, messages) .await .map(|_| ()) } enum VersionPolicy { Advance, PreserveForOwnedTurn(String), } /// Persist a message emitted as a side effect of the currently active Turn /// without invalidating that same Turn's captured session version. /// /// The task-local Turn ID is revalidated while holding the session lock. If it /// no longer owns the target session, the write advances `state_version` like /// any unrelated mutation so stale model work is still rejected. pub(super) async fn append_active_turn_message( session: &Arc>, message: ChatMessage, turn_id: String, ) -> Result<(), StorageError> { append_persisted_messages_inner( session, vec![message], VersionPolicy::PreserveForOwnedTurn(turn_id), ) .await .map(|_| ()) } pub(super) async fn append_persisted_messages_with_meta( session: &Arc>, messages: Vec, ) -> Result, StorageError> { append_persisted_messages_inner(session, messages, VersionPolicy::Advance).await } async fn append_persisted_messages_inner( session: &Arc>, messages: Vec, version_policy: VersionPolicy, ) -> Result, StorageError> { if messages.is_empty() { return Ok(Vec::new()); } let persistence_lock = { session.lock().await.persistence_lock.clone() }; let _persistence_guard = persistence_lock.lock().await; let message_ids: Vec<_> = messages.iter().map(|message| message.id.clone()).collect(); let (snapshots, advance_state_version): (Vec>, bool) = { let mut guard = session.lock().await; let advance_state_version = match &version_policy { VersionPolicy::Advance => true, VersionPolicy::PreserveForOwnedTurn(turn_id) => !guard.owns_active_turn(turn_id), }; let snapshots = messages .into_iter() .map(|message| { guard.add_message_in_memory_with_version(message, true, advance_state_version) }) .collect(); (snapshots, advance_state_version) }; let committed = snapshots .iter() .flatten() .map(|(_, _, message, _)| message.clone()) .collect(); if let Err(error) = persist_added_messages(snapshots).await { session .lock() .await .rollback_message_suffix_with_version(&message_ids, advance_state_version); return Err(error); } Ok(committed) } /// Publish `Completed` only after the supplied durable write succeeds. /// /// Keeping this ordering in one helper makes the user-visible terminal status /// impossible to publish optimistically before SQLite commits. pub(super) async fn finalize_turn_after_persistence( controller: &TurnController, usage: Option, persistence: F, ) -> Result where F: Future>, E: Display, { controller.begin_finalizing(); match persistence.await { Ok(value) => { controller.complete(usage); Ok(value) } Err(error) => { controller.fail(format!("failed to persist turn: {error}")); Err(error) } } } #[cfg(test)] mod tests { use super::*; use crate::config::LLMProviderConfig; use crate::memory::MemoryManager; use crate::session::{TurnController, TurnStatus}; use crate::tools::ToolRegistry; use std::collections::HashMap; use std::path::PathBuf; #[tokio::test] async fn completed_is_published_only_after_persistence_succeeds() { let (controller, _emitter, receiver) = TurnController::start("session", "message"); let persisted = Arc::new(std::sync::atomic::AtomicBool::new(false)); let persisted_in_future = persisted.clone(); let result: Result<(), String> = finalize_turn_after_persistence(&controller, None, async move { assert_eq!(receiver.borrow().status, TurnStatus::Running); assert_eq!( receiver.borrow().phase, crate::session::TurnPhase::Finalizing ); persisted_in_future.store(true, std::sync::atomic::Ordering::SeqCst); Ok(()) }) .await; assert!(result.is_ok()); assert!(persisted.load(std::sync::atomic::Ordering::SeqCst)); assert_eq!(controller.snapshot().status, TurnStatus::Completed); } #[tokio::test] async fn persistence_failure_never_publishes_completed() { let (controller, _emitter, _receiver) = TurnController::start("session", "message"); let result: Result<(), &str> = finalize_turn_after_persistence(&controller, None, async { Err("database down") }) .await; assert_eq!(result, Err("database down")); let snapshot = controller.snapshot(); assert_eq!(snapshot.status, TurnStatus::Failed); assert_eq!( snapshot.error.as_deref(), Some("failed to persist turn: database down") ); } #[tokio::test] async fn active_turn_side_effect_does_not_invalidate_its_session_version() { let dir = tempfile::tempdir().unwrap(); let storage = Arc::new( crate::storage::Storage::new(&dir.path().join("memory.db")) .await .unwrap(), ); let memory_manager = Arc::new(MemoryManager::new( storage, "test".to_string(), "test".to_string(), )); let config = LLMProviderConfig { provider_type: "openai".to_string(), name: "test".to_string(), base_url: "http://127.0.0.1".to_string(), api_key: "test".to_string(), extra_headers: HashMap::new(), model_id: "test".to_string(), temperature: None, max_tokens: None, model_extra: HashMap::new(), max_tool_iterations: 1, token_limit: 8_192, workspace_dir: PathBuf::from("."), input_types: vec!["text".to_string(), "image".to_string()], price_input_per_million: None, price_output_per_million: None, }; let session = Arc::new(Mutex::new( Session::new( crate::session::UnifiedSessionId::new("cli_chat", "chat", "dialog"), config, Arc::new(ToolRegistry::new()), None, String::new(), "test".to_string(), memory_manager, ) .await .unwrap(), )); let base_version = session.lock().await.state_version_for_test(); session.lock().await.set_active_turn_for_test("turn-1"); append_active_turn_message( &session, ChatMessage::assistant("sent screenshot"), "turn-1".to_string(), ) .await .unwrap(); { let guard = session.lock().await; assert_eq!(guard.state_version_for_test(), base_version); assert_eq!(guard.get_history().len(), 1); } append_active_turn_message( &session, ChatMessage::assistant("late screenshot"), "stale-turn".to_string(), ) .await .unwrap(); let guard = session.lock().await; assert_eq!(guard.state_version_for_test(), base_version + 1); assert_eq!(guard.get_history().len(), 2); } }