api_testing_core/websocket/
runner.rs1use 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
30fn 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 _ => 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 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 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}