use std::{collections::BTreeMap, time::Duration}; use crank_core::WebsocketTarget; use futures_util::{SinkExt, StreamExt}; use reqwest::header::{HeaderMap, HeaderName, HeaderValue}; use serde_json::Value; use tokio::time::{Instant, sleep}; use tokio_tungstenite::{ MaybeTlsStream, WebSocketStream, connect_async, tungstenite::{ Error as TungsteniteError, Message, client::IntoClientRequest, protocol::{CloseFrame, frame::coding::CloseCode}, }, }; use crate::{ HeartbeatPolicy, ReconnectPolicy, WebsocketAdapterError, WebsocketWindowRequest, WebsocketWindowResponse, }; #[derive(Clone, Debug)] pub struct WebsocketAdapter; impl Default for WebsocketAdapter { fn default() -> Self { Self::new() } } impl WebsocketAdapter { pub fn new() -> Self { Self } pub async fn execute_window( &self, target: &WebsocketTarget, request: &WebsocketWindowRequest, ) -> Result { let started_at = Instant::now(); let deadline = started_at + Duration::from_millis(request.window_duration_ms); let heartbeat = request .heartbeat_interval_ms .map(Duration::from_millis) .map(|interval| crate::HeartbeatPolicy { interval }); let reconnect = ReconnectPolicy { max_attempts: request.reconnect_max_attempts, backoff: Duration::from_millis(request.reconnect_backoff_ms), }; let mut attempts = 0_u32; let mut items = Vec::new(); let mut connected_headers = BTreeMap::new(); loop { let (mut stream, headers) = connect_websocket(target, request).await?; if connected_headers.is_empty() { connected_headers = headers; } send_subscribe_message(&mut stream, target).await?; let completed = collect_window( &mut stream, request.max_items, deadline, heartbeat.as_ref(), &mut items, ) .await?; if completed { send_unsubscribe_message(&mut stream, target).await?; return Ok(WebsocketWindowResponse { status_code: 101, headers: connected_headers, body: serde_json::json!({ "items": items, "done": true, }), }); } if attempts >= reconnect.max_attempts { return Err(WebsocketAdapterError::ReconnectExhausted); } attempts = attempts.saturating_add(1); reconnect_if_needed(&reconnect, attempts).await; } } } type WsStream = WebSocketStream>; pub async fn connect_websocket( target: &WebsocketTarget, request: &WebsocketWindowRequest, ) -> Result<(WsStream, BTreeMap), WebsocketAdapterError> { let mut client_request = target.url.as_str().into_client_request().map_err(|_| { WebsocketAdapterError::InvalidUrl { url: target.url.clone(), } })?; let headers = build_headers(target, request)?; for (name, value) in &headers { client_request .headers_mut() .insert(name.clone(), value.clone()); } if !target.subprotocols.is_empty() { let value = target.subprotocols.join(", "); let header_value = HeaderValue::from_str(&value) .map_err(|_| WebsocketAdapterError::InvalidSubprotocol { value })?; client_request .headers_mut() .insert("Sec-WebSocket-Protocol", header_value); } let (stream, response) = connect_async(client_request).await?; let response_headers = response .headers() .iter() .filter_map(|(name, value)| { value .to_str() .ok() .map(|value| (name.as_str().to_owned(), value.to_owned())) }) .collect(); Ok((stream, response_headers)) } pub async fn send_subscribe_message( stream: &mut WsStream, target: &WebsocketTarget, ) -> Result<(), WebsocketAdapterError> { let Some(payload) = target.subscribe_message_template.as_ref() else { return Ok(()); }; let message = serde_json::to_string(payload) .map_err(|_| WebsocketAdapterError::InvalidSubscribePayload)?; stream.send(Message::Text(message.into())).await?; Ok(()) } pub async fn send_unsubscribe_message( stream: &mut WsStream, target: &WebsocketTarget, ) -> Result<(), WebsocketAdapterError> { let Some(payload) = target.unsubscribe_message_template.as_ref() else { let _ = stream .close(Some(CloseFrame { code: CloseCode::Normal, reason: "completed".into(), })) .await; return Ok(()); }; let message = serde_json::to_string(payload) .map_err(|_| WebsocketAdapterError::InvalidUnsubscribePayload)?; if let Err(error) = stream.send(Message::Text(message.into())).await { if !matches!( error, TungsteniteError::ConnectionClosed | TungsteniteError::AlreadyClosed | TungsteniteError::Protocol( tokio_tungstenite::tungstenite::error::ProtocolError::SendAfterClosing ) ) { return Err(error.into()); } } let _ = stream .close(Some(CloseFrame { code: CloseCode::Normal, reason: "completed".into(), })) .await; Ok(()) } async fn collect_window( stream: &mut WsStream, max_items: Option, deadline: Instant, heartbeat: Option<&HeartbeatPolicy>, items: &mut Vec, ) -> Result { let mut heartbeat_deadline = heartbeat.map(|policy| Instant::now() + policy.interval); loop { if Instant::now() >= deadline { return Ok(!items.is_empty()); } let now = Instant::now(); let next_tick = heartbeat_deadline.unwrap_or(deadline); let sleep_until = std::cmp::min(next_tick, deadline); let wait = sleep_until.saturating_duration_since(now); let timer = sleep(wait); tokio::pin!(timer); tokio::select! { _ = &mut timer => { if heartbeat_deadline.is_some_and(|value| value <= Instant::now()) { heartbeat_tick(stream).await?; heartbeat_deadline = heartbeat.map(|policy| Instant::now() + policy.interval); continue; } return Ok(!items.is_empty()); } frame = read_next_frame(stream) => { match frame? { Some(value) => { items.push(value); if max_items.is_some_and(|limit| items.len() as u32 >= limit) { return Ok(true); } } None => return Ok(!items.is_empty()), } heartbeat_deadline = heartbeat.map(|policy| Instant::now() + policy.interval); } } } } pub async fn read_next_frame( stream: &mut WsStream, ) -> Result, WebsocketAdapterError> { loop { let Some(frame) = stream.next().await else { return Ok(None); }; match frame? { Message::Text(text) => return Ok(Some(decode_text_frame(text.as_ref())?)), Message::Binary(_) => return Err(WebsocketAdapterError::InvalidFramePayload), Message::Ping(payload) => { stream.send(Message::Pong(payload)).await?; } Message::Pong(_) => {} Message::Frame(_) => {} Message::Close(_) => return Ok(None), } } } pub fn decode_text_frame(text: &str) -> Result { serde_json::from_str(text).or_else(|_| Ok(Value::String(text.to_owned()))) } pub async fn heartbeat_tick(stream: &mut WsStream) -> Result<(), WebsocketAdapterError> { stream.send(Message::Ping(Vec::new().into())).await?; Ok(()) } pub async fn reconnect_if_needed(policy: &ReconnectPolicy, attempts: u32) { if attempts == 0 || policy.backoff.is_zero() { return; } sleep(policy.backoff).await; } fn build_headers( target: &WebsocketTarget, request: &WebsocketWindowRequest, ) -> Result { let mut headers = HeaderMap::new(); for (name, value) in target.static_headers.iter().chain(request.headers.iter()) { let header_name = HeaderName::try_from(name.as_str()).map_err(|_| { WebsocketAdapterError::InvalidHeaderName { header: name.clone(), } })?; let header_value = HeaderValue::try_from(value.as_str()).map_err(|_| { WebsocketAdapterError::InvalidHeaderValue { header: name.clone(), } })?; headers.insert(header_name, header_value); } Ok(headers) } #[cfg(test)] mod tests { use std::{collections::BTreeMap, sync::Arc}; use futures_util::{SinkExt, StreamExt}; use serde_json::{Value, json}; use tokio::{net::TcpListener, sync::Mutex}; use tokio_tungstenite::{ accept_hdr_async, tungstenite::handshake::server::{Request, Response}, }; use crate::{WebsocketAdapter, WebsocketWindowRequest}; use crank_core::WebsocketTarget; #[tokio::test] async fn collects_window_messages_and_sends_subscribe_payload() { let received = Arc::new(Mutex::new(Vec::new())); let target_url = spawn_server(received.clone(), false).await; let adapter = WebsocketAdapter::new(); let target = WebsocketTarget { url: target_url, subprotocols: vec!["events.v1".to_owned()], subscribe_message_template: Some(json!({"type":"subscribe","topic":"metrics"})), unsubscribe_message_template: Some(json!({"type":"unsubscribe"})), static_headers: BTreeMap::from([("x-test-env".to_owned(), "ci".to_owned())]), }; let response = adapter .execute_window( &target, &WebsocketWindowRequest { headers: BTreeMap::new(), window_duration_ms: 1_000, max_items: Some(3), heartbeat_interval_ms: None, reconnect_max_attempts: 0, reconnect_backoff_ms: 0, }, ) .await .unwrap(); assert_eq!(response.status_code, 101); assert_eq!(response.body["items"].as_array().unwrap().len(), 3); assert_eq!(response.body["items"][0]["seq"], 1); let received = received.lock().await.clone(); assert!( received .iter() .any(|value| value == &json!({"type":"subscribe","topic":"metrics"})) ); } #[tokio::test] async fn reconnects_when_socket_closes_early() { let received = Arc::new(Mutex::new(Vec::new())); let target_url = spawn_server(received, true).await; let adapter = WebsocketAdapter::new(); let target = WebsocketTarget { url: target_url, subprotocols: Vec::new(), subscribe_message_template: None, unsubscribe_message_template: None, static_headers: BTreeMap::new(), }; let response = adapter .execute_window( &target, &WebsocketWindowRequest { headers: BTreeMap::new(), window_duration_ms: 1_000, max_items: Some(3), heartbeat_interval_ms: None, reconnect_max_attempts: 2, reconnect_backoff_ms: 10, }, ) .await .unwrap(); assert_eq!(response.body["items"].as_array().unwrap().len(), 3); assert_eq!(response.body["items"][2]["seq"], 3); } async fn spawn_server(received: Arc>>, close_early: bool) -> String { let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); let addr = listener.local_addr().unwrap(); tokio::spawn(async move { let mut accepted = 0_u32; loop { let (stream, _) = listener.accept().await.unwrap(); let received = received.clone(); accepted = accepted.saturating_add(1); tokio::spawn(async move { let websocket = accept_hdr_async(stream, |request: &Request, mut response: Response| { if let Some(value) = request.headers().get("x-test-env") { assert_eq!(value, "ci"); } if let Some(value) = request.headers().get("sec-websocket-protocol") { response .headers_mut() .insert("sec-websocket-protocol", value.clone()); } Ok(response) }) .await .unwrap(); let (mut sink, mut source) = websocket.split(); let mut sent = 0_u32; if !close_early { if let Some(message) = source.next().await { let message = message.unwrap(); if let tokio_tungstenite::tungstenite::Message::Text(text) = message { if let Ok(value) = serde_json::from_str::(&text) { received.lock().await.push(value); } } } } let payloads = if close_early && accepted == 1 { Vec::new() } else { vec![json!({"seq": 1}), json!({"seq": 2}), json!({"seq": 3})] }; for payload in payloads { sink.send(tokio_tungstenite::tungstenite::Message::Text( payload.to_string().into(), )) .await .unwrap(); sent += 1; } if sent < 3 { let _ = sink.close().await; } }); } }); format!("ws://{}", addr) } }