feat(session): expose canonical public snapshots
This commit is contained in:
@@ -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(),
|
||||
|
||||
@@ -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 },
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 })
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user