soma_mcp_client/upstream/transport/
websocket.rs1use futures::{SinkExt, StreamExt};
2use rmcp::service::{RxJsonRpcMessage, TxJsonRpcMessage};
3use rmcp::transport::worker::{Worker, WorkerConfig, WorkerContext, WorkerQuitReason};
4use rmcp::{transport::worker::WorkerTransport, RoleClient};
5use tokio_tungstenite::connect_async_with_config;
6use tokio_tungstenite::tungstenite::client::IntoClientRequest;
7use tokio_tungstenite::tungstenite::protocol::{Message, WebSocketConfig};
8use tokio_tungstenite::tungstenite::{self};
9
10const DEFAULT_MAX_MESSAGE_SIZE: usize = 10 * 1024 * 1024;
11const DEFAULT_MAX_FRAME_SIZE: usize = 128 * 1024;
12
13#[derive(Debug, Clone, thiserror::Error)]
14pub enum WebSocketTransportError {
15 #[error("{0}")]
16 Message(String),
17}
18
19impl WebSocketTransportError {
20 fn new(message: impl Into<String>) -> Self {
21 Self::Message(message.into())
22 }
23}
24
25#[derive(Debug, Clone)]
26pub struct WebSocketTransportConfig {
27 pub url: String,
28 pub authorization: Option<String>,
29 pub max_message_size: usize,
30 pub max_frame_size: usize,
31}
32
33impl WebSocketTransportConfig {
34 #[must_use]
35 pub fn new(url: impl Into<String>) -> Self {
36 Self {
37 url: url.into(),
38 authorization: None,
39 max_message_size: DEFAULT_MAX_MESSAGE_SIZE,
40 max_frame_size: DEFAULT_MAX_FRAME_SIZE,
41 }
42 }
43
44 #[must_use]
45 pub fn with_authorization(mut self, authorization: Option<String>) -> Self {
46 self.authorization = authorization;
47 self
48 }
49}
50
51#[derive(Debug)]
52pub struct WebSocketClientWorker {
53 config: WebSocketTransportConfig,
54}
55
56impl WebSocketClientWorker {
57 #[must_use]
58 pub fn new(config: WebSocketTransportConfig) -> Self {
59 Self { config }
60 }
61}
62
63impl Worker for WebSocketClientWorker {
64 type Error = WebSocketTransportError;
65 type Role = RoleClient;
66
67 fn err_closed() -> Self::Error {
68 WebSocketTransportError::new("websocket transport is closed")
69 }
70
71 fn err_join(error: tokio::task::JoinError) -> Self::Error {
72 WebSocketTransportError::new(format!("websocket transport task failed: {error}"))
73 }
74
75 fn config(&self) -> WorkerConfig {
76 let mut config = WorkerConfig::default();
77 config.name = Some("upstream-websocket-client".to_owned());
78 config.channel_buffer_capacity = 32;
79 config
80 }
81
82 async fn run(
83 self,
84 mut context: WorkerContext<Self>,
85 ) -> Result<(), WorkerQuitReason<Self::Error>> {
86 let mut request = self
87 .config
88 .url
89 .clone()
90 .into_client_request()
91 .map_err(|error| {
92 WorkerQuitReason::fatal(
93 WebSocketTransportError::new(format!("invalid websocket request: {error}")),
94 "build websocket request",
95 )
96 })?;
97 if let Some(authorization) = &self.config.authorization {
98 let header =
99 tungstenite::http::HeaderValue::from_str(authorization).map_err(|error| {
100 WorkerQuitReason::fatal(
101 WebSocketTransportError::new(format!(
102 "invalid websocket authorization header: {error}"
103 )),
104 "build websocket authorization header",
105 )
106 })?;
107 request
108 .headers_mut()
109 .insert(tungstenite::http::header::AUTHORIZATION, header);
110 }
111
112 let mut websocket_config = WebSocketConfig::default();
113 websocket_config.max_message_size = Some(self.config.max_message_size);
114 websocket_config.max_frame_size = Some(self.config.max_frame_size);
115 websocket_config.accept_unmasked_frames = false;
116 let (socket, _) = connect_async_with_config(request, Some(websocket_config), false)
117 .await
118 .map_err(|error| {
119 WorkerQuitReason::fatal(
120 WebSocketTransportError::new(format!("websocket connect failed: {error}")),
121 "connect websocket upstream",
122 )
123 })?;
124 let (mut writer, mut reader) = socket.split();
125 let cancellation = context.cancellation_token.clone();
126
127 loop {
128 tokio::select! {
129 _ = cancellation.cancelled() => {
130 drop(writer.send(Message::Close(None)).await);
131 return Err(WorkerQuitReason::Cancelled);
132 }
133 inbound = reader.next() => match inbound {
134 Some(Ok(Message::Text(text))) => {
135 let message = decode_server_message(text.as_str()).map_err(|error| {
136 WorkerQuitReason::fatal(error, "decode websocket frame")
137 })?;
138 context.send_to_handler(message).await?;
139 }
140 Some(Ok(Message::Binary(_))) => {
141 return Err(WorkerQuitReason::fatal(
142 WebSocketTransportError::new("binary websocket frames are not supported"),
143 "decode websocket frame",
144 ));
145 }
146 Some(Ok(Message::Ping(payload))) => {
147 writer.send(Message::Pong(payload)).await.map_err(|error| {
148 WorkerQuitReason::fatal(
149 WebSocketTransportError::new(format!("websocket pong failed: {error}")),
150 "send websocket pong",
151 )
152 })?;
153 }
154 Some(Ok(Message::Pong(_))) | Some(Ok(Message::Frame(_))) => {}
155 Some(Ok(Message::Close(_))) | None => {
156 return Err(WorkerQuitReason::TransportClosed);
157 }
158 Some(Err(error)) => {
159 return Err(WorkerQuitReason::fatal(
160 WebSocketTransportError::new(format!("websocket receive failed: {error}")),
161 "receive websocket frame",
162 ));
163 }
164 },
165 outbound = context.recv_from_handler() => {
166 let outbound = outbound?;
167 let payload = encode_client_message(&outbound.message).map_err(|error| {
168 WorkerQuitReason::fatal(error, "encode websocket frame")
169 })?;
170 match writer.send(Message::Text(payload.into())).await {
171 Ok(()) => {
172 drop(outbound.responder.send(Ok(())));
173 }
174 Err(error) => {
175 let send_error = WebSocketTransportError::new(format!("websocket send failed: {error}"));
176 let cloned = send_error.clone();
177 drop(outbound.responder.send(Err(cloned)));
178 return Err(WorkerQuitReason::fatal(send_error, "send websocket frame"));
179 }
180 }
181 }
182 }
183 }
184 }
185}
186
187pub type WebSocketClientTransport = WorkerTransport<WebSocketClientWorker>;
188
189pub fn connect(config: WebSocketTransportConfig) -> WebSocketClientTransport {
190 WorkerTransport::spawn(WebSocketClientWorker::new(config))
191}
192
193pub fn encode_client_message(
194 message: &TxJsonRpcMessage<RoleClient>,
195) -> Result<String, WebSocketTransportError> {
196 serde_json::to_string(message).map_err(|error| {
197 WebSocketTransportError::new(format!("failed to encode json-rpc frame: {error}"))
198 })
199}
200
201pub fn decode_server_message(
202 payload: &str,
203) -> Result<RxJsonRpcMessage<RoleClient>, WebSocketTransportError> {
204 serde_json::from_str(payload).map_err(|error| {
205 WebSocketTransportError::new(format!("failed to decode json-rpc frame: {error}"))
206 })
207}
208
209#[cfg(test)]
210#[path = "websocket_tests.rs"]
211mod tests;