Files
crank/crates/crank-adapter-websocket/src/client.rs
T
2026-04-11 01:37:24 +03:00

538 lines
18 KiB
Rust

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,
};
enum WindowCollectionStatus {
WindowExpired,
MaxItemsReached,
}
#[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<WebsocketWindowResponse, WebsocketAdapterError> {
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 status = match collect_window(
&mut stream,
request.max_items,
deadline,
heartbeat.as_ref(),
&mut items,
)
.await
{
Ok(status) => status,
Err(WebsocketAdapterError::ClosedEarly) => {
if attempts >= reconnect.max_attempts {
return Err(WebsocketAdapterError::ReconnectExhausted);
}
attempts = attempts.saturating_add(1);
reconnect_if_needed(&reconnect, attempts).await;
continue;
}
Err(error) => return Err(error),
};
match status {
WindowCollectionStatus::WindowExpired => {
send_unsubscribe_message(&mut stream, target).await?;
return Ok(WebsocketWindowResponse {
status_code: 101,
headers: connected_headers,
body: serde_json::json!({
"items": items,
"done": false,
}),
});
}
WindowCollectionStatus::MaxItemsReached => {
send_unsubscribe_message(&mut stream, target).await?;
return Ok(WebsocketWindowResponse {
status_code: 101,
headers: connected_headers,
body: serde_json::json!({
"items": items,
"done": true,
}),
});
}
}
}
}
}
type WsStream = WebSocketStream<MaybeTlsStream<tokio::net::TcpStream>>;
pub async fn connect_websocket(
target: &WebsocketTarget,
request: &WebsocketWindowRequest,
) -> Result<(WsStream, BTreeMap<String, String>), 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<u32>,
deadline: Instant,
heartbeat: Option<&HeartbeatPolicy>,
items: &mut Vec<Value>,
) -> Result<WindowCollectionStatus, WebsocketAdapterError> {
let mut heartbeat_deadline = heartbeat.map(|policy| Instant::now() + policy.interval);
loop {
if Instant::now() >= deadline {
return Ok(WindowCollectionStatus::WindowExpired);
}
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(WindowCollectionStatus::WindowExpired);
}
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(WindowCollectionStatus::MaxItemsReached);
}
}
None => return Err(WebsocketAdapterError::ClosedEarly),
}
heartbeat_deadline = heartbeat.map(|policy| Instant::now() + policy.interval);
}
}
}
}
pub async fn read_next_frame(
stream: &mut WsStream,
) -> Result<Option<Value>, 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<Value, WebsocketAdapterError> {
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<HeaderMap, WebsocketAdapterError> {
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);
assert_eq!(response.body["done"], true);
}
#[tokio::test]
async fn reconnects_after_partial_close_without_marking_done_early() {
let target_url = spawn_partial_close_server().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"][0]["seq"], 1);
assert_eq!(response.body["items"][2]["seq"], 3);
assert_eq!(response.body["done"], true);
}
async fn spawn_server(received: Arc<Mutex<Vec<Value>>>, 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::<Value>(&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)
}
async fn spawn_partial_close_server() -> 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();
accepted = accepted.saturating_add(1);
tokio::spawn(async move {
let websocket =
accept_hdr_async(stream, |_request: &Request, response: Response| {
Ok(response)
})
.await
.unwrap();
let (mut sink, _source) = websocket.split();
let payloads = if accepted == 1 {
vec![json!({"seq": 1})]
} else {
vec![json!({"seq": 2}), json!({"seq": 3})]
};
for payload in payloads {
sink.send(tokio_tungstenite::tungstenite::Message::Text(
payload.to_string().into(),
))
.await
.unwrap();
}
let _ = sink.close().await;
});
}
});
format!("ws://{}", addr)
}
}