返回 CodeWhale
sse.rs
根目录 / crates / tui / src / conformance / sse.rs
1 //! Recorded-SSE family: provider wire bytes → normalized stream events.
2 //!
3 //! Each case replays a synthetic recording (`<case>.sse`, never a real
4 //! provider capture) from a loopback HTTP server into the real
5 //! `CodewhaleClient::create_message_stream` for one route, and pins what the
6 //! wire adapter yields. The first golden line is the request the adapter
7 //! sent (method and path only — the body is the prompt family's business);
8 //! every later line is one normalized [`super::stream_json`] event, or one
9 //! `stream_failure` / `open_failure` record with exact `detail` bytes.
10
11 use std::sync::{Arc, Mutex};
12 use std::time::Duration;
13
14 use futures_util::StreamExt;
15 use serde_json::{Value, json};
16 use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader};
17
18 use codewhale_models::{ContentBlock, Message, MessageRequest, Role};
19
20 use super::golden::{self, Failures};
21 use super::stream_json;
22 use crate::client::CodewhaleClient;
23 use crate::config::{Config, ProviderConfig, ProvidersConfig};
24 use crate::llm_client::LlmClient;
25 use crate::test_support::{EnvVarGuard, lock_test_env};
26
27 const FAMILY: &str = "sse";
28 const STREAM_DEADLINE: Duration = Duration::from_secs(20);
29
30 #[derive(Debug, Clone)]
31 struct ReceivedRequest {
32 method: String,
33 path: String,
34 }
35
36 #[tokio::test]
37 async fn harness_timeout_rejects_a_real_stalled_sse_connection() {
38 let mut case = golden::read_case(FAMILY, "anthropic_thinking_text_usage");
39 case["hold_open"] = json!(true);
40 case["framing"] = json!("chunked");
41 let outcome = replay_with_deadline(&case, Vec::new(), Duration::from_secs(1)).await;
42 assert!(
43 matches!(outcome, Err(error) if error.contains("harness timeout: SSE replay")),
44 "a stalled SSE connection became recordable output"
45 );
46 }
47
48 #[tokio::test]
49 async fn empty_sse_recording_cannot_be_qualified() {
50 let case = golden::read_case(FAMILY, "anthropic_thinking_text_usage");
51 assert!(
52 replay(&case, Vec::new()).await.is_err(),
53 "empty recording became a passing case"
54 );
55 }
56
57 /// One-shot-per-connection loopback server that answers every request with
58 /// the recording. `close` framing ends the body by closing the socket (a
59 /// provider that stops sending); `chunked` framing uses HTTP chunks, and
60 /// `transport_cut` drops the connection before the terminating chunk so the
61 /// client sees a transport error rather than a clean EOF.
62 struct LoopbackRecording {
63 base: String,
64 requests: Arc<Mutex<Vec<ReceivedRequest>>>,
65 task: tokio::task::JoinHandle<()>,
66 }
67
68 impl Drop for LoopbackRecording {
69 fn drop(&mut self) {
70 self.task.abort();
71 }
72 }
73
74 #[derive(Clone)]
75 struct Framing {
76 chunked: bool,
77 transport_cut: bool,
78 chunk_bytes: usize,
79 hold_open: bool,
80 }
81
82 impl LoopbackRecording {
83 async fn start(body: Vec<u8>, framing: Framing) -> Self {
84 let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
85 .await
86 .expect("bind loopback");
87 let base = format!("http://{}", listener.local_addr().expect("local addr"));
88 let requests = Arc::new(Mutex::new(Vec::new()));
89 let captured = Arc::clone(&requests);
90 let task = tokio::spawn(async move {
91 let mut connections = tokio::task::JoinSet::new();
92 loop {
93 let accepted = tokio::select! {
94 accepted = listener.accept() => accepted,
95 _ = connections.join_next(), if !connections.is_empty() => continue,
96 };
97 let Ok((socket, _)) = accepted else {
98 break;
99 };
100 let body = body.clone();
101 let framing = framing.clone();
102 let captured = Arc::clone(&captured);
103 connections.spawn(async move {
104 serve(socket, body, framing, captured).await;
105 });
106 }
107 });
108 Self {
109 base,
110 requests,
111 task,
112 }
113 }
114 }
115
116 async fn serve(
117 socket: tokio::net::TcpStream,
118 body: Vec<u8>,
119 framing: Framing,
120 captured: Arc<Mutex<Vec<ReceivedRequest>>>,
121 ) {
122 let mut reader = BufReader::new(socket);
123 let mut line = String::new();
124 if reader.read_line(&mut line).await.unwrap_or(0) == 0 {
125 return;
126 }
127 let mut parts = line.split_whitespace();
128 let method = parts.next().unwrap_or_default().to_string();
129 let path = parts.next().unwrap_or_default().to_string();
130 let mut content_length = 0usize;
131 loop {
132 line.clear();
133 if reader.read_line(&mut line).await.unwrap_or(0) == 0 {
134 return;
135 }
136 if line == "\r\n" || line == "\n" {
137 break;
138 }
139 if let Some((name, value)) = line.split_once(':')
140 && name.eq_ignore_ascii_case("content-length")
141 {
142 content_length = value.trim().parse().unwrap_or(0);
143 }
144 }
145 let mut request_body = vec![0u8; content_length];
146 if reader.read_exact(&mut request_body).await.is_err() {
147 return;
148 }
149 captured
150 .lock()
151 .expect("request log")
152 .push(ReceivedRequest { method, path });
153
154 let mut socket = reader.into_inner();
155 let head = if framing.chunked {
156 "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nCache-Control: no-cache\r\nTransfer-Encoding: chunked\r\nConnection: close\r\n\r\n"
157 } else {
158 "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nCache-Control: no-cache\r\nConnection: close\r\n\r\n"
159 };
160 if socket.write_all(head.as_bytes()).await.is_err() {
161 return;
162 }
163 let size = if framing.chunk_bytes == 0 {
164 body.len().max(1)
165 } else {
166 framing.chunk_bytes
167 };
168 for piece in body.chunks(size) {
169 let written = if framing.chunked {
170 let mut frame = format!("{:x}\r\n", piece.len()).into_bytes();
171 frame.extend_from_slice(piece);
172 frame.extend_from_slice(b"\r\n");
173 socket.write_all(&frame).await
174 } else {
175 socket.write_all(piece).await
176 };
177 if written.is_err() || socket.flush().await.is_err() {
178 return;
179 }
180 // Give the client a chance to observe a split line before the rest.
181 tokio::task::yield_now().await;
182 }
183 if framing.hold_open {
184 // A stalled provider for the harness timeout regression. The listener
185 // owns this task, so dropping the replay cancels the open socket too.
186 std::future::pending::<()>().await;
187 }
188 if framing.chunked && !framing.transport_cut {
189 let _ = socket.write_all(b"0\r\n\r\n").await;
190 }
191 let _ = socket.flush().await;
192 let _ = socket.shutdown().await;
193 }
194
195 fn client_for_route(route: &str, base: &str, model: &str) -> CodewhaleClient {
196 let _ = rustls::crypto::ring::default_provider().install_default();
197 let config = match route {
198 "anthropic" => Config {
199 provider: Some("anthropic".to_string()),
200 providers: Some(ProvidersConfig {
201 anthropic: ProviderConfig {
202 api_key: Some("conformance-key".to_string()),
203 base_url: Some(base.to_string()),
204 ..ProviderConfig::default()
205 },
206 ..ProvidersConfig::default()
207 }),
208 ..Config::default()
209 },
210 "openai" => Config {
211 provider: Some("openai".to_string()),
212 providers: Some(ProvidersConfig {
213 openai: ProviderConfig {
214 api_key: Some("conformance-key".to_string()),
215 base_url: Some(format!("{base}/v1")),
216 ..ProviderConfig::default()
217 },
218 ..ProvidersConfig::default()
219 }),
220 ..Config::default()
221 },
222 // Keep DeepSeek's semantic endpoint for route shaping; redirect only
223 // its transport through the existing test seam below.
224 "deepseek" => Config {
225 provider: Some("deepseek".to_string()),
226 ..Config::default()
227 }
228 .with_legacy_root(
229 Some("conformance-key".to_string()),
230 Some("https://api.deepseek.com/v1".to_string()),
231 ),
232 "openai-codex" => Config {
233 provider: Some("openai-codex".to_string()),
234 providers: Some(ProvidersConfig {
235 openai_codex: ProviderConfig {
236 base_url: Some(base.to_string()),
237 ..ProviderConfig::default()
238 },
239 ..ProvidersConfig::default()
240 }),
241 ..Config::default()
242 },
243 other => panic!("sse fixture names unknown route `{other}`"),
244 };
245 // The fixture model and the client must share the production resolver.
246 // Constructing from Config's default model can bind a different wire
247 // format and record an open_failure without exercising the recording.
248 let resolved = crate::route_runtime::resolve_runtime_route(
249 &config,
250 config.active_provider_identity().unwrap().provider,
251 Some(model),
252 )
253 .expect("resolve fixture route");
254 let mut client = CodewhaleClient::from_candidate(&resolved.config, &resolved.candidate)
255 .expect("fixture client");
256 if route == "deepseek" {
257 client.set_test_chat_transport_base_url(base.to_string());
258 }
259 client
260 }
261
262 fn request_for_case(case: &Value) -> MessageRequest {
263 let model = case["model"].as_str().expect("case.model").to_string();
264 let tools = case.get("tools").map(|tools| {
265 serde_json::from_value::<Vec<codewhale_models::Tool>>(tools.clone())
266 .expect("case.tools parse as Tool definitions")
267 });
268 MessageRequest {
269 model,
270 messages: vec![Message {
271 role: Role::User,
272 content: vec![ContentBlock::Text {
273 text: "conformance request".to_string(),
274 cache_control: None,
275 }],
276 }],
277 max_tokens: 512,
278 system: None,
279 tools,
280 tool_choice: None,
281 metadata: None,
282 thinking: None,
283 reasoning_effort: case["reasoning_effort"].as_str().map(str::to_string),
284 stream: Some(true),
285 temperature: None,
286 top_p: None,
287 }
288 }
289
290 async fn replay(case: &Value, recording: Vec<u8>) -> Result<Vec<Value>, String> {
291 replay_with_deadline(case, recording, STREAM_DEADLINE).await
292 }
293
294 async fn replay_with_deadline(
295 case: &Value,
296 recording: Vec<u8>,
297 deadline: Duration,
298 ) -> Result<Vec<Value>, String> {
299 let framing = Framing {
300 chunked: case["framing"].as_str() == Some("chunked"),
301 transport_cut: case["transport_cut"].as_bool().unwrap_or(false),
302 chunk_bytes: case["chunk_bytes"].as_u64().unwrap_or(0) as usize,
303 hold_open: case["hold_open"].as_bool().unwrap_or(false),
304 };
305 let server = LoopbackRecording::start(recording, framing).await;
306 let route = case["route"].as_str().expect("case.route");
307 let client = {
308 // Client construction reads provider credentials from the
309 // environment; only the Codex route needs one, and it is synthetic.
310 let _env = lock_test_env();
311 let _codex = EnvVarGuard::set("OPENAI_CODEX_ACCESS_TOKEN", "conformance-token");
312 let _legacy_codex = EnvVarGuard::remove("CODEX_ACCESS_TOKEN");
313 client_for_route(
314 route,
315 &server.base,
316 case["model"].as_str().expect("case.model"),
317 )
318 };
319
320 let mut lines = Vec::new();
321 let events = golden::complete_within("SSE replay", deadline, async {
322 let mut events = Vec::new();
323 match client.create_message_stream(request_for_case(case)).await {
324 Err(error) => events.push(json!({
325 "type": "open_failure",
326 "detail": format!("{error:#}"),
327 })),
328 Ok(mut stream) => {
329 while let Some(item) = stream.next().await {
330 match item {
331 Ok(event) => events.push(stream_json::to_json(&event)),
332 Err(error) => events.push(json!({
333 "type": "stream_failure",
334 "detail": format!("{error:#}"),
335 })),
336 }
337 }
338 }
339 }
340 events
341 })
342 .await?;
343 if events.is_empty() {
344 return Err("SSE replay produced no event or provider error".to_string());
345 }
346
347 let requests = server.requests.lock().expect("request log").clone();
348 if requests.is_empty()
349 || !events
350 .iter()
351 .any(|event| stream_json::from_json(event).is_ok())
352 {
353 return Err(
354 "SSE recording was not exercised: no loopback request or normalized stream event"
355 .to_string(),
356 );
357 }
358 lines.push(json!({
359 "request": requests
360 .iter()
361 .map(|request| json!({ "method": request.method, "path": request.path }))
362 .collect::<Vec<_>>(),
363 }));
364 lines.extend(events);
365 Ok(lines)
366 }
367
368 #[tokio::test]
369 async fn recorded_sse_streams_match_goldens() {
370 let dir = golden::family_dir(FAMILY);
371 let names = golden::case_names(FAMILY);
372 let mut failures = Failures::default();
373 for name in &names {
374 let case = golden::read_case(FAMILY, name);
375 let recording_name = case["recording"]
376 .as_str()
377 .map_or_else(|| format!("{name}.sse"), str::to_string);
378 let recording = std::fs::read(dir.join(&recording_name))
379 .unwrap_or_else(|error| panic!("read {recording_name}: {error}"));
380 let mut lines = match replay(&case, recording).await {
381 Ok(lines) => lines,
382 Err(error) => {
383 failures.push(name, error);
384 continue;
385 }
386 };
387 let mut masker = golden::Masker::new(&[]);
388 for line in &mut lines {
389 masker.value(line);
390 }
391 // Every event line must be the StreamEvent serde contract itself, so
392 // the events family can feed these goldens straight back in.
393 for line in lines.iter().skip(1) {
394 let kind = line["type"].as_str().unwrap_or_default();
395 if matches!(kind, "stream_failure" | "open_failure") {
396 continue;
397 }
398 if let Err(message) = stream_json::round_trips(line) {
399 failures.push(
400 name,
401 format!("normalized event is not StreamEvent: {message}"),
402 );
403 }
404 }
405 failures.record(
406 name,
407 golden::check_golden(
408 &dir.join(format!("{name}.golden.jsonl")),
409 &golden::jsonl(&lines),
410 ),
411 );
412 }
413 failures.finish(FAMILY, names.len());
414 }
415
415 lines RUST