1use std::collections::BTreeMap;
2
3use serde::{Deserialize, Serialize};
4use time::OffsetDateTime;
5
6use crate::{
7 events::{ThreadId, TurnId},
8 extension_state::{ExtensionStateRecord, ExtensionStoreScope},
9};
10
11pub const TURN_LIFECYCLE_EXTENSION_ID: &str = "roder.lifecycle";
12pub const TURN_LIFECYCLE_STATE_KEY: &str = "turn_lifecycle";
13pub const TURN_LIFECYCLE_CORRUPTION_STATE_KEY: &str = "turn_lifecycle_corruption";
14pub const TURN_LIFECYCLE_SCHEMA_VERSION: u32 = 1;
15
16#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
18#[serde(rename_all = "snake_case")]
19pub enum TurnLifecycleState {
20 Running,
21 InterruptRequested,
22 Interrupted,
23 Completed,
24 Failed,
25 RecoveryNeeded,
26}
27
28impl TurnLifecycleState {
29 pub fn is_terminal(self) -> bool {
30 matches!(
31 self,
32 Self::Interrupted | Self::Completed | Self::Failed | Self::RecoveryNeeded
33 )
34 }
35
36 pub fn requires_recovery(self) -> bool {
37 matches!(self, Self::Running | Self::InterruptRequested)
38 }
39}
40
41#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq)]
43#[serde(rename_all = "snake_case")]
44pub enum TurnCleanupState {
45 #[default]
46 NotRequested,
47 Requested,
48 Completed,
49 TimedOut,
50 Unknown,
51}
52
53#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq)]
57#[serde(rename_all = "snake_case")]
58pub enum TurnCleanupOwnership {
59 #[default]
62 RuntimeTaskOnly,
63 ProviderCleanupPending,
66 ProviderCleanupConfirmed,
69}
70
71#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
73#[serde(rename_all = "snake_case")]
74pub enum TurnLifecycleReason {
75 UserInterrupt,
76 Shutdown,
77 DeadlineExceeded,
78 ProviderFailure,
79 RuntimeRestart,
80 RuntimeFailure,
81}
82
83#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
85#[serde(rename_all = "camelCase")]
86pub struct TurnLifecycleRecord {
87 pub thread_id: ThreadId,
88 pub turn_id: TurnId,
89 pub state: TurnLifecycleState,
90 #[serde(default)]
91 pub cleanup: TurnCleanupState,
92 #[serde(default, skip_serializing_if = "Option::is_none")]
93 pub reason: Option<TurnLifecycleReason>,
94 #[serde(default)]
95 pub ownership: TurnCleanupOwnership,
96 #[serde(with = "time::serde::rfc3339")]
97 pub timestamp: OffsetDateTime,
98}
99
100impl TurnLifecycleRecord {
101 pub fn new(
102 thread_id: ThreadId,
103 turn_id: TurnId,
104 state: TurnLifecycleState,
105 cleanup: TurnCleanupState,
106 reason: Option<TurnLifecycleReason>,
107 timestamp: OffsetDateTime,
108 ) -> Self {
109 Self {
110 thread_id,
111 turn_id,
112 state,
113 cleanup,
114 reason,
115 ownership: TurnCleanupOwnership::default(),
116 timestamp,
117 }
118 }
119
120 pub fn with_ownership(mut self, ownership: TurnCleanupOwnership) -> Self {
121 self.ownership = ownership;
122 self
123 }
124
125 pub fn extension_state(&self) -> anyhow::Result<ExtensionStateRecord> {
126 Ok(ExtensionStateRecord {
127 extension_id: TURN_LIFECYCLE_EXTENSION_ID.to_string(),
128 key: TURN_LIFECYCLE_STATE_KEY.to_string(),
129 scope: ExtensionStoreScope::Turn {
130 thread_id: self.thread_id.clone(),
131 turn_id: self.turn_id.clone(),
132 },
133 schema_version: TURN_LIFECYCLE_SCHEMA_VERSION,
134 value: serde_json::to_value(self)?,
135 })
136 }
137
138 pub fn from_extension_state(record: &ExtensionStateRecord) -> anyhow::Result<Option<Self>> {
139 if record.extension_id != TURN_LIFECYCLE_EXTENSION_ID
140 || record.key != TURN_LIFECYCLE_STATE_KEY
141 {
142 return Ok(None);
143 }
144
145 anyhow::ensure!(
146 record.schema_version == TURN_LIFECYCLE_SCHEMA_VERSION,
147 "unsupported turn lifecycle schema version {}",
148 record.schema_version
149 );
150
151 let decoded: Self = serde_json::from_value(record.value.clone())?;
152 match &record.scope {
153 ExtensionStoreScope::Turn { thread_id, turn_id }
154 if thread_id == &decoded.thread_id && turn_id == &decoded.turn_id => {}
155 ExtensionStoreScope::Turn { .. } => anyhow::bail!(
156 "turn lifecycle record scope does not match its embedded thread and turn identifiers"
157 ),
158 _ => anyhow::bail!("turn lifecycle records must use turn scope"),
159 }
160
161 Ok(Some(decoded))
162 }
163}
164
165#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)]
167#[serde(rename_all = "camelCase")]
168pub struct TurnLifecycleSnapshot {
169 pub records: Vec<TurnLifecycleRecord>,
170 pub corrupt_record_count: usize,
171}
172
173#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)]
177#[serde(rename_all = "camelCase")]
178pub struct LifecycleMetricsSnapshot {
179 pub shutdown_drain_count: u64,
180 pub clean_shutdown_count: u64,
181 pub deadline_exceeded_count: u64,
182 pub persistence_failed_count: u64,
183 pub restart_reconciliation_count: u64,
184 pub lifecycle_persistence_failure_count: u64,
185 pub shutdown_drain_duration_ms_total: u64,
186 pub provider_cleanup_confirmed_count: u64,
187 pub provider_cleanup_timed_out_count: u64,
188 pub provider_cleanup_unknown_count: u64,
189}
190
191pub fn turn_lifecycle_corruption_marker(
195 thread_id: ThreadId,
196 corrupt_record_count: usize,
197) -> ExtensionStateRecord {
198 ExtensionStateRecord {
199 extension_id: TURN_LIFECYCLE_EXTENSION_ID.to_string(),
200 key: TURN_LIFECYCLE_CORRUPTION_STATE_KEY.to_string(),
201 scope: ExtensionStoreScope::Thread { thread_id },
202 schema_version: TURN_LIFECYCLE_SCHEMA_VERSION,
203 value: serde_json::json!({ "count": corrupt_record_count }),
204 }
205}
206
207pub fn latest_turn_lifecycle_records(
210 records: &[ExtensionStateRecord],
211) -> (BTreeMap<TurnId, TurnLifecycleRecord>, usize) {
212 let mut latest = BTreeMap::new();
213 let mut corrupt_record_count = 0;
214
215 for record in records {
216 match TurnLifecycleRecord::from_extension_state(record) {
217 Ok(Some(decoded)) => {
218 let should_replace =
219 latest
220 .get(&decoded.turn_id)
221 .is_none_or(|current: &TurnLifecycleRecord| {
222 current.timestamp <= decoded.timestamp
223 });
224 if should_replace {
225 latest.insert(decoded.turn_id.clone(), decoded);
226 }
227 }
228 Ok(None) => {}
229 Err(_) => corrupt_record_count += 1,
230 }
231 }
232
233 (latest, corrupt_record_count)
234}
235
236pub fn turn_lifecycle_snapshot(records: &[ExtensionStateRecord]) -> TurnLifecycleSnapshot {
237 let marker_count: usize = records
238 .iter()
239 .filter(|record| {
240 record.extension_id == TURN_LIFECYCLE_EXTENSION_ID
241 && record.key == TURN_LIFECYCLE_CORRUPTION_STATE_KEY
242 })
243 .map(|record| {
244 record
245 .value
246 .get("count")
247 .and_then(serde_json::Value::as_u64)
248 .and_then(|count| usize::try_from(count).ok())
249 .unwrap_or(1)
250 })
251 .sum();
252 let (records, corrupt_record_count) = latest_turn_lifecycle_records(records);
253
254 TurnLifecycleSnapshot {
255 records: records.into_values().collect(),
256 corrupt_record_count: corrupt_record_count + marker_count,
257 }
258}
259
260#[cfg(test)]
261mod tests {
262 use super::*;
263
264 fn record(state: TurnLifecycleState, timestamp: OffsetDateTime) -> TurnLifecycleRecord {
265 TurnLifecycleRecord {
266 thread_id: "thread-1".to_string(),
267 turn_id: "turn-1".to_string(),
268 state,
269 cleanup: TurnCleanupState::NotRequested,
270 reason: None,
271 ownership: TurnCleanupOwnership::RuntimeTaskOnly,
272 timestamp,
273 }
274 }
275
276 #[test]
277 fn lifecycle_record_round_trips_through_extension_state() {
278 let original = record(
279 TurnLifecycleState::InterruptRequested,
280 OffsetDateTime::UNIX_EPOCH,
281 );
282
283 let state = original.extension_state().expect("record should encode");
284 let decoded = TurnLifecycleRecord::from_extension_state(&state)
285 .expect("record should decode")
286 .expect("record should be recognized");
287
288 assert_eq!(decoded, original);
289 }
290
291 #[test]
292 fn ownership_defaults_for_legacy_lifecycle_records() {
293 let legacy = serde_json::json!({
294 "threadId": "thread-1",
295 "turnId": "turn-1",
296 "state": "interrupted",
297 "cleanup": "unknown",
298 "timestamp": "1970-01-01T00:00:00Z"
299 });
300
301 let record: TurnLifecycleRecord = serde_json::from_value(legacy).unwrap();
302
303 assert_eq!(record.ownership, TurnCleanupOwnership::RuntimeTaskOnly);
304 }
305
306 #[test]
307 fn latest_records_keep_newest_valid_transition_and_count_corruption() {
308 let earlier = record(TurnLifecycleState::Running, OffsetDateTime::UNIX_EPOCH);
309 let later = record(
310 TurnLifecycleState::Interrupted,
311 OffsetDateTime::UNIX_EPOCH + time::Duration::seconds(60),
312 );
313 let mut corrupt = later.extension_state().expect("record should encode");
314 corrupt.schema_version = 99;
315
316 let (records, corrupt_record_count) = latest_turn_lifecycle_records(&[
317 earlier.extension_state().expect("record should encode"),
318 corrupt,
319 later.extension_state().expect("record should encode"),
320 ]);
321
322 assert_eq!(corrupt_record_count, 1);
323 assert_eq!(records.get("turn-1"), Some(&later));
324 }
325
326 #[test]
327 fn only_non_terminal_states_require_recovery() {
328 assert!(TurnLifecycleState::Running.requires_recovery());
329 assert!(TurnLifecycleState::InterruptRequested.requires_recovery());
330 assert!(!TurnLifecycleState::Interrupted.requires_recovery());
331 assert!(!TurnLifecycleState::Completed.requires_recovery());
332 assert!(!TurnLifecycleState::Failed.requires_recovery());
333 assert!(!TurnLifecycleState::RecoveryNeeded.requires_recovery());
334 }
335
336 #[test]
337 fn corruption_marker_is_reflected_without_exposing_raw_record_data() {
338 let snapshot =
339 turn_lifecycle_snapshot(&[turn_lifecycle_corruption_marker("thread-1".to_string(), 2)]);
340
341 assert!(snapshot.records.is_empty());
342 assert_eq!(snapshot.corrupt_record_count, 2);
343 }
344}