diff --git a/crates/session-store/src/fs_store.rs b/crates/session-store/src/fs_store.rs index d51cefbc..0f17fb21 100644 --- a/crates/session-store/src/fs_store.rs +++ b/crates/session-store/src/fs_store.rs @@ -764,6 +764,62 @@ mod tests { .unwrap() .contains("account-1") ); + + let replay = store + .write_uploaded_file_with_context( + session_id, + "notes.txt", + "text/plain", + b"hello", + &context, + UploadedFileLimits::default(), + ) + .unwrap(); + assert_eq!(replay.artifact_id, reference.artifact_id); + assert!(matches!( + store.write_uploaded_file_with_context( + session_id, + "renamed.txt", + "text/plain", + b"hello", + &context, + UploadedFileLimits::default(), + ), + Err(StoreError::InvalidUploadedFileName) + )); + } + + #[test] + fn uploaded_file_exact_replay_succeeds_at_session_count_limit() { + let tmp = tempfile::TempDir::new().unwrap(); + let store = FsStore::new(tmp.path()).unwrap(); + let session_id = new_session_id(); + let limits = UploadedFileLimits { + max_file_bytes: 1, + max_session_bytes: crate::DEFAULT_MAX_SESSION_UPLOADED_FILES, + }; + let mut first = None; + for index in 0..crate::DEFAULT_MAX_SESSION_UPLOADED_FILES { + let reference = store + .write_uploaded_file( + session_id, + &format!("file-{index}.txt"), + "text/plain", + b"x", + limits, + ) + .unwrap(); + first.get_or_insert(reference); + } + + let replay = store + .write_uploaded_file(session_id, "file-0.txt", "text/plain", b"x", limits) + .unwrap(); + assert_eq!(replay.artifact_id, first.unwrap().artifact_id); + assert!(matches!( + store.write_uploaded_file(session_id, "overflow.txt", "text/plain", b"x", limits), + Err(StoreError::ArtifactQuotaExceeded) + )); } #[test] diff --git a/crates/session-store/src/lib.rs b/crates/session-store/src/lib.rs index 9e34c5aa..bf89e149 100644 --- a/crates/session-store/src/lib.rs +++ b/crates/session-store/src/lib.rs @@ -27,6 +27,7 @@ //! system_prompt: None, //! config: &config, //! history: Vec::new(), +//! user_segments: Vec::new(), //! })?; //! ``` diff --git a/crates/session-store/src/segment.rs b/crates/session-store/src/segment.rs index e041ba7c..163b2317 100644 --- a/crates/session-store/src/segment.rs +++ b/crates/session-store/src/segment.rs @@ -17,6 +17,33 @@ pub struct SegmentStartState<'a> { pub system_prompt: Option<&'a str>, pub config: &'a RequestConfig, pub history: Vec, + pub user_segments: Vec>, +} + +fn seed_entries( + ts: u64, + session_id: SessionId, + state: SegmentStartState<'_>, + forked_from: Option, + compacted_from: Option, +) -> Vec { + let entry = LogEntry::AnnotatedSegmentStart { + ts, + session_id, + system_prompt: state.system_prompt.map(String::from), + config: state.config.clone(), + history: state.history, + forked_from, + compacted_from, + }; + let mut entries = vec![entry]; + if !state.user_segments.is_empty() { + entries.push(LogEntry::InputSegmentsCheckpoint { + ts, + user_segments: state.user_segments, + }); + } + entries } /// Create a new session + initial segment, writing the initial @@ -42,16 +69,8 @@ pub fn create_segment_with_ids( segment_id: SegmentId, state: SegmentStartState<'_>, ) -> Result<(), StoreError> { - let entry = LogEntry::AnnotatedSegmentStart { - ts: segment_log::now_millis(), - session_id, - system_prompt: state.system_prompt.map(String::from), - config: state.config.clone(), - history: state.history.to_vec(), - forked_from: None, - compacted_from: None, - }; - store.append(session_id, segment_id, &entry) + let entries = seed_entries(segment_log::now_millis(), session_id, state, None, None); + store.create_segment(session_id, segment_id, &entries) } /// Create a compacted segment from an existing one. Inherits the source's @@ -68,19 +87,17 @@ pub fn create_compacted_segment( source_turn_count: usize, ) -> Result { let segment_id = crate::new_segment_id(); - let entry = LogEntry::AnnotatedSegmentStart { - ts: segment_log::now_millis(), - session_id: source_session_id, - system_prompt: state.system_prompt.map(String::from), - config: state.config.clone(), - history: state.history.to_vec(), - forked_from: None, - compacted_from: Some(SegmentOrigin { + let entries = seed_entries( + segment_log::now_millis(), + source_session_id, + state, + None, + Some(SegmentOrigin { segment_id: source_segment_id, at_turn_index: source_turn_count, }), - }; - store.append(source_session_id, segment_id, &entry)?; + ); + store.create_segment(source_session_id, segment_id, &entries)?; Ok(segment_id) } @@ -152,21 +169,19 @@ pub fn ensure_head_or_fork( } let source_segment_id = *segment_id; let fork_id = crate::new_segment_id(); - let entry = LogEntry::AnnotatedSegmentStart { - ts: segment_log::now_millis(), + let entries = seed_entries( + segment_log::now_millis(), session_id, - system_prompt: state.system_prompt.map(String::from), - config: state.config.clone(), - history: state.history.to_vec(), - forked_from: Some(SegmentOrigin { + state, + Some(SegmentOrigin { segment_id: source_segment_id, at_turn_index, }), - compacted_from: None, - }; - store.create_segment(session_id, fork_id, &[entry])?; + None, + ); + store.create_segment(session_id, fork_id, &entries)?; *segment_id = fork_id; - *entries_written = 1; + *entries_written = entries.len(); Ok(()) } @@ -425,16 +440,8 @@ pub fn fork( ) -> Result<(SessionId, SegmentId), StoreError> { let session_id = crate::new_session_id(); let fork_id = crate::new_segment_id(); - let entry = LogEntry::AnnotatedSegmentStart { - ts: segment_log::now_millis(), - session_id, - system_prompt: state.system_prompt.map(String::from), - config: state.config.clone(), - history: state.history.to_vec(), - forked_from: None, - compacted_from: None, - }; - store.create_segment(session_id, fork_id, &[entry])?; + let entries = seed_entries(segment_log::now_millis(), session_id, state, None, None); + store.create_segment(session_id, fork_id, &entries)?; store.copy_committed_uploaded_files(source_session_id, session_id)?; Ok((session_id, fork_id)) } diff --git a/crates/session-store/src/uploaded_file.rs b/crates/session-store/src/uploaded_file.rs index 291966bd..3910bd3a 100644 --- a/crates/session-store/src/uploaded_file.rs +++ b/crates/session-store/src/uploaded_file.rs @@ -228,9 +228,6 @@ pub(crate) fn write_uploaded_file( FileExt::lock_exclusive(&aggregate_lock)?; let (paste_bytes, _) = crate::paste_artifact::stored_paste_usage(dir)?; let (file_bytes, file_count) = stored_uploaded_file_usage(dir)?; - if file_count >= DEFAULT_MAX_SESSION_UPLOADED_FILES { - return Err(StoreError::ArtifactQuotaExceeded); - } let normalized_name = normalized_file_name(file_name); for entry in fs::read_dir(dir)? { let path = entry?.path(); @@ -242,10 +239,12 @@ pub(crate) fn write_uploaded_file( continue; } let stored: StoredUploadedFile = serde_json::from_slice(&fs::read(&path)?)?; - if stored.source_entry_id.is_none() - && normalized_file_name(&stored.file_name) == normalized_name - { - if stored.media_type == media_type + let same_context = context.is_some() && stored.upload_context.as_ref() == context; + let same_uncommitted_name = stored.source_entry_id.is_none() + && normalized_file_name(&stored.file_name) == normalized_name; + if same_context || same_uncommitted_name { + if stored.file_name == file_name + && stored.media_type == media_type && stored.byte_len == byte_len && stored.sha256 == sha256 && stored.upload_context.as_ref() == context @@ -270,6 +269,9 @@ pub(crate) fn write_uploaded_file( return Err(StoreError::InvalidUploadedFileName); } } + if file_count >= DEFAULT_MAX_SESSION_UPLOADED_FILES { + return Err(StoreError::ArtifactQuotaExceeded); + } if paste_bytes .checked_add(file_bytes) .and_then(|total| total.checked_add(byte_len)) diff --git a/crates/session-store/tests/session_test.rs b/crates/session-store/tests/session_test.rs index a8bddf1c..00f96e9d 100644 --- a/crates/session-store/tests/session_test.rs +++ b/crates/session-store/tests/session_test.rs @@ -237,6 +237,7 @@ async fn session_run_logs_entries() { system_prompt: worker.get_system_prompt(), config: worker.request_config(), history: annotated(&worker.history()), + user_segments: Vec::new(), }, ) .unwrap(); @@ -285,6 +286,7 @@ async fn session_restore_round_trip() { system_prompt: worker.get_system_prompt(), config: worker.request_config(), history: annotated(&worker.history()), + user_segments: Vec::new(), }, ) .unwrap(); @@ -324,6 +326,7 @@ async fn session_run_with_tool_call() { system_prompt: worker.get_system_prompt(), config: worker.request_config(), history: annotated(&worker.history()), + user_segments: Vec::new(), }, ) .unwrap(); @@ -359,6 +362,7 @@ async fn session_resume_after_pause() { system_prompt: worker.get_system_prompt(), config: worker.request_config(), history: annotated(&worker.history()), + user_segments: Vec::new(), }, ) .unwrap(); @@ -398,6 +402,7 @@ async fn session_fork_creates_new_session() { system_prompt: worker.get_system_prompt(), config: worker.request_config(), history: annotated(&worker.history()), + user_segments: Vec::new(), }, ) .unwrap(); @@ -405,6 +410,9 @@ async fn session_fork_creates_new_session() { let (worker, _) = run_and_persist(worker, &store, sid, segid, "Hello").await; let original_history_len = worker.history().len(); + let source_user_segments = session_store::restore(&store, sid, segid) + .unwrap() + .user_segments; let (fork_sid, fork_segid) = session_store::fork( &store, sid, @@ -412,22 +420,28 @@ async fn session_fork_creates_new_session() { system_prompt: worker.get_system_prompt(), config: worker.request_config(), history: annotated(&worker.history()), + user_segments: source_user_segments.clone(), }, ) .unwrap(); assert_ne!(fork_sid, sid, "`fork` mints a fresh Session"); - // Fork should have a SegmentStart with the current history + // Fork should have an annotated seed and typed input checkpoint. let fork_entries = store.read_all(fork_sid, fork_segid).unwrap(); - assert_eq!(fork_entries.len(), 1); + assert_eq!(fork_entries.len(), 2); assert!(matches!( &fork_entries[0], LogEntry::AnnotatedSegmentStart { .. } )); + assert!(matches!( + &fork_entries[1], + LogEntry::InputSegmentsCheckpoint { .. } + )); let fork_state = collect_state(&fork_entries); assert_eq!(fork_state.session_id, Some(fork_sid)); assert_eq!(fork_state.history.len(), original_history_len); + assert_eq!(fork_state.user_segments, source_user_segments); assert_eq!(fork_state.system_prompt.as_deref(), Some("System prompt")); } @@ -443,6 +457,7 @@ async fn session_fork_at_truncates_within_session() { system_prompt: worker.get_system_prompt(), config: worker.request_config(), history: annotated(&worker.history()), + user_segments: Vec::new(), }, ) .unwrap(); @@ -508,6 +523,7 @@ fn rewound_fork_preserves_uploaded_file_segments_in_snapshot() { system_prompt: Some("System prompt"), config: &config, history: Vec::new(), + user_segments: Vec::new(), }, ) .unwrap(); @@ -548,6 +564,31 @@ fn rewound_fork_preserves_uploaded_file_segments_in_snapshot() { SessionSnapshotEntryData::UserInput { segments: restored } if restored == &segments ))); + + let fork_state = collect_state(&fork_entries); + let (copied_session_id, copied_segment_id) = session_store::fork( + &store, + sid, + SegmentStartState { + system_prompt: fork_state.system_prompt.as_deref(), + config: &fork_state.config, + history: fork_state.annotated_history.clone(), + user_segments: fork_state.user_segments.clone(), + }, + ) + .unwrap(); + let copied_entries = store + .read_all(copied_session_id, copied_segment_id) + .unwrap(); + let copied_snapshot = session_store::public_snapshot::project_session_snapshot( + copied_session_id, + &copied_entries, + ); + assert!(copied_snapshot.entries.iter().any(|entry| matches!( + &entry.data, + SessionSnapshotEntryData::UserInput { segments: restored } + if restored == &segments + ))); } #[tokio::test] @@ -562,6 +603,7 @@ async fn session_config_changed_logged() { system_prompt: worker.get_system_prompt(), config: worker.request_config(), history: annotated(&worker.history()), + user_segments: Vec::new(), }, ) .unwrap(); @@ -595,6 +637,7 @@ async fn session_auto_forks_on_conflict() { system_prompt: worker_a.get_system_prompt(), config: worker_a.request_config(), history: annotated(&worker_a.history()), + user_segments: Vec::new(), }, ) .unwrap(); @@ -623,6 +666,7 @@ async fn session_auto_forks_on_conflict() { system_prompt: worker_a.get_system_prompt(), config: worker_a.request_config(), history: annotated(&worker_a.history()), + user_segments: Vec::new(), }, ) .unwrap(); @@ -682,6 +726,7 @@ async fn nested_past_fork_leaves_ancestors_immutable() { system_prompt: worker.get_system_prompt(), config: worker.request_config(), history: annotated(&worker.history()), + user_segments: Vec::new(), }, ) .unwrap();