Skip to main content

soma_mcp_client/upstream/transport/
websocket.rs

1use 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;