use async_trait::async_trait;
use futures_util::{SinkExt, StreamExt};
use tokio::sync::mpsc;
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use tokio_tungstenite::tungstenite::http::HeaderValue;
use tokio_tungstenite::tungstenite::protocol::Message;
use crate::error::{Error, Result};
use super::transport::ClientTransport;
#[derive(Debug, Clone, Default)]
pub struct WebSocketClientConfig {
pub protocol_version: Option<String>,
pub bearer: Option<String>,
pub headers: Vec<(String, String)>,
}
pub struct WebSocketClientTransport {
outgoing_tx: mpsc::Sender<String>,
incoming_rx: mpsc::Receiver<String>,
connected: bool,
}
impl WebSocketClientTransport {
pub async fn connect(url: &str) -> Result<Self> {
Self::connect_with_config(url, WebSocketClientConfig::default()).await
}
pub async fn connect_with_config(url: &str, config: WebSocketClientConfig) -> Result<Self> {
let mut request = url
.into_client_request()
.map_err(|e| Error::Transport(format!("invalid WebSocket URL: {e}")))?;
let mut subprotocols = Vec::new();
if let Some(version) = &config.protocol_version {
subprotocols.push(format!("mcp.version.{version}"));
}
if let Some(token) = &config.bearer {
subprotocols.push(format!("mcp.auth.{token}"));
}
if !subprotocols.is_empty() {
let value = subprotocols.join(", ");
let header = HeaderValue::from_str(&value).map_err(|e| {
Error::Transport(format!("subprotocol is not a valid header value: {e}"))
})?;
request
.headers_mut()
.insert("sec-websocket-protocol", header);
}
if let Some(token) = &config.bearer {
let header = HeaderValue::from_str(&format!("Bearer {token}"))
.map_err(|e| Error::Transport(format!("bearer token is not a header: {e}")))?;
request.headers_mut().insert("authorization", header);
}
for (name, value) in &config.headers {
let name: tokio_tungstenite::tungstenite::http::HeaderName = name
.parse()
.map_err(|e| Error::Transport(format!("invalid header name '{name}': {e}")))?;
let value = HeaderValue::from_str(value)
.map_err(|e| Error::Transport(format!("invalid header value: {e}")))?;
request.headers_mut().insert(name, value);
}
let (stream, _response) = tokio_tungstenite::connect_async(request)
.await
.map_err(|e| Error::Transport(format!("WebSocket connect failed: {e}")))?;
let (mut sink, mut source) = stream.split();
let (outgoing_tx, mut outgoing_rx) = mpsc::channel::<String>(64);
let (incoming_tx, incoming_rx) = mpsc::channel::<String>(64);
tokio::spawn(async move {
while let Some(message) = outgoing_rx.recv().await {
if sink.send(Message::Text(message.into())).await.is_err() {
break;
}
}
let _ = sink.close().await;
});
tokio::spawn(async move {
while let Some(message) = source.next().await {
match message {
Ok(Message::Text(text)) => {
if incoming_tx.send(text.to_string()).await.is_err() {
break;
}
}
Ok(Message::Close(_)) => break,
Ok(_) => continue,
Err(error) => {
tracing::debug!(%error, "WebSocket receive error");
break;
}
}
}
});
Ok(Self {
outgoing_tx,
incoming_rx,
connected: true,
})
}
}
#[async_trait]
impl ClientTransport for WebSocketClientTransport {
async fn send(&mut self, message: &str) -> Result<()> {
self.outgoing_tx
.send(message.to_string())
.await
.map_err(|_| Error::Transport("WebSocket connection closed".to_string()))
}
async fn recv(&mut self) -> Result<Option<String>> {
match self.incoming_rx.recv().await {
Some(message) => Ok(Some(message)),
None => {
self.connected = false;
Ok(None)
}
}
}
fn is_connected(&self) -> bool {
self.connected && !self.outgoing_tx.is_closed()
}
async fn close(&mut self) -> Result<()> {
self.connected = false;
self.incoming_rx.close();
Ok(())
}
fn supports_session_recovery(&self) -> bool {
false
}
}