Skip to main content

claude_codex/providers/codex/
continuation.rs

1use std::collections::HashMap;
2use std::sync::Mutex;
3
4use super::translate::request::{ResponsesInputItem, ResponsesRequest};
5
6const TTL_MS: u64 = 30 * 60 * 1000;
7const MAX_STATES: usize = 10_000;
8const MAX_SESSION_TRANSCRIPT_BYTES: u64 = 2_000_000;
9const MAX_TOTAL_TRANSCRIPT_BYTES: u64 = 20_000_000;
10
11#[derive(Clone)]
12struct ContinuationState {
13    response_id: String,
14    prompt_signature: String,
15    transcript: Vec<ResponsesInputItem>,
16    transcript_bytes: u64,
17    updated_at: u64,
18}
19
20static STATES: Mutex<Option<HashMap<String, ContinuationState>>> = Mutex::new(None);
21static TOTAL_TRANSCRIPT_BYTES: Mutex<u64> = Mutex::new(0);
22
23#[derive(Clone)]
24pub struct ContinuationCandidate {
25    pub previous_response_id: Option<String>,
26    pub input_delta: Option<Vec<ResponsesInputItem>>,
27    pub input_delta_count: usize,
28    pub disabled_reason: Option<String>,
29}
30
31fn now_ms() -> u64 {
32    std::time::SystemTime::now()
33        .duration_since(std::time::UNIX_EPOCH)
34        .unwrap_or_default()
35        .as_millis() as u64
36}
37
38pub fn continuation_candidate(
39    session_id: Option<&str>,
40    body: &ResponsesRequest,
41    enabled: bool,
42) -> ContinuationCandidate {
43    let now = now_ms();
44
45    if !enabled {
46        return ContinuationCandidate {
47            previous_response_id: None,
48            input_delta: None,
49            input_delta_count: body.input.len(),
50            disabled_reason: Some("disabled".to_string()),
51        };
52    }
53
54    let session_id = match session_id {
55        Some(s) => s,
56        None => {
57            return ContinuationCandidate {
58                previous_response_id: None,
59                input_delta: None,
60                input_delta_count: body.input.len(),
61                disabled_reason: Some("missing_session".to_string()),
62            };
63        }
64    };
65
66    let state = {
67        let guard = STATES.lock().unwrap();
68        guard.as_ref().and_then(|m| m.get(session_id).cloned())
69    };
70    let state = match state {
71        Some(s) if now - s.updated_at <= TTL_MS => s,
72        Some(_) => {
73            clear_continuation(Some(session_id));
74            return ContinuationCandidate {
75                previous_response_id: None,
76                input_delta: None,
77                input_delta_count: body.input.len(),
78                disabled_reason: Some("missing_state".to_string()),
79            };
80        }
81        None => {
82            return ContinuationCandidate {
83                previous_response_id: None,
84                input_delta: None,
85                input_delta_count: body.input.len(),
86                disabled_reason: Some("missing_state".to_string()),
87            };
88        }
89    };
90
91    let signature = prompt_signature(body);
92    if signature != state.prompt_signature {
93        clear_continuation(Some(session_id));
94        return ContinuationCandidate {
95            previous_response_id: None,
96            input_delta: None,
97            input_delta_count: body.input.len(),
98            disabled_reason: Some("prompt_changed".to_string()),
99        };
100    }
101
102    let suffix = input_suffix_after_prefix(&body.input, &state.transcript);
103    let suffix = match suffix {
104        Some(s) => s,
105        None => {
106            clear_continuation(Some(session_id));
107            return ContinuationCandidate {
108                previous_response_id: None,
109                input_delta: None,
110                input_delta_count: body.input.len(),
111                disabled_reason: Some("not_append_only".to_string()),
112            };
113        }
114    };
115
116    if suffix.is_empty() {
117        return ContinuationCandidate {
118            previous_response_id: None,
119            input_delta: None,
120            input_delta_count: 0,
121            disabled_reason: Some("empty_delta".to_string()),
122        };
123    }
124
125    ContinuationCandidate {
126        previous_response_id: Some(state.response_id),
127        input_delta: Some(suffix.clone()),
128        input_delta_count: suffix.len(),
129        disabled_reason: None,
130    }
131}
132
133pub fn record_continuation(
134    session_id: Option<&str>,
135    request_body: &ResponsesRequest,
136    response_id: Option<&str>,
137    output_items: &[ResponsesInputItem],
138) {
139    let session_id = match session_id {
140        Some(s) => s,
141        None => return,
142    };
143
144    let response_id = match response_id {
145        Some(id) => id.to_string(),
146        None => {
147            clear_continuation(Some(session_id));
148            return;
149        }
150    };
151
152    let mut transcript: Vec<ResponsesInputItem> = request_body.input.clone();
153    transcript.extend_from_slice(output_items);
154
155    let transcript_json = serde_json::to_string(&transcript).unwrap_or_default();
156    let transcript_bytes = transcript_json.len() as u64;
157
158    if transcript_bytes > MAX_SESSION_TRANSCRIPT_BYTES {
159        clear_continuation(Some(session_id));
160        return;
161    }
162
163    clear_continuation(Some(session_id));
164
165    let state = ContinuationState {
166        response_id,
167        prompt_signature: prompt_signature(request_body),
168        transcript,
169        transcript_bytes,
170        updated_at: now_ms(),
171    };
172
173    {
174        let mut guard = TOTAL_TRANSCRIPT_BYTES.lock().unwrap();
175        *guard += transcript_bytes;
176    }
177    {
178        let mut guard = STATES.lock().unwrap();
179        let map = guard.get_or_insert_with(HashMap::new);
180        map.insert(session_id.to_string(), state);
181    }
182    evict_oldest();
183}
184
185pub fn clear_continuation(session_id: Option<&str>) {
186    let session_id = match session_id {
187        Some(s) => s,
188        None => return,
189    };
190    let mut guard = STATES.lock().unwrap();
191    if let Some(map) = guard.as_mut()
192        && let Some(existing) = map.remove(session_id)
193    {
194        let mut bytes_guard = TOTAL_TRANSCRIPT_BYTES.lock().unwrap();
195        *bytes_guard = bytes_guard.saturating_sub(existing.transcript_bytes);
196    }
197}
198
199pub fn has_continuation_for_tests(session_id: &str) -> bool {
200    let guard = STATES.lock().unwrap();
201    guard.as_ref().is_some_and(|m| m.contains_key(session_id))
202}
203
204pub fn clear_all_continuations_for_tests() {
205    let mut guard = STATES.lock().unwrap();
206    *guard = None;
207    let mut bytes_guard = TOTAL_TRANSCRIPT_BYTES.lock().unwrap();
208    *bytes_guard = 0;
209}
210
211fn input_suffix_after_prefix(
212    input: &[ResponsesInputItem],
213    prefix: &[ResponsesInputItem],
214) -> Option<Vec<ResponsesInputItem>> {
215    if prefix.len() > input.len() {
216        return None;
217    }
218    for i in 0..prefix.len() {
219        let a = serde_json::to_value(&input[i]).unwrap_or_default();
220        let b = serde_json::to_value(&prefix[i]).unwrap_or_default();
221        if a != b {
222            return None;
223        }
224    }
225    Some(input[prefix.len()..].to_vec())
226}
227
228fn prompt_signature(body: &ResponsesRequest) -> String {
229    let value = serde_json::to_value(body).unwrap_or_default();
230    let obj = match value.as_object() {
231        Some(o) => o,
232        None => return String::new(),
233    };
234    let mut entries: Vec<(&String, &serde_json::Value)> =
235        obj.iter().filter(|(k, _)| *k != "input").collect();
236    entries.sort_by_key(|(a, _)| *a);
237    let mut sig = String::from("{");
238    for (i, (key, val)) in entries.iter().enumerate() {
239        if i > 0 {
240            sig.push(',');
241        }
242        sig.push_str(&format!("\"{}\":{}", key, stable_json(val)));
243    }
244    sig.push('}');
245    sig
246}
247
248fn stable_json(value: &serde_json::Value) -> String {
249    match value {
250        serde_json::Value::Null => "null".to_string(),
251        serde_json::Value::Bool(b) => b.to_string(),
252        serde_json::Value::Number(n) => n.to_string(),
253        serde_json::Value::String(s) => serde_json::to_string(s).unwrap_or_default(),
254        serde_json::Value::Array(arr) => {
255            let items: Vec<String> = arr.iter().map(stable_json).collect();
256            format!("[{}]", items.join(","))
257        }
258        serde_json::Value::Object(obj) => {
259            let mut entries: Vec<(&String, &serde_json::Value)> = obj.iter().collect();
260            entries.sort_by_key(|(a, _)| *a);
261            let items: Vec<String> = entries
262                .iter()
263                .map(|(k, v)| {
264                    format!(
265                        "{}:{}",
266                        serde_json::to_string(k).unwrap_or_default(),
267                        stable_json(v)
268                    )
269                })
270                .collect();
271            format!("{{{}}}", items.join(","))
272        }
273    }
274}
275
276fn evict_oldest() {
277    let mut guard = STATES.lock().unwrap();
278    let map = match guard.as_mut() {
279        Some(m) => m,
280        None => return,
281    };
282    let mut bytes_guard = TOTAL_TRANSCRIPT_BYTES.lock().unwrap();
283    while map.len() > MAX_STATES || *bytes_guard > MAX_TOTAL_TRANSCRIPT_BYTES {
284        let key = map.keys().next().cloned();
285        match key {
286            Some(k) => {
287                if let Some(existing) = map.remove(&k) {
288                    *bytes_guard = bytes_guard.saturating_sub(existing.transcript_bytes);
289                }
290            }
291            None => break,
292        }
293    }
294}
295
296#[cfg(test)]
297mod tests {
298    use super::*;
299    use serde_json::json;
300
301    fn request_with_input(
302        input: Vec<ResponsesInputItem>,
303        extra: Option<serde_json::Value>,
304    ) -> ResponsesRequest {
305        let mut fields = serde_json::Map::new();
306        fields.insert("model".into(), json!("gpt-5.5"));
307        fields.insert("input".into(), json!(input));
308        fields.insert("store".into(), json!(false));
309        fields.insert("stream".into(), json!(true));
310        fields.insert("text".into(), json!({"verbosity": "low"}));
311        fields.insert("parallel_tool_calls".into(), json!(true));
312        if let Some(extras) = extra
313            && let Some(obj) = extras.as_object()
314        {
315            for (k, v) in obj {
316                fields.insert(k.clone(), v.clone());
317            }
318        }
319        serde_json::from_value(serde_json::Value::Object(fields)).unwrap()
320    }
321
322    #[test]
323    fn continuation_behaviors() {
324        // All tests run in sequence to avoid global state interference
325
326        // disabled_when_not_enabled
327        clear_all_continuations_for_tests();
328        let input = vec![ResponsesInputItem::Message {
329            role: "user".to_string(),
330            content: vec![
331                super::super::translate::request::ResponsesContentPart::InputText {
332                    text: "one".to_string(),
333                },
334            ],
335        }];
336        let req = request_with_input(input, None);
337        let result = continuation_candidate(Some("s1"), &req, false);
338        assert_eq!(result.disabled_reason, Some("disabled".to_string()));
339        assert_eq!(result.input_delta_count, 1);
340
341        // missing_session
342        clear_all_continuations_for_tests();
343        let input = vec![ResponsesInputItem::Message {
344            role: "user".to_string(),
345            content: vec![
346                super::super::translate::request::ResponsesContentPart::InputText {
347                    text: "one".to_string(),
348                },
349            ],
350        }];
351        let req = request_with_input(input, None);
352        let result = continuation_candidate(None, &req, true);
353        assert_eq!(result.disabled_reason, Some("missing_session".to_string()));
354
355        // uses_previous_response_id_for_append_only
356        clear_all_continuations_for_tests();
357        let input = vec![ResponsesInputItem::Message {
358            role: "user".to_string(),
359            content: vec![
360                super::super::translate::request::ResponsesContentPart::InputText {
361                    text: "one".to_string(),
362                },
363            ],
364        }];
365        let req = request_with_input(input, None);
366        record_continuation(Some("s1"), &req, Some("resp_1"), &[]);
367
368        let input2 = vec![
369            ResponsesInputItem::Message {
370                role: "user".to_string(),
371                content: vec![
372                    super::super::translate::request::ResponsesContentPart::InputText {
373                        text: "one".to_string(),
374                    },
375                ],
376            },
377            ResponsesInputItem::Message {
378                role: "user".to_string(),
379                content: vec![
380                    super::super::translate::request::ResponsesContentPart::InputText {
381                        text: "two".to_string(),
382                    },
383                ],
384            },
385        ];
386        let req2 = request_with_input(input2, None);
387        let result = continuation_candidate(Some("s1"), &req2, true);
388        assert_eq!(result.previous_response_id, Some("resp_1".to_string()));
389        assert_eq!(result.input_delta_count, 1);
390
391        // clears_state_when_prompt_signature_changes
392        clear_all_continuations_for_tests();
393        let input = vec![ResponsesInputItem::Message {
394            role: "user".to_string(),
395            content: vec![
396                super::super::translate::request::ResponsesContentPart::InputText {
397                    text: "one".to_string(),
398                },
399            ],
400        }];
401        let req = request_with_input(input.clone(), None);
402        record_continuation(Some("s1"), &req, Some("resp_1"), &[]);
403
404        let req2 = request_with_input(input, Some(json!({"service_tier": "flex"})));
405        let result = continuation_candidate(Some("s1"), &req2, true);
406        assert_eq!(result.disabled_reason, Some("prompt_changed".to_string()));
407        assert!(!has_continuation_for_tests("s1"));
408
409        // clears_state_when_missing_response_id
410        clear_all_continuations_for_tests();
411        let input = vec![ResponsesInputItem::Message {
412            role: "user".to_string(),
413            content: vec![
414                super::super::translate::request::ResponsesContentPart::InputText {
415                    text: "one".to_string(),
416                },
417            ],
418        }];
419        let req = request_with_input(input.clone(), None);
420        record_continuation(Some("s1"), &req, Some("resp_1"), &[]);
421        assert!(has_continuation_for_tests("s1"));
422
423        record_continuation(Some("s1"), &req, None, &[]);
424        assert!(!has_continuation_for_tests("s1"));
425    }
426}