use std::path::Path;
use crate::error::SessionError;
use crate::event::SessionEvent;
use crate::log::SessionEventLog;
use crate::replay::ReplayEngine;
use crate::store::SessionStore;
#[derive(Debug, Clone)]
pub struct ForkResult {
pub new_session_id: String,
pub events_copied: usize,
}
pub struct ForkEngine;
impl ForkEngine {
#[tracing::instrument(name = "session.fork.run", skip_all, level = "info", fields(at_seq))]
pub async fn fork(
data_dir: &Path,
src_id: &str,
new_id: &str,
at_seq: Option<u64>,
store: &SessionStore,
) -> Result<ForkResult, SessionError> {
if store.get(src_id).await?.is_none() {
return Err(SessionError::NotFound(src_id.to_owned()));
}
let src_dir = crate::session_dir(data_dir, src_id);
let src_log = SessionEventLog::open(&src_dir).await?;
let all_events = src_log.read_all().await?;
let total = u64::try_from(all_events.len()).unwrap_or(u64::MAX);
let at_seq = at_seq.unwrap_or(total);
if at_seq > total {
return Err(SessionError::InvalidForkPoint(format!(
"at_seq={at_seq} exceeds source session's event count={total}"
)));
}
ReplayEngine::replay(&src_dir, Some(at_seq)).await?;
let take_n = usize::try_from(at_seq).unwrap_or(usize::MAX);
let to_copy: Vec<_> = all_events.iter().take(take_n).cloned().collect();
let (cwd, provider_name, model) = to_copy
.iter()
.find_map(|e| match &e.kind {
SessionEvent::SessionStarted {
cwd,
provider_name,
model,
..
} => Some((cwd.clone(), provider_name.clone(), model.clone())),
_ => None,
})
.unwrap_or_default();
let child_dir = crate::session_dir(data_dir, new_id);
let child_log = SessionEventLog::open(&child_dir).await?;
child_log
.append(
None,
None,
SessionEvent::SessionStarted {
session_id: new_id.to_owned(),
cwd,
provider_name,
model,
forked_from: Some((src_id.to_owned(), at_seq)),
},
)
.await?;
for envelope in &to_copy {
child_log
.append(envelope.turn_id, envelope.parent_seq, envelope.kind.clone())
.await?;
}
store.record_fork(new_id, src_id, at_seq).await?;
store
.update_seq(
new_id,
child_log.last_seq().unwrap_or(0),
to_copy.len() as u64 + 1,
)
.await?;
src_log
.append(
None,
None,
SessionEvent::ForkPoint {
new_session_id: new_id.to_owned(),
},
)
.await?;
Ok(ForkResult {
new_session_id: new_id.to_owned(),
events_copied: to_copy.len(),
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::store::SessionStore;
async fn make_pool() -> zeph_db::DbPool {
let config = zeph_db::DbConfig {
url: ":memory:".to_owned(),
..Default::default()
};
let pool = config
.connect()
.await
.expect("connect in-memory sqlite pool");
zeph_db::run_migrations(&pool)
.await
.expect("run migrations");
pool
}
async fn seed_parent(data_dir: &Path, store: &SessionStore, id: &str) {
store.create(id).await.unwrap();
let dir = crate::session_dir(data_dir, id);
let log = SessionEventLog::open(&dir).await.unwrap();
log.append(
None,
None,
SessionEvent::SessionStarted {
session_id: id.to_owned(),
cwd: "/repo".to_owned(),
provider_name: "claude".to_owned(),
model: "opus".to_owned(),
forked_from: None,
},
)
.await
.unwrap();
log.append(
None,
None,
SessionEvent::UserMessage {
text: "hello".to_owned(),
image_refs: vec![],
},
)
.await
.unwrap();
log.append(
None,
None,
SessionEvent::AssistantMessage {
parts: vec![zeph_llm::provider::MessagePart::Text {
text: "hi".to_owned(),
}],
},
)
.await
.unwrap();
log.append(
None,
None,
SessionEvent::UserMessage {
text: "second turn".to_owned(),
image_refs: vec![],
},
)
.await
.unwrap();
store
.update_seq(id, log.last_seq().unwrap(), 4)
.await
.unwrap();
}
#[tokio::test]
async fn test_fork_copies_events() {
let store = SessionStore::new(make_pool().await);
let data_dir = tempfile::tempdir().unwrap();
seed_parent(data_dir.path(), &store, "parent").await;
let result = ForkEngine::fork(data_dir.path(), "parent", "child", Some(3), &store)
.await
.unwrap();
assert_eq!(result.events_copied, 3);
assert_eq!(result.new_session_id, "child");
let child_dir = crate::session_dir(data_dir.path(), &result.new_session_id);
let child_log = SessionEventLog::open(&child_dir).await.unwrap();
let events = child_log.read_all().await.unwrap();
assert_eq!(events.len(), 4);
}
#[tokio::test]
async fn test_fork_provenance_metadata() {
let store = SessionStore::new(make_pool().await);
let data_dir = tempfile::tempdir().unwrap();
seed_parent(data_dir.path(), &store, "parent").await;
ForkEngine::fork(data_dir.path(), "parent", "child", Some(2), &store)
.await
.unwrap();
let meta = store.get("child").await.unwrap().unwrap();
assert_eq!(meta.forked_from.as_deref(), Some("parent"));
assert_eq!(meta.forked_at_seq, Some(2));
}
#[tokio::test]
async fn test_fork_appends_forkpoint_to_parent() {
let store = SessionStore::new(make_pool().await);
let data_dir = tempfile::tempdir().unwrap();
seed_parent(data_dir.path(), &store, "parent").await;
ForkEngine::fork(data_dir.path(), "parent", "child", Some(2), &store)
.await
.unwrap();
let parent_dir = crate::session_dir(data_dir.path(), "parent");
let parent_log = SessionEventLog::open(&parent_dir).await.unwrap();
let events = parent_log.read_all().await.unwrap();
assert!(matches!(
events.last().unwrap().kind,
SessionEvent::ForkPoint { .. }
));
}
#[tokio::test]
async fn test_fork_rejects_seq_beyond_source() {
let store = SessionStore::new(make_pool().await);
let data_dir = tempfile::tempdir().unwrap();
seed_parent(data_dir.path(), &store, "parent").await;
let err = ForkEngine::fork(data_dir.path(), "parent", "child", Some(100), &store)
.await
.unwrap_err();
assert!(matches!(err, SessionError::InvalidForkPoint(_)));
}
#[tokio::test]
async fn test_fork_rejects_unknown_source() {
let store = SessionStore::new(make_pool().await);
let data_dir = tempfile::tempdir().unwrap();
let err = ForkEngine::fork(data_dir.path(), "no-such", "child", Some(0), &store)
.await
.unwrap_err();
assert!(matches!(err, SessionError::NotFound(_)));
}
#[tokio::test]
async fn test_fork_none_copies_everything() {
let store = SessionStore::new(make_pool().await);
let data_dir = tempfile::tempdir().unwrap();
seed_parent(data_dir.path(), &store, "parent").await;
let result = ForkEngine::fork(data_dir.path(), "parent", "child", None, &store)
.await
.unwrap();
assert_eq!(result.events_copied, 4);
}
}