1use agent_client_protocol::{ConnectTo, Error, Lines, Role};
2use futures::{SinkExt, StreamExt, stream};
3use std::{io, time::Duration};
4use thiserror::Error;
5use tokio::io::{AsyncRead, AsyncWrite};
6use tokio::sync::mpsc;
7use tokio::time::{Instant, MissedTickBehavior, interval_at, timeout};
8use tokio_tungstenite::WebSocketStream;
9use tokio_tungstenite::tungstenite::protocol::{CloseFrame, frame::coding::CloseCode};
10use tokio_tungstenite::tungstenite::{self, Bytes, Message};
11use tokio_util::sync::PollSender;
12
13#[derive(Debug, Error)]
14pub enum WebSocketError {
15 #[error(transparent)]
16 WebSocket(#[from] tungstenite::Error),
17 #[error("unexpected raw WebSocket frame")]
18 UnexpectedFrame,
19 #[error("binary WebSocket messages are not supported")]
20 BinaryMessage,
21 #[error("WebSocket consumer cannot keep up")]
22 SlowConsumer,
23 #[error("WebSocket write deadline exceeded")]
24 WriteDeadline,
25}
26
27pub struct WebSocketTransport<T> {
29 socket: WebSocketStream<T>,
30}
31
32impl<T: AsyncRead + AsyncWrite + Unpin> WebSocketTransport<T> {
33 pub fn new(socket: WebSocketStream<T>) -> Self {
34 Self { socket }
35 }
36}
37
38impl<T, U> ConnectTo<U> for WebSocketTransport<T>
39where
40 T: AsyncRead + AsyncWrite + Unpin + Send + 'static,
41 U: Role,
42{
43 async fn connect_to(self, peer: impl ConnectTo<U::Counterpart>) -> Result<(), Error> {
44 let (to_tx, to_rx) = mpsc::channel::<String>(32);
45 let (from_tx, from_rx) = mpsc::channel::<io::Result<String>>(32);
46 let output = PollSender::new(to_tx).sink_map_err(|_| io::Error::from(io::ErrorKind::BrokenPipe));
47 let input = stream::unfold(from_rx, |mut rx| async { rx.recv().await.map(|item| (item, rx)) });
48 let connection_future = ConnectTo::<U>::connect_to(Lines::new(output, input), peer);
49 let socket_loop = self.run_socket_loop(to_rx, from_tx);
50 tokio::pin!(connection_future, socket_loop);
51 tokio::select! {
52 result = &mut connection_future => {
53 result?;
54 socket_loop.await.map_err(Error::into_internal_error)
55 }
56
57 result = &mut socket_loop => {
58 result.map_err(Error::into_internal_error)?;
59 connection_future.await
60 },
61 }
62 }
63}
64
65const WRITE_DEADLINE: Duration = Duration::from_secs(10);
66const KEEPALIVE_INTERVAL: Duration = Duration::from_secs(20);
67
68impl<T: AsyncRead + AsyncWrite + Unpin> WebSocketTransport<T> {
69 async fn run_socket_loop(
70 mut self,
71 mut to_rx: mpsc::Receiver<String>,
72 from_tx: mpsc::Sender<io::Result<String>>,
73 ) -> Result<(), WebSocketError> {
74 let mut keepalive = interval_at(Instant::now() + KEEPALIVE_INTERVAL, KEEPALIVE_INTERVAL);
75 keepalive.set_missed_tick_behavior(MissedTickBehavior::Skip);
76 loop {
77 tokio::select! {
78 _ = keepalive.tick() => {
79 self.write(Message::Ping(Bytes::new())).await?;
80 }
81 message = to_rx.recv() => {
82 let Some(text) = message else {
83 return self.finish_close().await;
84 };
85 self.write(Message::Text(text.into())).await?;
86 }
87 message = self.socket.next() => {
88 match message {
89 Some(Ok(Message::Text(text))) => {
90 from_tx.try_send(Ok(text.to_string())).map_err(|_| WebSocketError::SlowConsumer)?;
91 }
92 Some(Ok(Message::Ping(_))) => {
93 timeout(WRITE_DEADLINE, self.socket.flush()).await
95 .map_err(|_| WebSocketError::WriteDeadline)??;
96 }
97 Some(Ok(Message::Pong(_))) => {},
98 Some(Ok(Message::Close(_))) => {
99 let _ = timeout(WRITE_DEADLINE, self.socket.flush()).await;
100 return Ok(());
101 }
102 Some(Ok(Message::Binary(_))) => {
103 self.close(CloseCode::Unsupported).await;
104 return Err(WebSocketError::BinaryMessage);
105 }
106 Some(Ok(Message::Frame(_))) => return Err(WebSocketError::UnexpectedFrame),
107 Some(Err(tungstenite::Error::ConnectionClosed | tungstenite::Error::AlreadyClosed)) | None => return Ok(()),
108 Some(Err(error)) => {
109 if matches!(error, tungstenite::Error::Capacity(_)) {
110 self.close(CloseCode::Size).await;
111 }
112 return Err(error.into());
113 }
114 }
115 }
116 }
117 }
118 }
119
120 async fn finish_close(&mut self) -> Result<(), WebSocketError> {
121 timeout(WRITE_DEADLINE, async {
122 match self.socket.send(Message::Close(None)).await {
123 Ok(()) => {}
124 Err(tungstenite::Error::ConnectionClosed | tungstenite::Error::AlreadyClosed) => return Ok(()),
125 Err(error) => return Err(error.into()),
126 }
127 while let Some(message) = self.socket.next().await {
128 match message {
129 Ok(Message::Close(_)) | Err(tungstenite::Error::ConnectionClosed) => break,
130 Err(error) => return Err(error.into()),
131 _ => {}
132 }
133 }
134 Ok(())
135 })
136 .await
137 .map_err(|_| WebSocketError::WriteDeadline)?
138 }
139
140 async fn write(&mut self, message: Message) -> Result<(), WebSocketError> {
141 timeout(WRITE_DEADLINE, self.socket.send(message))
142 .await
143 .map_err(|_| WebSocketError::WriteDeadline)?
144 .map_err(WebSocketError::from)
145 }
146
147 async fn close(&mut self, code: CloseCode) {
148 let _ = self.write(Message::Close(Some(CloseFrame { code, reason: "".into() }))).await;
149 }
150}