feat(session): expose canonical public snapshots

This commit is contained in:
2026-08-30 00:41:51 +09:00
parent 40ac83e632
commit 88be87e03e
26 changed files with 1512 additions and 413 deletions
+14 -27
View File
@@ -84,10 +84,7 @@ impl WorkerHandle {
(entries, entry_rx, in_flight)
};
let event = Event::Snapshot {
entries: entries
.into_iter()
.map(|entry| serde_json::to_value(entry).expect("log entry serializes"))
.collect(),
session: session_store::public_snapshot::project_current_session_snapshot(&entries),
greeting: self.shared_state.greeting.clone(),
status: self.shared_state.get_status(),
in_flight,
@@ -1874,28 +1871,16 @@ where
St: Store,
{
match worker.rewind_to(target, expected_head_entries) {
Ok(applied) => match applied
.entries
.into_iter()
.map(serde_json::to_value)
.collect::<Result<Vec<_>, _>>()
{
Ok(entries) => {
let _ = event_tx.send(Event::RewindApplied {
entries,
input: applied.input,
summary: applied.summary,
});
true
}
Err(error) => {
let _ = event_tx.send(Event::Error {
code: ErrorCode::Internal,
message: format!("failed to encode rewind snapshot: {error}"),
});
false
}
},
Ok(applied) => {
let session =
session_store::public_snapshot::project_current_session_snapshot(&applied.entries);
let _ = event_tx.send(Event::RewindApplied {
session,
input: applied.input,
summary: applied.summary,
});
true
}
Err(err) => {
let _ = event_tx.send(Event::Error {
code: ErrorCode::InvalidRequest,
@@ -2101,7 +2086,9 @@ mod tests {
let mut writer = JsonLineWriter::new(w);
writer
.write(&Event::Snapshot {
entries: Vec::new(),
session: protocol::SessionSnapshot {
entries: Vec::new(),
},
greeting: protocol::Greeting {
worker_name: "parent".into(),
cwd: "/tmp".into(),
+18 -6
View File
@@ -1481,7 +1481,9 @@ mod tests {
let mut writer = JsonLineWriter::new(stream);
writer
.write(&Event::Snapshot {
entries: Vec::new(),
session: protocol::SessionSnapshot {
entries: Vec::new(),
},
greeting: protocol::Greeting {
worker_name: "target".into(),
cwd: "/tmp".into(),
@@ -1514,7 +1516,9 @@ mod tests {
.unwrap();
writer
.write(&Event::Snapshot {
entries: Vec::new(),
session: protocol::SessionSnapshot {
entries: Vec::new(),
},
greeting: protocol::Greeting {
worker_name: "target".into(),
cwd: "/tmp".into(),
@@ -1603,7 +1607,9 @@ mod tests {
let mut writer = JsonLineWriter::new(stream);
writer
.write(&Event::Snapshot {
entries: Vec::new(),
session: protocol::SessionSnapshot {
entries: Vec::new(),
},
greeting: protocol::Greeting {
worker_name: "target".into(),
cwd: "/tmp".into(),
@@ -1627,7 +1633,9 @@ mod tests {
let mut writer = JsonLineWriter::new(writer_half);
writer
.write(&Event::Snapshot {
entries: Vec::new(),
session: protocol::SessionSnapshot {
entries: Vec::new(),
},
greeting: protocol::Greeting {
worker_name: "target".into(),
cwd: "/tmp".into(),
@@ -1729,7 +1737,9 @@ mod tests {
.unwrap();
writer
.write(&Event::Snapshot {
entries: Vec::new(),
session: protocol::SessionSnapshot {
entries: Vec::new(),
},
greeting: protocol::Greeting {
worker_name: "alerted".into(),
cwd: "/tmp".into(),
@@ -1779,7 +1789,9 @@ mod tests {
let mut writer = JsonLineWriter::new(stream);
let _ = writer
.write(&Event::Snapshot {
entries: Vec::new(),
session: protocol::SessionSnapshot {
entries: Vec::new(),
},
greeting: protocol::Greeting {
worker_name: "child-live".into(),
cwd: "/tmp".into(),
@@ -6,7 +6,7 @@ use agen::tool::{Tool, ToolDefinition, ToolError, ToolMeta, ToolOutput};
use async_trait::async_trait;
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use session_store::{LogEntry, collect_state};
use session_store::LogEntry;
use super::manage_worker::{WORKER_CONTROL_SERVICE_ID, WorkerControlService};
use crate::feature::{
@@ -61,7 +61,7 @@ pub struct WorkerObservationSubject {
#[derive(Debug, Clone)]
pub struct WorkerSessionCapture {
pub segment_id: String,
pub entries: Vec<agen::HistoryEntry<crate::SessionHistoryMetadata>>,
pub session: protocol::SessionSnapshot,
}
impl WorkerSessionCapture {
@@ -69,17 +69,9 @@ impl WorkerSessionCapture {
segment_id: impl Into<String>,
log_entries: &[LogEntry],
) -> Result<Self, String> {
let segment_id = segment_id.into();
let state = collect_state(log_entries);
let parsed_segment_id = segment_id.parse().unwrap_or_default();
let entries = crate::session_history::restore_history_entries(
state.session_id.unwrap_or_default(),
parsed_segment_id,
log_entries,
)?;
Ok(Self {
segment_id,
entries,
segment_id: segment_id.into(),
session: session_store::public_snapshot::project_current_session_snapshot(log_entries),
})
}
}
@@ -115,7 +107,7 @@ struct WorkspaceWorkerObservationListResponse {
#[derive(Debug, Deserialize)]
struct WorkspaceWorkerObservationCaptureResponse {
segment_id: String,
entries: Vec<serde_json::Value>,
session: protocol::SessionSnapshot,
}
pub struct WorkspaceClientWorkerObservationProvider {
@@ -173,26 +165,9 @@ impl WorkerObservationProvider for WorkspaceClientWorkerObservationProvider {
let body = workspace_response_body(response)?;
let response = serde_json::from_str::<WorkspaceWorkerObservationCaptureResponse>(&body)
.map_err(|error| WorkerObservationError::Unavailable(error.to_string()))?;
let entries = response
.entries
.into_iter()
.map(|entry| {
serde_json::from_value(entry)
.map_err(|error| WorkerObservationError::Unavailable(error.to_string()))
})
.collect::<Result<Vec<session_store::LogEntry>, _>>()?;
let state = collect_state(&entries);
let segment_id = response.segment_id;
let parsed_segment_id = segment_id.parse().unwrap_or_default();
let typed_entries = crate::session_history::restore_history_entries(
state.session_id.unwrap_or_default(),
parsed_segment_id,
&entries,
)
.map_err(WorkerObservationError::Unavailable)?;
Ok(WorkerSessionCapture {
segment_id,
entries: typed_entries,
segment_id: response.segment_id,
session: response.session,
})
}
}
@@ -420,16 +395,9 @@ impl WorkerObservationProvider for SpawnedSubWorkerObservationProvider {
.get_internal(name)
.ok_or(WorkerObservationError::NotFound)?;
let entries = record.session.entries();
let state = collect_state(&entries);
let typed_entries = crate::session_history::restore_history_entries(
state.session_id.unwrap_or_default(),
Default::default(),
&entries,
)
.map_err(WorkerObservationError::Unavailable)?;
Ok(WorkerSessionCapture {
segment_id: format!("subworker:{name}"),
entries: typed_entries,
session: session_store::public_snapshot::project_current_session_snapshot(&entries),
})
}
}
@@ -699,9 +667,9 @@ async fn latest_view(
.capture_worker_session(subject)
.await
.map_err(tool_error)?;
Ok(SessionCapture::from_history_entries(
Ok(SessionCapture::from_session_snapshot(
capture.segment_id,
capture.entries,
capture.session,
))
}
@@ -799,16 +767,42 @@ mod tests {
.clone()
.into_iter()
.enumerate()
.map(|(index, item)| {
let mut metadata = crate::SessionHistoryMetadata::legacy_unknown();
metadata.entry_id =
session_store::LoggedSessionHistoryEntryId(format!("fake-{index:08}"));
agen::HistoryEntry::new(item, metadata)
.filter_map(|(index, item)| {
let data = match item {
Item::Message { role, content, .. } => {
let role = match role {
Role::User => protocol::SessionMessageRole::User,
Role::Assistant => protocol::SessionMessageRole::Assistant,
Role::System => return None,
};
protocol::SessionSnapshotEntryData::Message {
role,
content: content
.into_iter()
.map(|part| match part {
agen::ContentPart::Text { text } => {
protocol::SessionContentPart::Text { text }
}
agen::ContentPart::Refusal { refusal } => {
protocol::SessionContentPart::Refusal { refusal }
}
})
.collect(),
}
}
_ => return None,
};
Some(protocol::SessionSnapshotEntry {
entry_id: format!("fake-{index:08}"),
provenance: protocol::SessionEntryProvenance::LegacyUnknown,
derived_from: Vec::new(),
data,
})
})
.collect();
Ok(WorkerSessionCapture {
segment_id: "segment".to_string(),
entries,
session: protocol::SessionSnapshot { entries },
})
}
}
+2 -2
View File
@@ -335,7 +335,7 @@ enum InternalWorkerSessionCommand {
/// task; protocol access is consumed only by the owning parent registry.
#[derive(Debug, Clone)]
pub(crate) struct InternalWorkerSessionSnapshot {
pub entries: Vec<LogEntry>,
pub session: protocol::SessionSnapshot,
pub status: WorkerStatus,
pub error: Option<String>,
pub in_flight: InFlightSnapshot,
@@ -402,7 +402,7 @@ impl InternalWorkerSessionHandle {
(entries, snapshot_from_guard(&guard))
};
InternalWorkerSessionSnapshot {
entries,
session: session_store::public_snapshot::project_current_session_snapshot(&entries),
status: match self.status() {
InternalWorkerSessionStatus::Running => WorkerStatus::Running,
InternalWorkerSessionStatus::Paused => WorkerStatus::Paused,
+3 -2
View File
@@ -30,8 +30,9 @@ pub fn subscribe_worker_protocol_session(handle: &WorkerHandle) -> WorkerProtoco
pub fn live_log_entry_event(entry: LogEntry) -> Option<Event> {
match entry {
entry @ (LogEntry::SegmentStart { .. } | LogEntry::AnnotatedSegmentStart { .. }) => {
let value = serde_json::to_value(&entry).expect("LogEntry is Serialize");
Some(Event::SegmentRotated { entry: value })
let session =
session_store::public_snapshot::project_current_session_snapshot(&[entry]);
Some(Event::SegmentRotated { session })
}
LogEntry::UserInput { segments, .. } | LogEntry::AnnotatedUserInput { segments, .. } => {
Some(Event::UserMessage { segments })
+104
View File
@@ -8,6 +8,10 @@ use std::sync::Arc;
use crate::session_history::{SessionHistoryMetadata, WorkerHistoryProvenance};
use agen::{HistoryEntry, Item, Role};
use protocol::{
SessionContentPart, SessionEntryProvenance, SessionMessageRole, SessionSnapshot,
SessionSnapshotEntryData,
};
use serde::{Deserialize, Serialize};
const DEFAULT_SEARCH_LIMIT: usize = 20;
@@ -225,6 +229,66 @@ pub(crate) struct SessionCapture {
}
impl SessionCapture {
pub(crate) fn from_session_snapshot(
segment_id: impl Into<String>,
snapshot: SessionSnapshot,
) -> Self {
let entries = snapshot
.entries
.into_iter()
.filter_map(|entry| {
let item = match entry.data {
SessionSnapshotEntryData::UserInput { segments } => {
Item::user_message(protocol::Segment::flatten_to_text(&segments))
}
SessionSnapshotEntryData::Message { role, content } => {
let role = match role {
SessionMessageRole::User => Role::User,
SessionMessageRole::Assistant => Role::Assistant,
};
Item::Message {
id: None,
role,
content: content
.into_iter()
.map(|part| match part {
SessionContentPart::Text { text } => {
agen::ContentPart::Text { text }
}
SessionContentPart::Refusal { refusal } => {
agen::ContentPart::Refusal { refusal }
}
})
.collect(),
status: None,
}
}
SessionSnapshotEntryData::ToolCall {
call_id,
name,
arguments,
} => Item::tool_call(call_id, name, arguments),
SessionSnapshotEntryData::ToolResult {
call_id,
summary,
content,
is_error,
attachments: _,
} => Item::tool_result_item(call_id, summary, content, is_error),
// Observation deliberately excludes system items and
// controller errors from model-visible session evidence.
SessionSnapshotEntryData::SystemItem { .. }
| SessionSnapshotEntryData::RunError { .. } => return None,
};
Some(HistoryEntry::new(
item,
public_snapshot_metadata(entry.entry_id, entry.provenance),
))
})
.collect();
Self::from_history_entries(segment_id, entries)
}
pub(crate) fn new(segment_id: impl Into<String>, items: Vec<Item>) -> Self {
let entries = items
.into_iter()
@@ -543,6 +607,46 @@ impl SessionCapture {
}
}
fn public_snapshot_metadata(
entry_id: String,
provenance: SessionEntryProvenance,
) -> SessionHistoryMetadata {
let worker = session_store::LoggedWorkerSubject {
workspace_id: None,
runtime_id: None,
worker_id: "public-session-snapshot".to_owned(),
};
let origin = match provenance {
SessionEntryProvenance::HumanInput => WorkerHistoryProvenance::HumanInput {
account_id: "public-session-snapshot".to_owned(),
},
SessionEntryProvenance::WorkerInput => WorkerHistoryProvenance::WorkerInput {
actor: worker.clone(),
},
SessionEntryProvenance::FlowInstruction => WorkerHistoryProvenance::FlowInstruction {
selector: "public-session-snapshot".to_owned(),
definition_id: "public-session-snapshot".to_owned(),
definition_revision: 0,
instance_id: "public-session-snapshot".to_owned(),
state_id: "public-session-snapshot".to_owned(),
},
SessionEntryProvenance::BackendInstruction => {
WorkerHistoryProvenance::BackendInstruction { operation_id: None }
}
SessionEntryProvenance::ModelOutput => WorkerHistoryProvenance::ModelOutput {
worker: worker.clone(),
},
SessionEntryProvenance::ToolOutput => WorkerHistoryProvenance::ToolOutput { worker },
SessionEntryProvenance::DerivedSummary => WorkerHistoryProvenance::DerivedSummary,
SessionEntryProvenance::LegacyUnknown => WorkerHistoryProvenance::LegacyUnknown,
};
SessionHistoryMetadata {
entry_id: session_store::LoggedSessionHistoryEntryId(entry_id),
origin,
derivation: None,
}
}
fn message_reference_kind(
origin: &WorkerHistoryProvenance,
provider_role: &Role,
+14 -1
View File
@@ -278,7 +278,20 @@ mod tests {
fn snapshot(entries: Vec<serde_json::Value>) -> Event {
Event::Snapshot {
entries,
session: protocol::SessionSnapshot {
entries: entries
.into_iter()
.enumerate()
.map(|(index, value)| protocol::SessionSnapshotEntry {
entry_id: format!("test-{index}"),
provenance: protocol::SessionEntryProvenance::LegacyUnknown,
derived_from: Vec::new(),
data: protocol::SessionSnapshotEntryData::RunError {
message: value.to_string(),
},
})
.collect(),
},
greeting: Greeting {
worker_name: "server".into(),
cwd: "/tmp".into(),
+2 -6
View File
@@ -806,11 +806,7 @@ fn internal_worker_snapshot(
InternalWorkerSnapshot {
worker,
revision,
entries: snapshot
.entries
.into_iter()
.filter_map(|entry| serde_json::to_value(entry).ok())
.collect(),
session: snapshot.session,
status: snapshot.status,
error: snapshot.error,
in_flight: snapshot.in_flight,
@@ -1045,7 +1041,7 @@ mod tests {
let snapshots = registry.internal_worker_snapshots();
assert_eq!(snapshots.len(), 1);
assert_eq!(snapshots[0].revision, 2);
assert_eq!(snapshots[0].entries.len(), 1);
assert_eq!(snapshots[0].session.entries.len(), 1);
record.session.emit_test_text_delta("partial");
let streamed = tokio::time::timeout(Duration::from_secs(1), parent_rx.recv())
+6 -6
View File
@@ -960,9 +960,7 @@ mod tests {
use crate::WorkspaceId;
use agen::llm_client::event::{Event as LlmEvent, ResponseStatus, StatusEvent};
use agen::llm_client::types::ContentPart;
use agen::llm_client::{ClientError, LlmClient, Request};
use agen::{Item, Role};
use async_trait::async_trait;
use futures::Stream;
use manifest::{AuthRef, ModelManifest, SchemeKind, WorkerManifest};
@@ -1252,9 +1250,11 @@ extract_threshold = 4000
)
.await
.unwrap();
assert!(first_capture.entries.iter().map(|entry| &entry.item).any(|item| {
matches!(item, Item::Message { role: Role::Assistant, content, .. } if content.iter().any(|part| matches!(part, ContentPart::Text { text } if text.contains("reviewed"))))
}));
assert!(
serde_json::to_string(&first_capture.session)
.unwrap()
.contains("reviewed")
);
let send = (crate::spawn::comm_tools::sub_worker_send_tool(registry.clone()))().1;
send.execute(
@@ -1274,7 +1274,7 @@ extract_threshold = 4000
)
.await
.unwrap();
assert!(latest_capture.entries.len() > first_capture.entries.len());
assert!(latest_capture.session.entries.len() > first_capture.session.entries.len());
fail_requests.store(true, Ordering::SeqCst);
send.execute(
+22 -27
View File
@@ -839,26 +839,26 @@ async fn snapshot_includes_user_input_for_in_flight_turn() {
loop {
let event = reader.next::<Event>().await.unwrap().unwrap();
match event {
Event::Snapshot { entries, .. } => {
// Walk the entries, find a `LogEntry::UserInput` and
// confirm its segments flatten to our submitted text.
let mut found = false;
for value in &entries {
let entry: session_store::LogEntry =
serde_json::from_value(value.clone()).expect("LogEntry deserialise");
if let session_store::LogEntry::UserInput { segments, .. }
| session_store::LogEntry::AnnotatedUserInput { segments, .. } = entry
{
let text = protocol::Segment::flatten_to_text(&segments);
if text == "hello in-flight" {
found = true;
break;
}
Event::Snapshot { session, .. } => {
let found = session.entries.iter().any(|entry| match &entry.data {
protocol::SessionSnapshotEntryData::UserInput { segments } => {
protocol::Segment::flatten_to_text(segments) == "hello in-flight"
}
}
protocol::SessionSnapshotEntryData::Message {
role: protocol::SessionMessageRole::User,
content,
} => content.iter().any(|part| {
matches!(
part,
protocol::SessionContentPart::Text { text }
if text == "hello in-flight"
)
}),
_ => false,
});
assert!(
found,
"snapshot must carry the in-flight UserInput entry: {entries:?}"
"snapshot must carry the in-flight UserInput entry: {session:?}"
);
return;
}
@@ -2410,17 +2410,12 @@ async fn snapshot_contains_user_input(handle: &WorkerHandle, needle: &str) -> bo
loop {
let event = reader.next::<Event>().await.unwrap().unwrap();
match event {
Event::Snapshot { entries, .. } => {
return entries.into_iter().any(|value| {
let entry: session_store::LogEntry =
serde_json::from_value(value).expect("LogEntry deserialise");
match entry {
session_store::LogEntry::UserInput { segments, .. }
| session_store::LogEntry::AnnotatedUserInput { segments, .. } => {
protocol::Segment::flatten_to_text(&segments).contains(needle)
}
_ => false,
Event::Snapshot { session, .. } => {
return session.entries.into_iter().any(|entry| match entry.data {
protocol::SessionSnapshotEntryData::UserInput { segments } => {
protocol::Segment::flatten_to_text(&segments).contains(needle)
}
_ => false,
});
}
Event::Alert(_) => continue,
@@ -203,12 +203,12 @@ async fn session_start_state_captures_rendered_prompt() {
.unwrap();
let first = entries.first().expect("at least one entry");
match first {
LogEntry::SegmentStart { system_prompt, .. } => {
LogEntry::AnnotatedSegmentStart { system_prompt, .. } => {
let sp = system_prompt.as_deref().expect("system prompt set");
assert!(sp.starts_with("hello"));
assert!(sp.contains(&pwd.display().to_string()));
}
other => panic!("expected SegmentStart as first entry, got {other:?}"),
other => panic!("expected AnnotatedSegmentStart as first entry, got {other:?}"),
}
}