session-log-segments実装

This commit is contained in:
2026-04-29 22:42:10 +09:00
parent 8a9e3b4fe3
commit e3b36371e9
14 changed files with 310 additions and 65 deletions
+1
View File
@@ -15,6 +15,7 @@ uuid = { version = "1", features = ["v7", "serde"] }
thiserror = "2.0"
sha2 = "0.11.0"
hex = "0.4.3"
protocol = { version = "0.1.0", path = "../protocol" }
[dev-dependencies]
tokio = { version = "1.49", features = ["macros", "rt-multi-thread", "fs", "io-util"] }
+1 -1
View File
@@ -39,7 +39,7 @@ pub use logged_item::{LoggedContentPart, LoggedItem, LoggedRole, from_logged, to
pub use session::{
SessionStartState, create_compacted_session, create_session, create_session_with_id,
ensure_head_or_fork, fork, fork_at, restore, save_config_changed, save_delta, save_extension,
save_run_completed, save_run_errored, save_turn_end, save_usage,
save_run_completed, save_run_errored, save_turn_end, save_usage, save_user_input,
};
pub use llm_worker::UsageRecord;
pub use llm_worker::llm_client::types::{ContentPart, Item, Role};
+31 -12
View File
@@ -11,6 +11,7 @@ use crate::store::{Store, StoreError};
use llm_worker::WorkerResult;
use llm_worker::llm_client::RequestConfig;
use llm_worker::llm_client::types::Item;
use protocol::Segment;
/// State snapshot for creating a SessionStart entry.
pub struct SessionStartState<'a> {
@@ -138,10 +139,37 @@ pub async fn ensure_head_or_fork(
Ok(())
}
/// Log a `UserInput` entry from the original typed `Vec<Segment>`.
///
/// Submit-time entry. Pod calls this at the head of a `Run` turn before
/// the worker pushes its flattened user message into history; replay
/// derives the worker `Item::user_message` from these segments via
/// [`Segment::flatten_to_text`].
pub async fn save_user_input(
store: &impl Store,
session_id: SessionId,
head_hash: &mut Option<EntryHash>,
segments: Vec<Segment>,
) -> Result<(), StoreError> {
append_entry(
store,
session_id,
head_hash,
LogEntry::UserInput {
ts: session_log::now_millis(),
segments,
},
)
.await
}
/// Log the history delta — new items added since the previous snapshot.
///
/// Classifies items into UserInput, AssistantItems, ToolResults, and
/// HookInjectedItems entries automatically.
/// Classifies items into AssistantItems, ToolResults, and HookInjectedItems
/// entries automatically. User messages are skipped because they are
/// persisted upfront via [`save_user_input`] at submit time; the worker
/// pushes a flattened copy into its history that arrives here in
/// `new_items` and would otherwise produce a duplicate `UserInput` entry.
pub async fn save_delta(
store: &impl Store,
session_id: SessionId,
@@ -158,16 +186,7 @@ pub async fn save_delta(
while i < new_items.len() {
let item = &new_items[i];
if item.is_user_message() {
append_entry(
store,
session_id,
head_hash,
LogEntry::UserInput {
ts,
item: new_items[i].clone(),
},
)
.await?;
// Already persisted by save_user_input at submit time.
i += 1;
} else if item.is_tool_result() {
let start = i;
+102 -11
View File
@@ -10,6 +10,7 @@
use llm_worker::llm_client::types::{Item, RequestConfig};
use llm_worker::{UsageRecord, WorkerResult};
use protocol::Segment;
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
@@ -111,8 +112,12 @@ pub enum LogEntry {
compacted_from: Option<SessionOrigin>,
},
/// User input pushed to history (worker.rs:229).
UserInput { ts: u64, item: Item },
/// User input accepted at submit time. Carries the original typed
/// `Vec<Segment>` so clients can re-render typed atoms (paste chips,
/// file/knowledge refs, workflow invocations) on session restore.
/// Replay flattens these into a `Item::user_message` for the worker
/// history; the worker layer never sees segments directly.
UserInput { ts: u64, segments: Vec<Segment> },
/// Assistant response items added to history (worker.rs:1040-1041).
AssistantItems { ts: u64, items: Vec<LoggedItem> },
@@ -209,6 +214,13 @@ pub struct RestoredState {
/// `LogEntry::Extension` を replay 順に積んだもの。`(domain, payload)`。
/// session-store は domain を不透明扱いし、各ドメインが自前で fold する。
pub extensions: Vec<(String, serde_json::Value)>,
/// User submissions in original typed form, in submit order.
/// One entry per `LogEntry::UserInput`; the K-th entry corresponds to
/// the K-th `Item::user_message` derived during replay (modulo
/// pre-compaction history seeded via `SessionStart.history`, whose
/// original segments are not preserved). Used by clients to re-render
/// typed atoms (paste chips, refs) on session restore.
pub user_segments: Vec<Vec<Segment>>,
}
/// Replay a sequence of hashed entries to reconstruct worker state.
@@ -222,6 +234,7 @@ pub fn collect_state(entries: &[HashedEntry]) -> RestoredState {
head_hash: None,
usage_history: Vec::new(),
extensions: Vec::new(),
user_segments: Vec::new(),
};
for hashed in entries {
@@ -238,8 +251,10 @@ pub fn collect_state(entries: &[HashedEntry]) -> RestoredState {
state.config = config.clone();
state.history = history.iter().cloned().map(Item::from).collect();
}
LogEntry::UserInput { item, .. } => {
state.history.push(item.clone());
LogEntry::UserInput { segments, .. } => {
let text = Segment::flatten_to_text(segments);
state.history.push(Item::user_message(text));
state.user_segments.push(segments.clone());
}
LogEntry::AssistantItems { items, .. } => {
state
@@ -365,7 +380,7 @@ mod tests {
},
LogEntry::UserInput {
ts: 2000,
item: Item::user_message("Hello"),
segments: vec![Segment::text("Hello")],
},
LogEntry::AssistantItems {
ts: 3000,
@@ -400,7 +415,7 @@ mod tests {
},
LogEntry::UserInput {
ts: 2000,
item: Item::user_message("Check weather"),
segments: vec![Segment::text("Check weather")],
},
LogEntry::AssistantItems {
ts: 3000,
@@ -460,7 +475,7 @@ mod tests {
},
LogEntry::UserInput {
ts: 2000,
item: Item::user_message("Hello"),
segments: vec![Segment::text("Hello")],
},
];
let chain_a = build_chain(&raw);
@@ -473,11 +488,11 @@ mod tests {
fn different_content_produces_different_hash() {
let entry_a = LogEntry::UserInput {
ts: 1000,
item: Item::user_message("Hello"),
segments: vec![Segment::text("Hello")],
};
let entry_b = LogEntry::UserInput {
ts: 1000,
item: Item::user_message("World"),
segments: vec![Segment::text("World")],
};
let hash_a = compute_hash(None, &entry_a);
let hash_b = compute_hash(None, &entry_b);
@@ -497,7 +512,7 @@ mod tests {
},
LogEntry::UserInput {
ts: 2000,
item: Item::user_message("hi"),
segments: vec![Segment::text("hi")],
},
LogEntry::LlmUsage {
ts: 2100,
@@ -545,7 +560,7 @@ mod tests {
},
LogEntry::UserInput {
ts: 2000,
item: Item::user_message("hi"),
segments: vec![Segment::text("hi")],
},
]);
let state = collect_state(&entries);
@@ -658,4 +673,80 @@ mod tests {
let parsed = EntryHash::from_hex(&hex).unwrap();
assert_eq!(hash, parsed);
}
/// Mixed segments survive a JSON round-trip through `LogEntry::UserInput`,
/// and `collect_state` derives `Item::user_message` from the flattened
/// text while preserving the original segments separately. This covers
/// the segments → flatten → Item replay path from the ticket.
#[test]
fn replay_user_input_segments_round_trip() {
let segments = vec![
Segment::Text {
content: "see ".into(),
},
Segment::Paste {
id: 1,
chars: 12,
lines: 2,
content: "line1\nline2".into(),
},
Segment::FileRef {
path: "src/main.rs".into(),
},
];
let entry = LogEntry::UserInput {
ts: 4242,
segments: segments.clone(),
};
// Hash + JSON round-trip preserves the variant byte-for-byte.
let json = serde_json::to_string(&entry).unwrap();
let parsed: LogEntry = serde_json::from_str(&json).unwrap();
let entries = build_chain(&[
LogEntry::SessionStart {
ts: 1,
system_prompt: None,
config: RequestConfig::default(),
history: vec![],
forked_from: None,
compacted_from: None,
},
parsed,
]);
let state = collect_state(&entries);
// Worker history gets a flattened user_message item.
assert_eq!(state.history.len(), 1);
match &state.history[0] {
Item::Message { role, content, .. } => {
assert!(matches!(role, llm_worker::Role::User));
assert_eq!(content.len(), 1);
match &content[0] {
llm_worker::ContentPart::Text { text } => {
assert_eq!(
text,
"see line1\nline2[unresolved file ref: src/main.rs]"
);
}
other => panic!("unexpected content: {other:?}"),
}
}
other => panic!("unexpected variant: {other:?}"),
}
// Segments survive verbatim for client-side restore.
assert_eq!(state.user_segments.len(), 1);
assert_eq!(state.user_segments[0].len(), 3);
match &state.user_segments[0][1] {
Segment::Paste {
id,
chars,
lines,
content,
} => {
assert_eq!(*id, 1);
assert_eq!(*chars, 12);
assert_eq!(*lines, 2);
assert_eq!(content, "line1\nline2");
}
other => panic!("expected Paste, got {other:?}"),
}
}
}
+2 -2
View File
@@ -21,7 +21,7 @@ async fn round_trip_write_and_read() {
},
LogEntry::UserInput {
ts: 2000,
item: Item::user_message("Hello"),
segments: vec![protocol::Segment::text("Hello")],
},
LogEntry::AssistantItems {
ts: 3000,
@@ -210,7 +210,7 @@ async fn read_head_hash_returns_last_entry_hash() {
},
LogEntry::UserInput {
ts: 2000,
item: Item::user_message("Hello"),
segments: vec![protocol::Segment::text("Hello")],
},
]);
+16 -8
View File
@@ -99,6 +99,18 @@ async fn run_and_persist(
head_hash: &mut Option<EntryHash>,
input: &str,
) -> (Worker<MockLlmClient>, llm_worker::WorkerResult) {
// Mirror Pod's run-entry contract: log the user input as segments
// before the worker pushes its flattened user_message; save_delta
// skips the resulting user_message item to avoid double-write.
session_store::save_user_input(
store,
session_id,
head_hash,
vec![protocol::Segment::text(input)],
)
.await
.unwrap();
let history_before = worker.history().len();
let mut locked = worker.lock();
@@ -458,7 +470,7 @@ async fn session_auto_forks_on_conflict() {
// Simulate another Pod writing to the same session behind our back
let extra_entry = LogEntry::UserInput {
ts: 9999,
item: Item::user_message("Interloper"),
segments: vec![protocol::Segment::text("Interloper")],
};
let current_head = store.read_head_hash(original_sid).await.unwrap();
let hash = session_store::compute_hash(current_head.as_ref(), &extra_entry);
@@ -492,12 +504,8 @@ async fn session_auto_forks_on_conflict() {
// Original session should still have the interloper entry
let original_entries = store.read_all(original_sid).await.unwrap();
let has_interloper = original_entries.iter().any(|e| {
if let LogEntry::UserInput { item, .. } = &e.entry {
item.is_user_message()
} else {
false
}
});
let has_interloper = original_entries
.iter()
.any(|e| matches!(&e.entry, LogEntry::UserInput { .. }));
assert!(has_interloper);
}