feat: add websocket upstream adapter
This commit is contained in:
@@ -0,0 +1,442 @@
|
||||
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<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 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<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<bool, WebsocketAdapterError> {
|
||||
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<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);
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,30 @@
|
||||
use serde_json::Value;
|
||||
use thiserror::Error;
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum WebsocketAdapterError {
|
||||
#[error("invalid websocket url: {url}")]
|
||||
InvalidUrl { url: String },
|
||||
#[error("invalid websocket header name {header}")]
|
||||
InvalidHeaderName { header: String },
|
||||
#[error("invalid websocket header value for {header}")]
|
||||
InvalidHeaderValue { header: String },
|
||||
#[error("invalid websocket subprotocol {value}")]
|
||||
InvalidSubprotocol { value: String },
|
||||
#[error("websocket connect failed")]
|
||||
Connect(#[from] tokio_tungstenite::tungstenite::Error),
|
||||
#[error("websocket window expired before collecting any items")]
|
||||
WindowExpired,
|
||||
#[error("websocket stream produced malformed frame payload")]
|
||||
InvalidFramePayload,
|
||||
#[error("websocket endpoint returned close frame before collection completed")]
|
||||
ClosedEarly,
|
||||
#[error("websocket reconnect policy exhausted")]
|
||||
ReconnectExhausted,
|
||||
#[error("websocket upstream returned invalid subscribe payload")]
|
||||
InvalidSubscribePayload,
|
||||
#[error("websocket upstream returned invalid unsubscribe payload")]
|
||||
InvalidUnsubscribePayload,
|
||||
#[error("websocket upstream status {status}")]
|
||||
UnexpectedStatus { status: u16, body: Value },
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
mod client;
|
||||
mod error;
|
||||
mod session;
|
||||
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
|
||||
pub use client::WebsocketAdapter;
|
||||
pub use error::WebsocketAdapterError;
|
||||
pub use session::{HeartbeatPolicy, ReconnectPolicy};
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, Default)]
|
||||
pub struct WebsocketWindowRequest {
|
||||
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
|
||||
pub headers: BTreeMap<String, String>,
|
||||
pub window_duration_ms: u64,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub max_items: Option<u32>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub heartbeat_interval_ms: Option<u64>,
|
||||
#[serde(default)]
|
||||
pub reconnect_max_attempts: u32,
|
||||
#[serde(default)]
|
||||
pub reconnect_backoff_ms: u64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
pub struct WebsocketWindowResponse {
|
||||
pub status_code: u16,
|
||||
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
|
||||
pub headers: BTreeMap<String, String>,
|
||||
pub body: Value,
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
use std::time::Duration;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct ReconnectPolicy {
|
||||
pub max_attempts: u32,
|
||||
pub backoff: Duration,
|
||||
}
|
||||
|
||||
impl Default for ReconnectPolicy {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
max_attempts: 0,
|
||||
backoff: Duration::from_millis(0),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct HeartbeatPolicy {
|
||||
pub interval: Duration,
|
||||
}
|
||||
Reference in New Issue
Block a user