feat: add websocket upstream adapter

This commit is contained in:
a.tolmachev
2026-04-06 13:23:45 +03:00
parent b5f80c5d2f
commit 45ea011b7f
27 changed files with 978 additions and 30 deletions
+19
View File
@@ -0,0 +1,19 @@
[package]
name = "crank-adapter-websocket"
edition.workspace = true
license.workspace = true
rust-version.workspace = true
version.workspace = true
[dependencies]
crank-core = { path = "../crank-core" }
futures-util = "0.3"
reqwest.workspace = true
serde.workspace = true
serde_json.workspace = true
thiserror.workspace = true
tokio = { workspace = true, features = ["net", "time"] }
tokio-tungstenite.workspace = true
[dev-dependencies]
tokio = { workspace = true, features = ["macros", "net", "rt-multi-thread", "time"] }
@@ -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 },
}
+35
View File
@@ -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,
}
+2 -1
View File
@@ -30,7 +30,8 @@ pub use observability::{
pub use operation::{
ConfigExport, ExecutionConfig, GeneratedDraft, GeneratedDraftStatus, GraphqlTarget,
GrpcProtocolOptions, GrpcTarget, Operation, OperationStatus, ProtocolOptions, RestTarget,
RetryPolicy, Samples, Target, ToolDescription, ToolExample,
RetryPolicy, Samples, Target, ToolDescription, ToolExample, WebsocketProtocolOptions,
WebsocketTarget,
};
pub use protocol::{AuthKind, ExportMode, GraphqlOperationType, HttpMethod, Protocol};
pub use secret::{Secret, SecretKind, SecretStatus, SecretVersion};
+53 -2
View File
@@ -50,12 +50,26 @@ pub struct GrpcTarget {
pub descriptor_set_b64: String,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct WebsocketTarget {
pub url: String,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub subprotocols: Vec<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub subscribe_message_template: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub unsubscribe_message_template: Option<Value>,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub static_headers: BTreeMap<String, String>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum Target {
Rest(RestTarget),
Graphql(GraphqlTarget),
Grpc(GrpcTarget),
Websocket(WebsocketTarget),
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize, Default)]
@@ -68,10 +82,22 @@ pub struct GrpcProtocolOptions {
pub use_tls: bool,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize, Default)]
pub struct WebsocketProtocolOptions {
#[serde(skip_serializing_if = "Option::is_none")]
pub heartbeat_interval_ms: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub reconnect_max_attempts: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub reconnect_backoff_ms: Option<u64>,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize, Default)]
pub struct ProtocolOptions {
#[serde(skip_serializing_if = "Option::is_none")]
pub grpc: Option<GrpcProtocolOptions>,
#[serde(skip_serializing_if = "Option::is_none")]
pub websocket: Option<WebsocketProtocolOptions>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
@@ -209,6 +235,7 @@ mod tests {
operation::{
ConfigExport, ExecutionConfig, GraphqlTarget, Operation, OperationStatus,
ProtocolOptions, RestTarget, Samples, Target, ToolDescription, ToolExample,
WebsocketProtocolOptions, WebsocketTarget,
},
protocol::{AuthKind, ExportMode, GraphqlOperationType, HttpMethod, Protocol},
};
@@ -244,6 +271,23 @@ mod tests {
assert_eq!(value["operation_type"], "mutation");
}
#[test]
fn websocket_target_serializes_templates() {
let target = Target::Websocket(WebsocketTarget {
url: "wss://events.example.com/stream".to_owned(),
subprotocols: vec!["graphql-transport-ws".to_owned()],
subscribe_message_template: Some(json!({ "type": "subscribe" })),
unsubscribe_message_template: Some(json!({ "type": "unsubscribe" })),
static_headers: BTreeMap::from([("x-env".to_owned(), "test".to_owned())]),
});
let value = serde_json::to_value(target).unwrap();
assert_eq!(value["kind"], "websocket");
assert_eq!(value["subprotocols"][0], "graphql-transport-ws");
assert_eq!(value["subscribe_message_template"]["type"], "subscribe");
}
#[test]
fn operation_exposes_local_domain_helpers() {
let operation = Operation {
@@ -269,7 +313,14 @@ mod tests {
retry_policy: None,
auth_profile_ref: Some(AuthProfileId::new("auth_01")),
headers: BTreeMap::new(),
protocol_options: Some(ProtocolOptions::default()),
protocol_options: Some(ProtocolOptions {
grpc: None,
websocket: Some(WebsocketProtocolOptions {
heartbeat_interval_ms: Some(5_000),
reconnect_max_attempts: Some(3),
reconnect_backoff_ms: Some(250),
}),
}),
streaming: None,
},
tool_description: ToolDescription {
+19
View File
@@ -8,6 +8,7 @@ pub enum Protocol {
Rest,
Graphql,
Grpc,
Websocket,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
@@ -49,6 +50,7 @@ impl Protocol {
Self::Rest => true,
Self::Graphql => matches!(mode, ExecutionMode::Unary),
Self::Grpc => true,
Self::Websocket => !matches!(mode, ExecutionMode::Unary),
}
}
@@ -57,6 +59,7 @@ impl Protocol {
Self::Rest => true,
Self::Graphql => matches!(behavior, TransportBehavior::RequestResponse),
Self::Grpc => !matches!(behavior, TransportBehavior::DeferredResult),
Self::Websocket => !matches!(behavior, TransportBehavior::RequestResponse),
}
}
}
@@ -79,4 +82,20 @@ mod tests {
assert!(Protocol::Grpc.supports_execution_mode(ExecutionMode::Session));
assert!(!Protocol::Grpc.supports_transport_behavior(TransportBehavior::DeferredResult));
}
#[test]
fn websocket_requires_streaming_modes() {
assert!(!Protocol::Websocket.supports_execution_mode(ExecutionMode::Unary));
assert!(Protocol::Websocket.supports_execution_mode(ExecutionMode::Window));
assert!(Protocol::Websocket.supports_execution_mode(ExecutionMode::Session));
assert!(Protocol::Websocket.supports_execution_mode(ExecutionMode::AsyncJob));
assert!(
!Protocol::Websocket.supports_transport_behavior(TransportBehavior::RequestResponse)
);
assert!(Protocol::Websocket.supports_transport_behavior(TransportBehavior::ServerStream));
assert!(
Protocol::Websocket.supports_transport_behavior(TransportBehavior::StatefulSession)
);
assert!(Protocol::Websocket.supports_transport_behavior(TransportBehavior::DeferredResult));
}
}
+1
View File
@@ -3958,6 +3958,7 @@ fn target_summary(target: &Target) -> (String, String) {
.to_owned(),
),
Target::Grpc(grpc) => (grpc.server_addr.clone(), grpc.method.clone()),
Target::Websocket(websocket) => (websocket.url.clone(), "SUBSCRIBE".to_owned()),
}
}
+2
View File
@@ -9,6 +9,7 @@ version.workspace = true
crank-adapter-graphql = { path = "../crank-adapter-graphql" }
crank-adapter-grpc = { path = "../crank-adapter-grpc" }
crank-adapter-rest = { path = "../crank-adapter-rest" }
crank-adapter-websocket = { path = "../crank-adapter-websocket" }
crank-core = { path = "../crank-core" }
crank-mapping = { path = "../crank-mapping" }
crank-schema = { path = "../crank-schema" }
@@ -21,3 +22,4 @@ axum.workspace = true
crank-adapter-grpc = { path = "../crank-adapter-grpc", features = ["test-support"] }
futures-util = "0.3"
tokio.workspace = true
tokio-tungstenite.workspace = true
+3
View File
@@ -1,6 +1,7 @@
use crank_adapter_graphql::GraphqlAdapterError;
use crank_adapter_grpc::GrpcAdapterError;
use crank_adapter_rest::RestAdapterError;
use crank_adapter_websocket::WebsocketAdapterError;
use crank_core::{ExecutionMode, Protocol};
use crank_mapping::MappingError;
use crank_schema::SchemaError;
@@ -18,6 +19,8 @@ pub enum RuntimeError {
GrpcAdapter(#[from] GrpcAdapterError),
#[error(transparent)]
RestAdapter(#[from] RestAdapterError),
#[error(transparent)]
WebsocketAdapter(#[from] WebsocketAdapterError),
#[error("protocol {protocol:?} is not supported by runtime")]
UnsupportedProtocol { protocol: Protocol },
#[error("operation {operation_id} does not define streaming config")]
+211 -5
View File
@@ -3,6 +3,7 @@ use std::collections::BTreeMap;
use crank_adapter_graphql::{GraphqlAdapter, GraphqlRequest};
use crank_adapter_grpc::{GrpcAdapter, GrpcRequest, GrpcWindowRequest};
use crank_adapter_rest::{RestAdapter, RestRequest, RestWindowRequest};
use crank_adapter_websocket::{WebsocketAdapter, WebsocketWindowRequest};
use crank_core::{ExecutionMode, Target, TransportBehavior};
use serde_json::{Map, Value, json};
@@ -15,6 +16,7 @@ pub struct RuntimeExecutor {
graphql_adapter: GraphqlAdapter,
grpc_adapter: GrpcAdapter,
rest_adapter: RestAdapter,
websocket_adapter: WebsocketAdapter,
}
impl Default for RuntimeExecutor {
@@ -29,6 +31,7 @@ impl RuntimeExecutor {
graphql_adapter: GraphqlAdapter::new(),
grpc_adapter: GrpcAdapter::new(),
rest_adapter: RestAdapter::new(),
websocket_adapter: WebsocketAdapter::new(),
}
}
@@ -63,8 +66,10 @@ impl RuntimeExecutor {
let adapter_response = if matches!(
streaming.transport_behavior,
TransportBehavior::ServerStream
) && matches!(operation.target, Target::Rest(_) | Target::Grpc(_))
{
) && matches!(
operation.target,
Target::Rest(_) | Target::Grpc(_) | Target::Websocket(_)
) {
self.execute_window_adapter(operation, prepared_request)
.await?
} else {
@@ -200,6 +205,10 @@ impl RuntimeExecutor {
data: Value::Null,
})
}
Target::Websocket(_) => Err(RuntimeError::UnsupportedExecutionMode {
operation_id: operation.operation_id.as_str().to_owned(),
mode: ExecutionMode::Unary,
}),
}
}
@@ -274,6 +283,46 @@ impl RuntimeExecutor {
data: Value::Null,
})
}
Target::Websocket(target) => {
let Some(streaming) = operation.execution_config.streaming.as_ref() else {
return Err(RuntimeError::MissingStreamingConfig {
operation_id: operation.operation_id.as_str().to_owned(),
});
};
let websocket_options = operation
.execution_config
.protocol_options
.as_ref()
.and_then(|options| options.websocket.as_ref());
let request = WebsocketWindowRequest {
headers: merge_headers(
&BTreeMap::new(),
&operation.execution_config.headers,
&prepared_request.headers,
),
window_duration_ms: streaming.window_duration_ms.unwrap_or_default(),
max_items: streaming.max_items.map(|value| value.saturating_add(1)),
heartbeat_interval_ms: websocket_options
.and_then(|options| options.heartbeat_interval_ms),
reconnect_max_attempts: websocket_options
.and_then(|options| options.reconnect_max_attempts)
.unwrap_or_default(),
reconnect_backoff_ms: websocket_options
.and_then(|options| options.reconnect_backoff_ms)
.unwrap_or_default(),
};
let response = self
.websocket_adapter
.execute_window(target, &request)
.await?;
Ok(AdapterResponse {
status_code: response.status_code,
headers: response.headers,
body: response.body,
data: Value::Null,
})
}
_ => self.execute_adapter(operation, prepared_request).await,
}
}
@@ -391,14 +440,16 @@ mod tests {
use crank_core::{
AggregationMode, DescriptorId, ExecutionConfig, ExecutionMode, GeneratedDraft,
GeneratedDraftStatus, GraphqlOperationType, GraphqlTarget, GrpcTarget, HttpMethod,
Operation, OperationId, OperationStatus, Protocol, RestTarget, Samples, StreamingConfig,
Target, ToolDescription, ToolExample, ToolFamilyConfig, TransportBehavior,
Operation, OperationId, OperationStatus, Protocol, ProtocolOptions, RestTarget, Samples,
StreamingConfig, Target, ToolDescription, ToolExample, ToolFamilyConfig, TransportBehavior,
WebsocketProtocolOptions, WebsocketTarget,
};
use crank_mapping::{MappingRule, MappingSet};
use crank_schema::{Schema, SchemaKind};
use futures_util::stream;
use futures_util::{SinkExt, StreamExt, stream};
use serde_json::{Value, json};
use tokio::net::TcpListener;
use tokio_tungstenite::accept_async;
use crate::{RuntimeError, RuntimeExecutor, RuntimeOperation};
@@ -464,6 +515,26 @@ mod tests {
assert!(!result.window_complete);
}
#[tokio::test]
async fn executes_websocket_window_mode_end_to_end() {
let target_url = spawn_websocket_server().await;
let executor = RuntimeExecutor::new();
let operation =
test_websocket_window_operation(&target_url, AggregationMode::RawItems, Some(3));
let result = executor
.execute_window(&operation, &json!({}))
.await
.unwrap();
assert_eq!(result.items.len(), 3);
assert_eq!(result.items[0], json!({ "seq": 1, "value": 101 }));
assert_eq!(result.items[2], json!({ "seq": 3, "value": 103 }));
assert!(result.window_complete);
assert!(!result.truncated);
assert!(!result.has_more);
}
#[tokio::test]
async fn rejects_invalid_input_shape() {
let base_url = spawn_runtime_server().await;
@@ -652,6 +723,43 @@ mod tests {
format!("http://{}", address)
}
async fn spawn_websocket_server() -> String {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
tokio::spawn(async move {
loop {
let (stream, _) = listener.accept().await.unwrap();
tokio::spawn(async move {
let websocket = accept_async(stream).await.unwrap();
let (mut sink, mut source) = websocket.split();
if let Some(message) = source.next().await {
let message = message.unwrap();
if !matches!(message, tokio_tungstenite::tungstenite::Message::Text(_)) {
return;
}
}
for payload in [
json!({ "seq": 1, "value": 101 }),
json!({ "seq": 2, "value": 102 }),
json!({ "seq": 3, "value": 103 }),
] {
sink.send(tokio_tungstenite::tungstenite::Message::Text(
payload.to_string().into(),
))
.await
.unwrap();
}
let _ = sink.close().await;
});
}
});
format!("ws://{}", address)
}
async fn create_lead(Json(payload): Json<Value>) -> (axum::http::StatusCode, Json<Value>) {
let should_fail = payload
.get("fail")
@@ -1168,6 +1276,104 @@ mod tests {
operation
}
fn test_websocket_window_operation(
target_url: &str,
aggregation_mode: AggregationMode,
max_items: Option<u32>,
) -> RuntimeOperation {
RuntimeOperation::from(Operation {
id: OperationId::new("op_websocket_window_runtime"),
name: "telemetry_window_ws".to_owned(),
display_name: "Telemetry Window WebSocket".to_owned(),
category: "ops".to_owned(),
protocol: Protocol::Websocket,
status: OperationStatus::Published,
version: 1,
target: Target::Websocket(WebsocketTarget {
url: target_url.to_owned(),
subprotocols: Vec::new(),
subscribe_message_template: Some(json!({"type":"subscribe","topic":"telemetry"})),
unsubscribe_message_template: Some(
json!({"type":"unsubscribe","topic":"telemetry"}),
),
static_headers: BTreeMap::new(),
}),
input_schema: Schema {
kind: SchemaKind::Object,
description: None,
required: true,
nullable: false,
default_value: None,
fields: BTreeMap::new(),
items: None,
enum_values: Vec::new(),
variants: Vec::new(),
},
output_schema: object_schema("ignored", SchemaKind::String),
input_mapping: MappingSet {
rules: vec![MappingRule {
source: "$.mcp".to_owned(),
target: "$.request.body".to_owned(),
required: false,
default_value: None,
transform: None,
condition: None,
notes: None,
}],
},
output_mapping: MappingSet { rules: Vec::new() },
execution_config: ExecutionConfig {
timeout_ms: 1_000,
retry_policy: None,
auth_profile_ref: None,
headers: BTreeMap::new(),
protocol_options: Some(ProtocolOptions {
grpc: None,
websocket: Some(WebsocketProtocolOptions {
heartbeat_interval_ms: Some(100),
reconnect_max_attempts: Some(1),
reconnect_backoff_ms: Some(10),
}),
}),
streaming: Some(StreamingConfig {
mode: ExecutionMode::Window,
transport_behavior: TransportBehavior::ServerStream,
window_duration_ms: Some(2_000),
poll_interval_ms: None,
upstream_timeout_ms: Some(1_000),
idle_timeout_ms: None,
max_session_lifetime_ms: None,
max_items,
max_bytes: None,
aggregation_mode,
summary_path: None,
items_path: Some("$.items".to_owned()),
cursor_path: None,
status_path: None,
done_path: Some("$.done".to_owned()),
redacted_paths: Vec::new(),
truncate_item_fields: false,
max_field_length: None,
drop_duplicates: false,
sampling_rate: None,
tool_family: ToolFamilyConfig::default(),
}),
},
tool_description: ToolDescription {
title: "Telemetry Window WebSocket".to_owned(),
description: "Collects bounded WebSocket telemetry.".to_owned(),
tags: vec!["websocket".to_owned(), "stream".to_owned()],
examples: Vec::new(),
},
samples: None,
generated_draft: None,
config_export: None,
created_at: "2026-04-06T12:00:00Z".to_owned(),
updated_at: "2026-04-06T12:00:00Z".to_owned(),
published_at: Some("2026-04-06T12:00:00Z".to_owned()),
})
}
fn test_window_sse_operation(
base_url: &str,
aggregation_mode: AggregationMode,