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 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 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 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 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 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}