use std::future::Future;
use bytes::{
Bytes,
BytesMut,
};
use tokio::io::{
AsyncRead,
AsyncReadExt,
AsyncWrite,
AsyncWriteExt,
};
#[derive(Debug)]
pub enum TransportError {
FastWebSocket(fastwebsockets::WebSocketError),
Io(std::io::Error),
}
impl std::fmt::Display for TransportError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::FastWebSocket(e) => write!(f, "fastwebsockets: {e}"),
Self::Io(e) => write!(f, "transport I/O: {e}"),
}
}
}
impl std::error::Error for TransportError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::FastWebSocket(e) => Some(e),
Self::Io(e) => Some(e),
}
}
}
impl From<fastwebsockets::WebSocketError> for TransportError {
fn from(e: fastwebsockets::WebSocketError) -> Self {
Self::FastWebSocket(e)
}
}
impl From<std::io::Error> for TransportError {
fn from(e: std::io::Error) -> Self {
Self::Io(e)
}
}
pub enum WsMessage {
Binary(Bytes),
Ping(Bytes),
Close,
}
pub trait WsFrameWriter: Send {
fn write_binary(
&mut self, payload: Bytes,
) -> impl Future<Output = Result<(), TransportError>> + Send;
fn write_pong(
&mut self, payload: Bytes,
) -> impl Future<Output = Result<(), TransportError>> + Send;
fn flush(&mut self) -> impl Future<Output = Result<(), TransportError>> + Send;
fn close(&mut self) -> impl Future<Output = Result<(), TransportError>> + Send;
}
pub trait WsFrameReader: Send {
fn read_message(
&mut self,
) -> impl Future<Output = Option<Result<WsMessage, TransportError>>> + Send;
}
pub struct FastWsWriter<S: AsyncWrite + Unpin + Send> {
ws: fastwebsockets::WebSocketWrite<S>,
}
impl<S: AsyncWrite + Unpin + Send> WsFrameWriter for FastWsWriter<S> {
async fn write_binary(&mut self, payload: Bytes) -> Result<(), TransportError> {
let frame = fastwebsockets::Frame::binary(fastwebsockets::Payload::Borrowed(&payload));
self.ws.write_frame(frame).await?;
Ok(())
}
async fn write_pong(&mut self, payload: Bytes) -> Result<(), TransportError> {
let frame = fastwebsockets::Frame::pong(fastwebsockets::Payload::Borrowed(&payload));
self.ws.write_frame(frame).await?;
Ok(())
}
async fn flush(&mut self) -> Result<(), TransportError> {
Ok(())
}
async fn close(&mut self) -> Result<(), TransportError> {
let frame = fastwebsockets::Frame::close(1000, &[]);
self.ws.write_frame(frame).await?;
Ok(())
}
}
pub struct FastWsReader<S: AsyncRead + Unpin + Send> {
ws: fastwebsockets::FragmentCollectorRead<S>,
}
impl<S: AsyncRead + Unpin + Send> FastWsReader<S> {
pub const fn new(ws: fastwebsockets::FragmentCollectorRead<S>) -> Self {
Self { ws }
}
}
impl<S: AsyncRead + Unpin + Send> WsFrameReader for FastWsReader<S> {
async fn read_message(&mut self) -> Option<Result<WsMessage, TransportError>> {
loop {
let frame = match self
.ws
.read_frame(&mut |_| async { Ok::<(), fastwebsockets::WebSocketError>(()) })
.await
{
Ok(f) => f,
Err(fastwebsockets::WebSocketError::ConnectionClosed) => return None,
Err(e) => return Some(Err(TransportError::FastWebSocket(e))),
};
match frame.opcode {
fastwebsockets::OpCode::Binary => {
let bytes = payload_to_bytes(frame.payload);
return Some(Ok(WsMessage::Binary(bytes)));
}
fastwebsockets::OpCode::Ping => {
let bytes = payload_to_bytes(frame.payload);
return Some(Ok(WsMessage::Ping(bytes)));
}
fastwebsockets::OpCode::Close => return Some(Ok(WsMessage::Close)),
fastwebsockets::OpCode::Pong => {}
fastwebsockets::OpCode::Text | fastwebsockets::OpCode::Continuation => {
tracing::warn!(
opcode = ?frame.opcode,
"fastwebsockets reader: unexpected opcode from K8s apiserver, \
absorbing (SPDY tunnel expects only Binary frames)"
);
}
}
}
}
}
fn payload_to_bytes(payload: fastwebsockets::Payload<'_>) -> Bytes {
match payload {
fastwebsockets::Payload::Bytes(bm) => bm.freeze(),
fastwebsockets::Payload::Owned(v) => Bytes::from(v),
fastwebsockets::Payload::Borrowed(b) => Bytes::copy_from_slice(b),
fastwebsockets::Payload::BorrowedMut(b) => Bytes::copy_from_slice(b),
}
}
pub fn split_fastws<S>(
stream: S,
) -> (
FastWsWriter<tokio::io::WriteHalf<S>>,
FastWsReader<tokio::io::ReadHalf<S>>,
)
where
S: AsyncRead + AsyncWrite + Unpin + Send,
{
let (read_half, write_half) = tokio::io::split(stream);
let (mut ws_read, ws_write) =
fastwebsockets::after_handshake_split(read_half, write_half, fastwebsockets::Role::Client);
ws_read.set_auto_pong(false);
ws_read.set_auto_close(false);
let frag_read = fastwebsockets::FragmentCollectorRead::new(ws_read);
(FastWsWriter { ws: ws_write }, FastWsReader::new(frag_read))
}
const RAW_READ_BUF_SIZE: usize = 16 * 1024;
pub struct RawSpdyWriter<W: AsyncWrite + Unpin + Send> {
inner: W,
}
impl<W: AsyncWrite + Unpin + Send> RawSpdyWriter<W> {
pub const fn new(inner: W) -> Self {
Self { inner }
}
}
impl<W: AsyncWrite + Unpin + Send> WsFrameWriter for RawSpdyWriter<W> {
async fn write_binary(&mut self, payload: Bytes) -> Result<(), TransportError> {
self.inner.write_all(&payload).await?;
Ok(())
}
async fn write_pong(&mut self, _payload: Bytes) -> Result<(), TransportError> {
Ok(())
}
async fn flush(&mut self) -> Result<(), TransportError> {
self.inner.flush().await?;
Ok(())
}
async fn close(&mut self) -> Result<(), TransportError> {
self.inner.shutdown().await?;
Ok(())
}
}
pub struct RawSpdyReader<R: AsyncRead + Unpin + Send> {
inner: R,
buf: BytesMut,
}
impl<R: AsyncRead + Unpin + Send> RawSpdyReader<R> {
pub fn new(inner: R) -> Self {
Self {
inner,
buf: BytesMut::with_capacity(RAW_READ_BUF_SIZE),
}
}
}
impl<R: AsyncRead + Unpin + Send> WsFrameReader for RawSpdyReader<R> {
async fn read_message(&mut self) -> Option<Result<WsMessage, TransportError>> {
if self.buf.capacity() == self.buf.len() {
self.buf.reserve(RAW_READ_BUF_SIZE);
}
match self.inner.read_buf(&mut self.buf).await {
Ok(0) => None,
Ok(_) => {
let chunk = self.buf.split().freeze();
Some(Ok(WsMessage::Binary(chunk)))
}
Err(e) => Some(Err(TransportError::Io(e))),
}
}
}
pub fn split_raw_spdy<S>(
stream: S,
) -> (
RawSpdyWriter<tokio::io::WriteHalf<S>>,
RawSpdyReader<tokio::io::ReadHalf<S>>,
)
where
S: AsyncRead + AsyncWrite + Unpin + Send,
{
let (read_half, write_half) = tokio::io::split(stream);
(
RawSpdyWriter::new(write_half),
RawSpdyReader::new(read_half),
)
}