Skip to main content

atman_runtime/event_log/
replay.rs

1use std::collections::{HashMap, HashSet, VecDeque};
2use std::io::BufRead;
3use std::path::Path;
4
5use crate::event::{Event, FlowRunId};
6use crate::event_log::reader::{
7    ReplayRecord, context_snapshot_from_records, read_replay_records, scan_replay_records,
8};
9use crate::message::Message;
10use crate::projection::message_window::{
11    TranscriptEntry, apply_attachment_degradation, apply_envelope_to_messages,
12    message_belongs_to_root, project_transcript_records,
13};
14use crate::session::{ContextSnapshot, SessionOpenError};
15
16pub trait TranscriptReplayObserver {
17    fn observe(&mut self, entry: TranscriptEntry);
18}
19
20impl<F> TranscriptReplayObserver for F
21where
22    F: FnMut(TranscriptEntry),
23{
24    fn observe(&mut self, entry: TranscriptEntry) {
25        self(entry);
26    }
27}
28
29pub struct ReplayBundle {
30    pub last_seq: Option<u64>,
31    pub compacted_messages: Vec<(u64, Message)>,
32    pub all_messages: Vec<(u64, Message)>,
33    pub context: ContextSnapshot,
34    pub deferred_form_answers: Vec<crate::form::DeferredFormAnswer>,
35}
36
37#[derive(Debug, Default)]
38pub(crate) struct FlowOwnership {
39    pub known: HashSet<FlowRunId>,
40    pub spawned: HashSet<FlowRunId>,
41}
42
43impl FlowOwnership {
44    fn from_records(records: &[ReplayRecord]) -> Self {
45        let mut known = HashSet::new();
46        let mut spawned = HashSet::new();
47        let mut children = HashMap::<FlowRunId, Vec<FlowRunId>>::new();
48        for record in records {
49            let Event::FlowStart {
50                run_id,
51                parent_run_id,
52                spawned: is_spawned,
53                ..
54            } = &record.envelope.event
55            else {
56                continue;
57            };
58            known.insert(run_id.clone());
59            if let Some(parent) = parent_run_id {
60                children
61                    .entry(parent.clone())
62                    .or_default()
63                    .push(run_id.clone());
64            }
65            if *is_spawned {
66                spawned.insert(run_id.clone());
67            }
68        }
69        let mut queue = spawned.iter().cloned().collect::<VecDeque<_>>();
70        while let Some(parent) = queue.pop_front() {
71            let Some(descendants) = children.get(&parent) else {
72                continue;
73            };
74            for descendant in descendants {
75                if spawned.insert(descendant.clone()) {
76                    queue.push_back(descendant.clone());
77                }
78            }
79        }
80        Self { known, spawned }
81    }
82}
83
84pub struct SessionReplay;
85
86impl SessionReplay {
87    pub fn from_path(
88        path: &Path,
89        observer: Option<&mut dyn TranscriptReplayObserver>,
90    ) -> Result<ReplayBundle, SessionOpenError> {
91        let records = read_replay_records(path)?;
92        Ok(Self::from_records(records, observer))
93    }
94
95    pub fn from_reader<R: BufRead>(
96        reader: R,
97        observer: Option<&mut dyn TranscriptReplayObserver>,
98    ) -> std::io::Result<ReplayBundle> {
99        let records = scan_replay_records(reader)?;
100        Ok(Self::from_records(records, observer))
101    }
102
103    fn from_records(
104        records: Vec<ReplayRecord>,
105        observer: Option<&mut dyn TranscriptReplayObserver>,
106    ) -> ReplayBundle {
107        let ownership = FlowOwnership::from_records(&records);
108        let mut compacted_messages = Vec::new();
109        let mut compacted_positions = HashMap::new();
110        let mut all_messages = Vec::new();
111        let mut all_positions = HashMap::new();
112        let mut deferred_answers = Vec::new();
113        for record in &records {
114            apply_envelope_to_messages(
115                &record.envelope,
116                &ownership.spawned,
117                &mut compacted_messages,
118                &mut compacted_positions,
119            );
120            match &record.envelope.event {
121                Event::DeferredFormRecorded { answer } => {
122                    if !deferred_answers
123                        .iter()
124                        .any(|pending: &crate::form::DeferredFormAnswer| {
125                            pending.prompt_id == answer.prompt_id
126                        })
127                    {
128                        deferred_answers.push(answer.clone());
129                    }
130                }
131                Event::DeferredFormApplied {
132                    prompt_id, message, ..
133                } => {
134                    deferred_answers.retain(|pending| &pending.prompt_id != prompt_id);
135                    all_positions.insert(record.envelope.seq, all_messages.len());
136                    all_messages.push((record.envelope.seq, message.clone()));
137                }
138                Event::UserMsg {
139                    message,
140                    flow_run_id,
141                    ..
142                }
143                | Event::AssistantMsg {
144                    message,
145                    flow_run_id,
146                    ..
147                }
148                | Event::ToolResultMsg {
149                    message,
150                    flow_run_id,
151                    ..
152                }
153                | Event::SystemMsg {
154                    message,
155                    flow_run_id,
156                    ..
157                } if message_belongs_to_root(flow_run_id.as_ref(), &ownership.spawned) => {
158                    all_positions.insert(record.envelope.seq, all_messages.len());
159                    all_messages.push((record.envelope.seq, message.clone()));
160                }
161                Event::AttachmentDegraded {
162                    message_seq,
163                    part_index,
164                    file_basename,
165                    reason,
166                    ..
167                } => {
168                    apply_attachment_degradation(
169                        &mut all_messages,
170                        &all_positions,
171                        *message_seq,
172                        *part_index,
173                        file_basename,
174                        reason,
175                    );
176                }
177                _ => {}
178            }
179        }
180        let context = context_snapshot_from_records(&records);
181        let last_seq = records.last().map(|record| record.envelope.seq);
182        if let Some(observer) = observer {
183            for entry in project_transcript_records(&records, &ownership) {
184                observer.observe(entry);
185            }
186        }
187        ReplayBundle {
188            last_seq,
189            compacted_messages,
190            all_messages,
191            context,
192            deferred_form_answers: deferred_answers,
193        }
194    }
195}
196
197pub fn transcript_from_envelopes(
198    envelopes: &[crate::event::EventEnvelope],
199) -> Vec<TranscriptEntry> {
200    let records = envelopes
201        .iter()
202        .cloned()
203        .map(|envelope| ReplayRecord {
204            persisted_ts: Some(envelope.ts),
205            envelope,
206        })
207        .collect::<Vec<_>>();
208    let ownership = FlowOwnership::from_records(&records);
209    project_transcript_records(&records, &ownership)
210}
211
212#[cfg(test)]
213mod tests {
214    use super::*;
215
216    #[test]
217    fn reader_parses_each_nonempty_line_once() {
218        let event = crate::event::EventEnvelope::new(
219            1,
220            Event::TurnStart {
221                turn_id: crate::event::TurnId::now(),
222            },
223        );
224        let line = serde_json::to_string(&event).unwrap();
225        let input = format!("{line}\nnot-json\n\n{line}\n");
226        crate::event_log::reader::reset_parse_attempts();
227
228        let bundle = SessionReplay::from_reader(input.as_bytes(), None).unwrap();
229
230        assert_eq!(crate::event_log::reader::parse_attempts(), 3);
231        assert_eq!(bundle.last_seq, Some(1));
232    }
233}