fix: preserve attachment replay evidence
This commit is contained in:
@@ -764,6 +764,62 @@ mod tests {
|
|||||||
.unwrap()
|
.unwrap()
|
||||||
.contains("account-1")
|
.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]
|
#[test]
|
||||||
|
|||||||
@@ -27,6 +27,7 @@
|
|||||||
//! system_prompt: None,
|
//! system_prompt: None,
|
||||||
//! config: &config,
|
//! config: &config,
|
||||||
//! history: Vec::new(),
|
//! history: Vec::new(),
|
||||||
|
//! user_segments: Vec::new(),
|
||||||
//! })?;
|
//! })?;
|
||||||
//! ```
|
//! ```
|
||||||
|
|
||||||
|
|||||||
@@ -17,6 +17,33 @@ pub struct SegmentStartState<'a> {
|
|||||||
pub system_prompt: Option<&'a str>,
|
pub system_prompt: Option<&'a str>,
|
||||||
pub config: &'a RequestConfig,
|
pub config: &'a RequestConfig,
|
||||||
pub history: Vec<LoggedHistoryEntry>,
|
pub history: Vec<LoggedHistoryEntry>,
|
||||||
|
pub user_segments: Vec<Vec<Segment>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
fn seed_entries(
|
||||||
|
ts: u64,
|
||||||
|
session_id: SessionId,
|
||||||
|
state: SegmentStartState<'_>,
|
||||||
|
forked_from: Option<SegmentOrigin>,
|
||||||
|
compacted_from: Option<SegmentOrigin>,
|
||||||
|
) -> Vec<LogEntry> {
|
||||||
|
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
|
/// Create a new session + initial segment, writing the initial
|
||||||
@@ -42,16 +69,8 @@ pub fn create_segment_with_ids(
|
|||||||
segment_id: SegmentId,
|
segment_id: SegmentId,
|
||||||
state: SegmentStartState<'_>,
|
state: SegmentStartState<'_>,
|
||||||
) -> Result<(), StoreError> {
|
) -> Result<(), StoreError> {
|
||||||
let entry = LogEntry::AnnotatedSegmentStart {
|
let entries = seed_entries(segment_log::now_millis(), session_id, state, None, None);
|
||||||
ts: segment_log::now_millis(),
|
store.create_segment(session_id, segment_id, &entries)
|
||||||
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)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Create a compacted segment from an existing one. Inherits the source's
|
/// 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,
|
source_turn_count: usize,
|
||||||
) -> Result<SegmentId, StoreError> {
|
) -> Result<SegmentId, StoreError> {
|
||||||
let segment_id = crate::new_segment_id();
|
let segment_id = crate::new_segment_id();
|
||||||
let entry = LogEntry::AnnotatedSegmentStart {
|
let entries = seed_entries(
|
||||||
ts: segment_log::now_millis(),
|
segment_log::now_millis(),
|
||||||
session_id: source_session_id,
|
source_session_id,
|
||||||
system_prompt: state.system_prompt.map(String::from),
|
state,
|
||||||
config: state.config.clone(),
|
None,
|
||||||
history: state.history.to_vec(),
|
Some(SegmentOrigin {
|
||||||
forked_from: None,
|
|
||||||
compacted_from: Some(SegmentOrigin {
|
|
||||||
segment_id: source_segment_id,
|
segment_id: source_segment_id,
|
||||||
at_turn_index: source_turn_count,
|
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)
|
Ok(segment_id)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -152,21 +169,19 @@ pub fn ensure_head_or_fork(
|
|||||||
}
|
}
|
||||||
let source_segment_id = *segment_id;
|
let source_segment_id = *segment_id;
|
||||||
let fork_id = crate::new_segment_id();
|
let fork_id = crate::new_segment_id();
|
||||||
let entry = LogEntry::AnnotatedSegmentStart {
|
let entries = seed_entries(
|
||||||
ts: segment_log::now_millis(),
|
segment_log::now_millis(),
|
||||||
session_id,
|
session_id,
|
||||||
system_prompt: state.system_prompt.map(String::from),
|
state,
|
||||||
config: state.config.clone(),
|
Some(SegmentOrigin {
|
||||||
history: state.history.to_vec(),
|
|
||||||
forked_from: Some(SegmentOrigin {
|
|
||||||
segment_id: source_segment_id,
|
segment_id: source_segment_id,
|
||||||
at_turn_index,
|
at_turn_index,
|
||||||
}),
|
}),
|
||||||
compacted_from: None,
|
None,
|
||||||
};
|
);
|
||||||
store.create_segment(session_id, fork_id, &[entry])?;
|
store.create_segment(session_id, fork_id, &entries)?;
|
||||||
*segment_id = fork_id;
|
*segment_id = fork_id;
|
||||||
*entries_written = 1;
|
*entries_written = entries.len();
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -425,16 +440,8 @@ pub fn fork(
|
|||||||
) -> Result<(SessionId, SegmentId), StoreError> {
|
) -> Result<(SessionId, SegmentId), StoreError> {
|
||||||
let session_id = crate::new_session_id();
|
let session_id = crate::new_session_id();
|
||||||
let fork_id = crate::new_segment_id();
|
let fork_id = crate::new_segment_id();
|
||||||
let entry = LogEntry::AnnotatedSegmentStart {
|
let entries = seed_entries(segment_log::now_millis(), session_id, state, None, None);
|
||||||
ts: segment_log::now_millis(),
|
store.create_segment(session_id, fork_id, &entries)?;
|
||||||
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])?;
|
|
||||||
store.copy_committed_uploaded_files(source_session_id, session_id)?;
|
store.copy_committed_uploaded_files(source_session_id, session_id)?;
|
||||||
Ok((session_id, fork_id))
|
Ok((session_id, fork_id))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -228,9 +228,6 @@ pub(crate) fn write_uploaded_file(
|
|||||||
FileExt::lock_exclusive(&aggregate_lock)?;
|
FileExt::lock_exclusive(&aggregate_lock)?;
|
||||||
let (paste_bytes, _) = crate::paste_artifact::stored_paste_usage(dir)?;
|
let (paste_bytes, _) = crate::paste_artifact::stored_paste_usage(dir)?;
|
||||||
let (file_bytes, file_count) = stored_uploaded_file_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);
|
let normalized_name = normalized_file_name(file_name);
|
||||||
for entry in fs::read_dir(dir)? {
|
for entry in fs::read_dir(dir)? {
|
||||||
let path = entry?.path();
|
let path = entry?.path();
|
||||||
@@ -242,10 +239,12 @@ pub(crate) fn write_uploaded_file(
|
|||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
let stored: StoredUploadedFile = serde_json::from_slice(&fs::read(&path)?)?;
|
let stored: StoredUploadedFile = serde_json::from_slice(&fs::read(&path)?)?;
|
||||||
if stored.source_entry_id.is_none()
|
let same_context = context.is_some() && stored.upload_context.as_ref() == context;
|
||||||
&& normalized_file_name(&stored.file_name) == normalized_name
|
let same_uncommitted_name = stored.source_entry_id.is_none()
|
||||||
{
|
&& normalized_file_name(&stored.file_name) == normalized_name;
|
||||||
if stored.media_type == media_type
|
if same_context || same_uncommitted_name {
|
||||||
|
if stored.file_name == file_name
|
||||||
|
&& stored.media_type == media_type
|
||||||
&& stored.byte_len == byte_len
|
&& stored.byte_len == byte_len
|
||||||
&& stored.sha256 == sha256
|
&& stored.sha256 == sha256
|
||||||
&& stored.upload_context.as_ref() == context
|
&& stored.upload_context.as_ref() == context
|
||||||
@@ -270,6 +269,9 @@ pub(crate) fn write_uploaded_file(
|
|||||||
return Err(StoreError::InvalidUploadedFileName);
|
return Err(StoreError::InvalidUploadedFileName);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
if file_count >= DEFAULT_MAX_SESSION_UPLOADED_FILES {
|
||||||
|
return Err(StoreError::ArtifactQuotaExceeded);
|
||||||
|
}
|
||||||
if paste_bytes
|
if paste_bytes
|
||||||
.checked_add(file_bytes)
|
.checked_add(file_bytes)
|
||||||
.and_then(|total| total.checked_add(byte_len))
|
.and_then(|total| total.checked_add(byte_len))
|
||||||
|
|||||||
@@ -237,6 +237,7 @@ async fn session_run_logs_entries() {
|
|||||||
system_prompt: worker.get_system_prompt(),
|
system_prompt: worker.get_system_prompt(),
|
||||||
config: worker.request_config(),
|
config: worker.request_config(),
|
||||||
history: annotated(&worker.history()),
|
history: annotated(&worker.history()),
|
||||||
|
user_segments: Vec::new(),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
.unwrap();
|
.unwrap();
|
||||||
@@ -285,6 +286,7 @@ async fn session_restore_round_trip() {
|
|||||||
system_prompt: worker.get_system_prompt(),
|
system_prompt: worker.get_system_prompt(),
|
||||||
config: worker.request_config(),
|
config: worker.request_config(),
|
||||||
history: annotated(&worker.history()),
|
history: annotated(&worker.history()),
|
||||||
|
user_segments: Vec::new(),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
.unwrap();
|
.unwrap();
|
||||||
@@ -324,6 +326,7 @@ async fn session_run_with_tool_call() {
|
|||||||
system_prompt: worker.get_system_prompt(),
|
system_prompt: worker.get_system_prompt(),
|
||||||
config: worker.request_config(),
|
config: worker.request_config(),
|
||||||
history: annotated(&worker.history()),
|
history: annotated(&worker.history()),
|
||||||
|
user_segments: Vec::new(),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
.unwrap();
|
.unwrap();
|
||||||
@@ -359,6 +362,7 @@ async fn session_resume_after_pause() {
|
|||||||
system_prompt: worker.get_system_prompt(),
|
system_prompt: worker.get_system_prompt(),
|
||||||
config: worker.request_config(),
|
config: worker.request_config(),
|
||||||
history: annotated(&worker.history()),
|
history: annotated(&worker.history()),
|
||||||
|
user_segments: Vec::new(),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
.unwrap();
|
.unwrap();
|
||||||
@@ -398,6 +402,7 @@ async fn session_fork_creates_new_session() {
|
|||||||
system_prompt: worker.get_system_prompt(),
|
system_prompt: worker.get_system_prompt(),
|
||||||
config: worker.request_config(),
|
config: worker.request_config(),
|
||||||
history: annotated(&worker.history()),
|
history: annotated(&worker.history()),
|
||||||
|
user_segments: Vec::new(),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
.unwrap();
|
.unwrap();
|
||||||
@@ -405,6 +410,9 @@ async fn session_fork_creates_new_session() {
|
|||||||
let (worker, _) = run_and_persist(worker, &store, sid, segid, "Hello").await;
|
let (worker, _) = run_and_persist(worker, &store, sid, segid, "Hello").await;
|
||||||
|
|
||||||
let original_history_len = worker.history().len();
|
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(
|
let (fork_sid, fork_segid) = session_store::fork(
|
||||||
&store,
|
&store,
|
||||||
sid,
|
sid,
|
||||||
@@ -412,22 +420,28 @@ async fn session_fork_creates_new_session() {
|
|||||||
system_prompt: worker.get_system_prompt(),
|
system_prompt: worker.get_system_prompt(),
|
||||||
config: worker.request_config(),
|
config: worker.request_config(),
|
||||||
history: annotated(&worker.history()),
|
history: annotated(&worker.history()),
|
||||||
|
user_segments: source_user_segments.clone(),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
.unwrap();
|
.unwrap();
|
||||||
assert_ne!(fork_sid, sid, "`fork` mints a fresh Session");
|
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();
|
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!(
|
assert!(matches!(
|
||||||
&fork_entries[0],
|
&fork_entries[0],
|
||||||
LogEntry::AnnotatedSegmentStart { .. }
|
LogEntry::AnnotatedSegmentStart { .. }
|
||||||
));
|
));
|
||||||
|
assert!(matches!(
|
||||||
|
&fork_entries[1],
|
||||||
|
LogEntry::InputSegmentsCheckpoint { .. }
|
||||||
|
));
|
||||||
|
|
||||||
let fork_state = collect_state(&fork_entries);
|
let fork_state = collect_state(&fork_entries);
|
||||||
assert_eq!(fork_state.session_id, Some(fork_sid));
|
assert_eq!(fork_state.session_id, Some(fork_sid));
|
||||||
assert_eq!(fork_state.history.len(), original_history_len);
|
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"));
|
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(),
|
system_prompt: worker.get_system_prompt(),
|
||||||
config: worker.request_config(),
|
config: worker.request_config(),
|
||||||
history: annotated(&worker.history()),
|
history: annotated(&worker.history()),
|
||||||
|
user_segments: Vec::new(),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
.unwrap();
|
.unwrap();
|
||||||
@@ -508,6 +523,7 @@ fn rewound_fork_preserves_uploaded_file_segments_in_snapshot() {
|
|||||||
system_prompt: Some("System prompt"),
|
system_prompt: Some("System prompt"),
|
||||||
config: &config,
|
config: &config,
|
||||||
history: Vec::new(),
|
history: Vec::new(),
|
||||||
|
user_segments: Vec::new(),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
.unwrap();
|
.unwrap();
|
||||||
@@ -548,6 +564,31 @@ fn rewound_fork_preserves_uploaded_file_segments_in_snapshot() {
|
|||||||
SessionSnapshotEntryData::UserInput { segments: restored }
|
SessionSnapshotEntryData::UserInput { segments: restored }
|
||||||
if restored == &segments
|
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]
|
#[tokio::test]
|
||||||
@@ -562,6 +603,7 @@ async fn session_config_changed_logged() {
|
|||||||
system_prompt: worker.get_system_prompt(),
|
system_prompt: worker.get_system_prompt(),
|
||||||
config: worker.request_config(),
|
config: worker.request_config(),
|
||||||
history: annotated(&worker.history()),
|
history: annotated(&worker.history()),
|
||||||
|
user_segments: Vec::new(),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
.unwrap();
|
.unwrap();
|
||||||
@@ -595,6 +637,7 @@ async fn session_auto_forks_on_conflict() {
|
|||||||
system_prompt: worker_a.get_system_prompt(),
|
system_prompt: worker_a.get_system_prompt(),
|
||||||
config: worker_a.request_config(),
|
config: worker_a.request_config(),
|
||||||
history: annotated(&worker_a.history()),
|
history: annotated(&worker_a.history()),
|
||||||
|
user_segments: Vec::new(),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
.unwrap();
|
.unwrap();
|
||||||
@@ -623,6 +666,7 @@ async fn session_auto_forks_on_conflict() {
|
|||||||
system_prompt: worker_a.get_system_prompt(),
|
system_prompt: worker_a.get_system_prompt(),
|
||||||
config: worker_a.request_config(),
|
config: worker_a.request_config(),
|
||||||
history: annotated(&worker_a.history()),
|
history: annotated(&worker_a.history()),
|
||||||
|
user_segments: Vec::new(),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
.unwrap();
|
.unwrap();
|
||||||
@@ -682,6 +726,7 @@ async fn nested_past_fork_leaves_ancestors_immutable() {
|
|||||||
system_prompt: worker.get_system_prompt(),
|
system_prompt: worker.get_system_prompt(),
|
||||||
config: worker.request_config(),
|
config: worker.request_config(),
|
||||||
history: annotated(&worker.history()),
|
history: annotated(&worker.history()),
|
||||||
|
user_segments: Vec::new(),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
.unwrap();
|
.unwrap();
|
||||||
|
|||||||
Reference in New Issue
Block a user