返回 CodeWhale
sse_turn_recovery.rs
根目录 / crates / tui / src / core / engine / tests / sse_turn_recovery.rs
1 //! #5769: a terminal partial SSE loss must not strand the next manual turn.
2 //! These use the real HTTP client and one EngineHandle throughout, with no
3 //! approval, restart, SyncSession, or synthetic continuation between turns.
4
5 use super::*;
6 use serde_json::{Value, json};
7 use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader};
8 use tokio::sync::Notify;
9
10 const FIRST_USER: &str = "FIRST_REAL_USER";
11 const SECOND_USER: &str = "SECOND_REAL_USER";
12 const NEXT_ANSWER: &str = "NEXT_TURN_COMPLETED";
13 const SESSION_ID: &str = "sse-recovery-same-session";
14 const CONTROL_TIMEOUT: Duration = Duration::from_secs(5);
15 const FIXTURE_INPUT_TOKENS: u32 = 13;
16 const FIXTURE_OUTPUT_TOKENS: u32 = 5;
17
18 #[derive(Clone, Copy)]
19 enum Failure {
20 TruncatedBody,
21 StalledBody,
22 }
23
24 struct LoopbackSse {
25 base_url: String,
26 requests: Arc<StdMutex<Vec<Value>>>,
27 stalled_connection_closed: Arc<Notify>,
28 task: tokio::task::JoinHandle<()>,
29 }
30
31 impl Drop for LoopbackSse {
32 fn drop(&mut self) {
33 // The server owns its connection tasks through JoinSet: aborting it
34 // also closes any still-held response, including assertion failures.
35 self.task.abort();
36 }
37 }
38
39 impl LoopbackSse {
40 async fn start(failure: Failure) -> Self {
41 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
42 let base_url = format!("http://{}/v1", listener.local_addr().unwrap());
43 let requests = Arc::new(StdMutex::new(Vec::new()));
44 let captured = Arc::clone(&requests);
45 let stalled_connection_closed = Arc::new(Notify::new());
46 let closed = Arc::clone(&stalled_connection_closed);
47 let task = tokio::spawn(async move {
48 let mut connections = tokio::task::JoinSet::new();
49 loop {
50 tokio::select! {
51 accepted = listener.accept() => {
52 let (socket, _) = accepted.unwrap();
53 let captured = Arc::clone(&captured);
54 let closed = Arc::clone(&closed);
55 connections.spawn(async move {
56 serve_response(socket, failure, captured, closed).await;
57 });
58 }
59 completed = connections.join_next(), if !connections.is_empty() => {
60 completed.unwrap().expect("loopback connection task");
61 }
62 }
63 }
64 });
65 Self {
66 base_url,
67 requests,
68 stalled_connection_closed,
69 task,
70 }
71 }
72 }
73
74 async fn serve_response(
75 socket: tokio::net::TcpStream,
76 failure: Failure,
77 captured: Arc<StdMutex<Vec<Value>>>,
78 closed: Arc<Notify>,
79 ) {
80 let mut reader = BufReader::new(socket);
81 let mut line = String::new();
82 reader.read_line(&mut line).await.unwrap();
83 assert_eq!(line.trim(), "POST /v1/chat/completions HTTP/1.1");
84 let mut content_length = None;
85 loop {
86 line.clear();
87 assert!(reader.read_line(&mut line).await.unwrap() > 0);
88 if line == "\r\n" {
89 break;
90 }
91 if let Some((name, value)) = line.split_once(':')
92 && name.eq_ignore_ascii_case("content-length")
93 {
94 content_length = Some(value.trim().parse::<usize>().unwrap());
95 }
96 }
97 let length = content_length.expect("request body length");
98 assert!(length <= 1024 * 1024);
99 let mut body = vec![0; length];
100 reader.read_exact(&mut body).await.unwrap();
101 let request: Value = serde_json::from_slice(&body).unwrap();
102 let is_second_turn = request["messages"]
103 .as_array()
104 .unwrap()
105 .iter()
106 .any(|message| {
107 message["role"] == "user" && message["content"].to_string().contains(SECOND_USER)
108 });
109 let call = {
110 let mut requests = captured.lock().unwrap();
111 requests.push(request.clone());
112 requests.len()
113 };
114 let text = if is_second_turn {
115 NEXT_ANSWER.to_string()
116 } else {
117 format!("partial-{call}")
118 };
119 let frame = json!({
120 "id": "loopback-sse", "object": "chat.completion.chunk", "model": request["model"],
121 "choices": [{"index": 0, "delta": {"content": text},
122 "finish_reason": if is_second_turn { Some("stop") } else { None }}]
123 });
124 let response = if is_second_turn {
125 // OpenAI-compatible streaming usage arrives in a final choices-empty
126 // frame. Keep the fixture token-only so it proves event separation
127 // without retaining user prompt text anywhere beyond the test request.
128 let usage = json!({
129 "id": "loopback-sse", "object": "chat.completion.chunk", "model": request["model"],
130 "choices": [],
131 "usage": {
132 "prompt_tokens": FIXTURE_INPUT_TOKENS,
133 "completion_tokens": FIXTURE_OUTPUT_TOKENS,
134 },
135 });
136 format!("data: {frame}\n\ndata: {usage}\n\ndata: [DONE]\n\n")
137 } else {
138 format!("data: {frame}\n\n")
139 };
140 // Declaring an unmet length makes a socket close a real reqwest decode
141 // error, rather than a clean EOF or a canned ModelClient error string.
142 let declared_length = if is_second_turn {
143 response.len()
144 } else {
145 1024 * 1024
146 };
147 let mut socket = reader.into_inner();
148 socket.write_all(format!("HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nContent-Length: {declared_length}\r\nConnection: close\r\n\r\n{response}").as_bytes()).await.unwrap();
149 if !is_second_turn && matches!(failure, Failure::StalledBody) {
150 // Do not close or expire the response: cancellation must release the
151 // client connection before any production stream timeout elapses.
152 let mut byte = [0];
153 assert_eq!(socket.read(&mut byte).await.unwrap(), 0);
154 closed.notify_one();
155 } else {
156 socket.shutdown().await.unwrap();
157 }
158 }
159
160 async fn next_event(handle: &EngineHandle) -> Event {
161 handle
162 .rx_event
163 .write()
164 .await
165 .recv()
166 .await
167 .expect("engine event stream remains open")
168 }
169
170 async fn finish_turn(handle: &EngineHandle, events: &mut Vec<Event>) {
171 tokio::time::timeout(model_turn_event_timeout(), async {
172 loop {
173 let event = next_event(handle).await;
174 let terminal = matches!(event, Event::TurnComplete { .. });
175 events.push(event);
176 if terminal {
177 break;
178 }
179 }
180 })
181 .await
182 .expect("same engine must settle the turn");
183 }
184
185 async fn send_user(handle: &EngineHandle, config: &Config, content: &str) {
186 let mut op = external_user_message_op(content, AppMode::Agent, config);
187 if let Op::SendMessage(TurnSpec { allow_shell, .. }) = &mut op {
188 *allow_shell = false;
189 }
190 tokio::time::timeout(CONTROL_TIMEOUT, handle.send(op))
191 .await
192 .expect("same engine mailbox must accept the next user turn")
193 .unwrap();
194 }
195
196 fn terminal_status(events: &[Event]) -> TurnOutcomeStatus {
197 match events.last().unwrap() {
198 Event::TurnComplete { status, .. } => *status,
199 event => panic!("expected terminal TurnComplete, got {event:?}"),
200 }
201 }
202
203 fn terminal_diagnostics(events: &[Event]) -> &crate::tool_inspection::TurnStopDiagnostics {
204 events
205 .iter()
206 .find_map(|event| match event {
207 Event::ToolRequestSnapshot { snapshot } => snapshot.terminal.as_ref(),
208 _ => None,
209 })
210 .expect("terminal request diagnostics")
211 }
212
213 fn retry_status_count(events: &[Event]) -> usize {
214 events
215 .iter()
216 .filter(|event| matches!(event, Event::Status { message } if message.starts_with("Retry attempt: stream-resume ")))
217 .count()
218 }
219
220 async fn verify_next_user_turn_after_loss(failure: Failure) {
221 let server = LoopbackSse::start(failure).await;
222 let workspace = tempdir().unwrap();
223 let config = Config {
224 provider: Some("custom".to_string()),
225 default_text_model: Some(crate::config::DEFAULT_TEXT_MODEL.to_string()),
226 ..Config::default()
227 }
228 .with_legacy_root(
229 Some("synthetic-loopback-key".to_string()),
230 Some(server.base_url.clone()),
231 );
232 let (engine, handle) = Engine::new(
233 EngineConfig {
234 max_steps: 1,
235 terminal_chrome_enabled: true,
236 session_id: Some(SESSION_ID.to_string()),
237 ..deterministic_engine_config(workspace.path())
238 },
239 &config,
240 );
241 let task = tokio::spawn(engine.run());
242 send_user(&handle, &config, FIRST_USER).await;
243 let mut first = Vec::new();
244 if matches!(failure, Failure::StalledBody) {
245 tokio::time::timeout(model_turn_event_timeout(), async {
246 loop {
247 let event = next_event(&handle).await;
248 let partial =
249 matches!(&event, Event::MessageDelta { content, .. } if content == "partial-1");
250 first.push(event);
251 if partial {
252 break;
253 }
254 }
255 })
256 .await
257 .expect("the stalled stream must deliver its partial response");
258 handle.cancel();
259 tokio::time::timeout(CONTROL_TIMEOUT, finish_turn(&handle, &mut first))
260 .await
261 .expect("cancel must interrupt an open response without waiting for its idle timeout");
262 tokio::time::timeout(CONTROL_TIMEOUT, server.stalled_connection_closed.notified())
263 .await
264 .expect("cancel must release the HTTP response connection");
265 assert_eq!(terminal_status(&first), TurnOutcomeStatus::Interrupted);
266 } else {
267 finish_turn(&handle, &mut first).await;
268 assert_eq!(terminal_status(&first), TurnOutcomeStatus::Failed);
269 assert!(
270 first
271 .iter()
272 .any(|event| matches!(event, Event::Error { envelope, .. }
273 if envelope.category == crate::error_taxonomy::ErrorCategory::Network
274 && envelope.message.contains("error decoding response body")))
275 );
276 }
277 let partial_count = match failure {
278 Failure::TruncatedBody => usize::try_from(super::super::MAX_STREAM_RETRIES).unwrap() + 1,
279 Failure::StalledBody => 1,
280 };
281 assert_eq!(
282 server.requests.lock().unwrap().len(),
283 partial_count,
284 "the failed/cancelled turn must settle before the new user turn; no healthy same-turn retry"
285 );
286 let first_terminal = terminal_diagnostics(&first);
287 assert_eq!(
288 usize::try_from(first_terminal.model_requests_started).unwrap(),
289 partial_count,
290 "terminal parent-request count must match POSTs observed by the loopback"
291 );
292 let expected_resumes = match failure {
293 Failure::TruncatedBody => super::super::MAX_STREAM_RETRIES,
294 Failure::StalledBody => 0,
295 };
296 assert_eq!(first_terminal.stream_resumes, expected_resumes);
297 assert_eq!(first_terminal.transparent_stream_retries, 0);
298 assert_eq!(
299 retry_status_count(&first),
300 expected_resumes as usize,
301 "each admitted resume keeps a receipt; the next user turn must start with none"
302 );
303 assert!(
304 !first
305 .iter()
306 .any(|event| matches!(event, Event::TurnUsage { .. })),
307 "a stream without a provider usage frame must not fabricate token usage"
308 );
309
310 // Submit the NEXT real user message immediately after terminal settlement,
311 // using the original handle. There is no reconstruction or --continue.
312 send_user(&handle, &config, SECOND_USER).await;
313 let mut second = Vec::new();
314 finish_turn(&handle, &mut second).await;
315 assert_eq!(terminal_status(&second), TurnOutcomeStatus::Completed);
316 assert!(matches!(
317 second.last(),
318 Some(Event::TurnComplete { error: None, .. })
319 ));
320 assert!(second.iter().any(
321 |event| matches!(event, Event::MessageDelta { content, .. } if content == NEXT_ANSWER)
322 ));
323 let second_terminal = terminal_diagnostics(&second);
324 assert_eq!(second_terminal.model_requests_started, 1);
325 assert_eq!(second_terminal.stream_resumes, 0);
326 assert_eq!(second_terminal.transparent_stream_retries, 0);
327 assert_eq!(retry_status_count(&second), 0);
328 let usage_receipts = second
329 .iter()
330 .filter_map(|event| match event {
331 Event::TurnUsage { usage, .. } => Some(usage),
332 _ => None,
333 })
334 .collect::<Vec<_>>();
335 assert_eq!(usage_receipts.len(), 1);
336 assert_eq!(usage_receipts[0].input_tokens, FIXTURE_INPUT_TOKENS);
337 assert_eq!(usage_receipts[0].output_tokens, FIXTURE_OUTPUT_TOKENS);
338
339 for event in first.iter().chain(&second) {
340 assert!(
341 !matches!(
342 event,
343 Event::ApprovalRequired { .. }
344 | Event::ElevationRequired { .. }
345 | Event::UserInputRequired { .. }
346 ),
347 "this regression has no pending approval or user-input gate: {event:?}"
348 );
349 if let Event::SessionUpdated { session_id, .. } = event {
350 assert_eq!(session_id, SESSION_ID);
351 }
352 }
353 let turn_ids: Vec<_> = first
354 .iter()
355 .chain(&second)
356 .filter_map(|event| match event {
357 Event::TurnStarted { turn_id, .. } => Some(turn_id),
358 _ => None,
359 })
360 .collect();
361 assert_eq!(turn_ids.len(), 2);
362 assert_ne!(
363 turn_ids[0], turn_ids[1],
364 "second completion must be a distinct user turn"
365 );
366 let requests = server.requests.lock().unwrap().clone();
367 assert_eq!(requests.len(), partial_count + 1);
368 assert_eq!(
369 requests.len() - partial_count,
370 1,
371 "the clean second user turn must issue exactly one loopback POST"
372 );
373 let replay = &requests.last().unwrap()["messages"];
374 let replay_text = replay.to_string();
375 for fragment in [FIRST_USER.to_string(), SECOND_USER.to_string()]
376 .into_iter()
377 .chain((1..=partial_count).map(|index| format!("partial-{index}")))
378 {
379 assert_eq!(
380 replay_text.matches(&fragment).count(),
381 1,
382 "missing or duplicated history fragment {fragment}: {replay}"
383 );
384 }
385 assert_eq!(
386 replay
387 .as_array()
388 .unwrap()
389 .iter()
390 .filter(|message| message["role"] == "user")
391 .count(),
392 2
393 );
394
395 let (tx, rx) = tokio::sync::oneshot::channel();
396 handle
397 .send(Op::GetSessionSnapshot {
398 tx: Arc::new(StdMutex::new(Some(tx))),
399 })
400 .await
401 .unwrap();
402 let snapshot = tokio::time::timeout(CONTROL_TIMEOUT, rx)
403 .await
404 .expect("session snapshot stays responsive")
405 .unwrap();
406 let persisted = serde_json::to_string(&snapshot.messages).unwrap();
407 for fragment in [FIRST_USER, SECOND_USER, NEXT_ANSWER] {
408 assert_eq!(persisted.matches(fragment).count(), 1);
409 }
410 for index in 1..=partial_count {
411 assert_eq!(persisted.matches(&format!("partial-{index}")).count(), 1);
412 }
413 let (tx, rx) = tokio::sync::oneshot::channel();
414 handle
415 .send(Op::GetProviderRuntimeStatus {
416 tx: Arc::new(StdMutex::new(Some(tx))),
417 })
418 .await
419 .unwrap();
420 let readiness = tokio::time::timeout(CONTROL_TIMEOUT, rx)
421 .await
422 .expect("provider readiness stays responsive")
423 .unwrap();
424 assert_eq!(
425 readiness.active_provider_requests, 0,
426 "no stream request permit may leak into idle state"
427 );
428 handle.send(Op::Shutdown).await.unwrap();
429 tokio::time::timeout(CONTROL_TIMEOUT, task)
430 .await
431 .expect("shutdown after loss must not block")
432 .unwrap();
433 assert!(
434 !server.task.is_finished(),
435 "loopback server must not have panicked"
436 );
437 }
438
439 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
440 async fn terminal_partial_sse_loss_accepts_next_user_turn_on_same_engine() {
441 verify_next_user_turn_after_loss(Failure::TruncatedBody).await;
442 }
443
444 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
445 async fn cancelled_partial_sse_releases_connection_and_accepts_next_user_turn() {
446 verify_next_user_turn_after_loss(Failure::StalledBody).await;
447 }
448
448 lines RUST