Skip to main content

api_testing_core/websocket/
runner.rs

1use std::net::TcpStream;
2use std::sync::mpsc;
3use std::thread;
4use std::time::Duration;
5
6use anyhow::Context;
7use serde::Serialize;
8use tungstenite::WebSocket;
9use tungstenite::client::IntoClientRequest;
10use tungstenite::http::{HeaderName, HeaderValue};
11use tungstenite::stream::MaybeTlsStream;
12use tungstenite::{Message, connect};
13
14use crate::Result;
15use crate::websocket::schema::{WebsocketExpect, WebsocketRequestFile, WebsocketStep};
16
17#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
18pub struct WebsocketTranscriptEntry {
19    pub direction: String,
20    pub payload: String,
21}
22
23#[derive(Debug, Clone, PartialEq, Eq)]
24pub struct WebsocketExecutedRequest {
25    pub target: String,
26    pub transcript: Vec<WebsocketTranscriptEntry>,
27    pub last_received: Option<String>,
28}
29
30/// Apply a read timeout to the underlying TCP socket so per-step `read()`
31/// calls cannot hang forever. Pass `None` to clear the timeout.
32fn set_socket_read_timeout(
33    socket: &mut WebSocket<MaybeTlsStream<TcpStream>>,
34    timeout: Option<Duration>,
35) -> std::io::Result<()> {
36    let tcp = match socket.get_mut() {
37        MaybeTlsStream::Plain(stream) => stream,
38        MaybeTlsStream::Rustls(stream) => stream.get_mut(),
39        // Forward-compat: other variants (e.g. NativeTls if a future feature is
40        // enabled) do not expose a TcpStream we can configure here.
41        _ => return Ok(()),
42    };
43    tcp.set_read_timeout(timeout)
44}
45
46fn parse_message_text(message: Message) -> String {
47    match message {
48        Message::Text(t) => t.to_string(),
49        Message::Binary(b) => String::from_utf8_lossy(&b).to_string(),
50        Message::Ping(b) => format!("<PING:{}>", String::from_utf8_lossy(&b)),
51        Message::Pong(b) => format!("<PONG:{}>", String::from_utf8_lossy(&b)),
52        Message::Close(frame) => match frame {
53            Some(f) => format!("<CLOSE:{}:{}>", f.code, f.reason),
54            None => "<CLOSE>".to_string(),
55        },
56        Message::Frame(_) => "<FRAME>".to_string(),
57    }
58}
59
60fn apply_expect(expect: Option<&WebsocketExpect>, text: &str, path: &str) -> Result<()> {
61    if let Some(expect) = expect {
62        crate::websocket::expect::evaluate_text_expect(expect, text, path)?;
63    }
64    Ok(())
65}
66
67pub fn execute_websocket_request(
68    request_file: &WebsocketRequestFile,
69    target_override: &str,
70    bearer_token: Option<&str>,
71) -> Result<WebsocketExecutedRequest> {
72    let target = if !target_override.trim().is_empty() {
73        target_override.trim().to_string()
74    } else if let Some(url) = request_file.request.url.as_deref() {
75        url.to_string()
76    } else {
77        anyhow::bail!("websocket target URL is empty (set request.url or pass --url/--env)");
78    };
79
80    let mut request = target
81        .as_str()
82        .into_client_request()
83        .context("invalid websocket target URL")?;
84
85    for (key, value) in &request_file.request.headers {
86        let header_name = HeaderName::from_bytes(key.as_bytes())
87            .with_context(|| format!("invalid websocket header name: {key}"))?;
88        let header_value = HeaderValue::from_str(value)
89            .with_context(|| format!("invalid websocket header value for {key}"))?;
90        request.headers_mut().insert(header_name, header_value);
91    }
92
93    if let Some(token) = bearer_token {
94        let has_auth = request_file
95            .request
96            .headers
97            .iter()
98            .any(|(k, _)| k.eq_ignore_ascii_case("authorization"));
99        if !has_auth {
100            request.headers_mut().insert(
101                HeaderName::from_static("authorization"),
102                HeaderValue::from_str(&format!("Bearer {token}"))
103                    .context("invalid bearer token for Authorization header")?,
104            );
105        }
106    }
107
108    let (mut socket, _resp) = match request_file.request.connect_timeout_seconds {
109        Some(secs) => {
110            let (tx, rx) = mpsc::channel();
111            thread::spawn(move || {
112                let _ = tx.send(connect(request));
113            });
114            match rx.recv_timeout(Duration::from_secs(secs)) {
115                Ok(Ok(pair)) => pair,
116                Ok(Err(err)) => {
117                    return Err(anyhow::Error::new(err))
118                        .with_context(|| format!("failed to connect websocket target '{target}'"));
119                }
120                Err(_) => {
121                    anyhow::bail!(
122                        "websocket connect timed out after {secs}s for target '{target}'"
123                    );
124                }
125            }
126        }
127        None => connect(request)
128            .with_context(|| format!("failed to connect websocket target '{target}'"))?,
129    };
130
131    let mut transcript: Vec<WebsocketTranscriptEntry> = Vec::new();
132    let mut last_received: Option<String> = None;
133
134    for (idx, step) in request_file.request.steps.iter().enumerate() {
135        match step {
136            WebsocketStep::Send { text } => {
137                socket
138                    .send(Message::Text(text.clone().into()))
139                    .with_context(|| format!("websocket send failed at step {idx}"))?;
140                transcript.push(WebsocketTranscriptEntry {
141                    direction: "send".to_string(),
142                    payload: text.clone(),
143                });
144            }
145            WebsocketStep::Receive {
146                timeout_seconds,
147                expect,
148            } => {
149                let timeout = timeout_seconds.map(Duration::from_secs);
150                if let Some(dur) = timeout {
151                    set_socket_read_timeout(&mut socket, Some(dur)).with_context(|| {
152                        format!("failed to set read timeout for websocket step {idx}")
153                    })?;
154                }
155                let read_result = socket.read();
156                if timeout.is_some() {
157                    let _ = set_socket_read_timeout(&mut socket, None);
158                }
159                let message = read_result.with_context(|| match timeout_seconds {
160                    Some(secs) => {
161                        format!("websocket receive timed out after {secs}s at step {idx}")
162                    }
163                    None => format!("websocket receive failed at step {idx}"),
164                })?;
165                let text = parse_message_text(message);
166                apply_expect(
167                    expect.as_ref(),
168                    &text,
169                    &format!("websocket steps[{idx}].expect"),
170                )?;
171                transcript.push(WebsocketTranscriptEntry {
172                    direction: "receive".to_string(),
173                    payload: text.clone(),
174                });
175                last_received = Some(text);
176            }
177            WebsocketStep::Close => {
178                let _ = socket.close(None);
179                transcript.push(WebsocketTranscriptEntry {
180                    direction: "close".to_string(),
181                    payload: String::new(),
182                });
183            }
184        }
185    }
186
187    Ok(WebsocketExecutedRequest {
188        target,
189        transcript,
190        last_received,
191    })
192}
193
194#[cfg(test)]
195mod tests {
196    use std::net::TcpListener;
197    use std::thread;
198
199    use pretty_assertions::assert_eq;
200    use tempfile::TempDir;
201    use tungstenite::Message;
202
203    use super::*;
204    use crate::websocket::schema::WebsocketRequestFile;
205
206    fn spawn_echo_server() -> (String, thread::JoinHandle<()>) {
207        let listener = TcpListener::bind("127.0.0.1:0").expect("bind websocket listener");
208        let addr = listener.local_addr().expect("listener addr");
209
210        let handle = thread::spawn(move || {
211            let (stream, _) = listener.accept().expect("accept websocket stream");
212            let mut ws = tungstenite::accept(stream).expect("accept websocket handshake");
213            loop {
214                match ws.read() {
215                    Ok(Message::Text(text)) => {
216                        let response = if text.trim() == "ping" {
217                            "{\"ok\":true}".to_string()
218                        } else {
219                            text.to_string()
220                        };
221                        ws.send(Message::Text(response.into()))
222                            .expect("send response");
223                    }
224                    Ok(Message::Close(_)) => {
225                        let _ = ws.close(None);
226                        break;
227                    }
228                    Ok(_) => {}
229                    Err(_) => break,
230                }
231            }
232        });
233
234        (format!("ws://{addr}"), handle)
235    }
236
237    fn spawn_silent_server() -> (String, thread::JoinHandle<()>) {
238        let listener = TcpListener::bind("127.0.0.1:0").expect("bind silent listener");
239        let addr = listener.local_addr().expect("listener addr");
240
241        let handle = thread::spawn(move || {
242            let (stream, _) = listener.accept().expect("accept stream");
243            let mut ws = tungstenite::accept(stream).expect("accept handshake on silent server");
244            // Block on read until the client closes the connection — never
245            // sending anything ourselves so the client read times out.
246            let _ = ws.read();
247            let _ = ws.close(None);
248        });
249
250        (format!("ws://{addr}"), handle)
251    }
252
253    #[test]
254    fn websocket_runner_step_receive_timeout_surfaces_error() {
255        let tmp = TempDir::new().expect("tmp");
256        let request_path = tmp.path().join("silent.ws.json");
257
258        let (url, handle) = spawn_silent_server();
259
260        std::fs::write(
261            &request_path,
262            serde_json::to_vec_pretty(&serde_json::json!({
263                "url": url,
264                "steps": [
265                    {"type": "receive", "timeoutSeconds": 1}
266                ]
267            }))
268            .expect("serialize request"),
269        )
270        .expect("write request");
271
272        let loaded = WebsocketRequestFile::load(&request_path).expect("load request");
273        let started = std::time::Instant::now();
274        let err =
275            execute_websocket_request(&loaded, "", None).expect_err("expected receive to time out");
276        let elapsed = started.elapsed();
277
278        let msg = format!("{err:#}");
279        assert!(
280            msg.contains("timed out after 1s"),
281            "expected timeout message; got: {msg}"
282        );
283        // Cushion for OS scheduling and TLS handshake — keep well under the
284        // default integration-test budget.
285        assert!(
286            elapsed < Duration::from_secs(5),
287            "expected timeout to fire within 5s; got {elapsed:?}"
288        );
289
290        let _ = handle.join();
291    }
292
293    #[test]
294    fn websocket_runner_executes_send_receive_steps() {
295        let tmp = TempDir::new().expect("tmp");
296        let request_path = tmp.path().join("echo.ws.json");
297
298        let (url, handle) = spawn_echo_server();
299
300        std::fs::write(
301            &request_path,
302            serde_json::to_vec_pretty(&serde_json::json!({
303                "url": url,
304                "steps": [
305                    {"type": "send", "text": "ping"},
306                    {"type": "receive", "expect": {"jq": ".ok == true"}},
307                    {"type": "close"}
308                ]
309            }))
310            .expect("serialize request"),
311        )
312        .expect("write request");
313
314        let loaded = WebsocketRequestFile::load(&request_path).expect("load request");
315        let executed = execute_websocket_request(&loaded, "", None).expect("execute websocket");
316
317        assert_eq!(executed.transcript.len(), 3);
318        assert_eq!(executed.transcript[0].direction, "send");
319        assert_eq!(executed.transcript[1].direction, "receive");
320        assert_eq!(executed.last_received.as_deref(), Some("{\"ok\":true}"));
321
322        handle.join().expect("join websocket server");
323    }
324}