PicoBot/src/session/persistence.rs

291 lines
9.7 KiB
Rust

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<Option<MessagePersistSnapshot>>,
usage: Option<&crate::storage::TurnUsageRecord>,
) -> 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(());
};
if let Some(usage) = usage {
storage
.persist_turn_batch_with_retry(&session_id, &messages, &final_meta, usage)
.await
} else {
storage
.persist_message_batch_with_retry(&session_id, &messages, &final_meta)
.await
}
}
pub(super) async fn append_persisted_messages(
session: &Arc<Mutex<Session>>,
messages: Vec<ChatMessage>,
) -> 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<Mutex<Session>>,
message: ChatMessage,
turn_id: String,
) -> Result<(), StorageError> {
append_persisted_messages_inner(
session,
vec![message],
VersionPolicy::PreserveForOwnedTurn(turn_id),
None,
)
.await
.map(|_| ())
}
pub(super) async fn append_persisted_messages_with_meta(
session: &Arc<Mutex<Session>>,
messages: Vec<ChatMessage>,
) -> Result<Vec<crate::storage::message::MessageMeta>, StorageError> {
append_persisted_messages_inner(session, messages, VersionPolicy::Advance, None).await
}
pub(super) async fn append_persisted_turn_messages(
session: &Arc<Mutex<Session>>,
messages: Vec<ChatMessage>,
usage: crate::storage::TurnUsageRecord,
) -> Result<Vec<crate::storage::message::MessageMeta>, StorageError> {
append_persisted_messages_inner(session, messages, VersionPolicy::Advance, Some(usage)).await
}
async fn append_persisted_messages_inner(
session: &Arc<Mutex<Session>>,
messages: Vec<ChatMessage>,
version_policy: VersionPolicy,
usage: Option<crate::storage::TurnUsageRecord>,
) -> Result<Vec<crate::storage::message::MessageMeta>, 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<Option<MessagePersistSnapshot>>, 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, usage.as_ref()).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<F, T, E>(
controller: &TurnController,
usage: Option<Usage>,
persistence: F,
) -> Result<T, E>
where
F: Future<Output = Result<T, E>>,
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);
}
}