use std::io;
use std::pin::Pin;
use futures_util::sink::SinkExt;
use minarrow::{Field, Table, TableV};
use tokio_tungstenite::connect_async;
use tokio::net::TcpListener;
use crate::compression::Compression;
use crate::enums::IPCMessageProtocol;
use crate::models::sinks::table_sink::TableSink64;
use crate::models::streams::websocket::WsWrite;
use crate::models::transports::websocket::WebSocketTransport;
use crate::traits::transport_writer::IPCTransportWriter;
type WsAsyncWrite = Box<dyn tokio::io::AsyncWrite + Send + Sync + Unpin + 'static>;
pub struct WebSocketTableWriter {
sink: TableSink64<WsWrite<WsAsyncWrite>>,
}
impl WebSocketTableWriter {
pub async fn connect(
url: &str,
schema: Vec<Field>,
compression: Option<Compression>,
) -> io::Result<Self> {
let ws_write = Self::plain_ws_write(url).await?;
let sink = TableSink64::new(ws_write, schema, IPCMessageProtocol::Stream, compression)?;
Ok(Self { sink })
}
async fn plain_ws_write(url: &str) -> io::Result<WsWrite<WsAsyncWrite>> {
let (_read_half, write_half) = WebSocketTransport::connect(url).await?;
let boxed: WsAsyncWrite = Box::new(write_half);
let (_shared, ws_write) = WsWrite::new_client(boxed);
Ok(ws_write)
}
pub async fn accept(
listener: &TcpListener,
schema: Vec<Field>,
compression: Option<Compression>,
) -> io::Result<Self> {
let (read_half, write_half) = WebSocketTransport::accept(listener).await?;
Self::from_halves(read_half, write_half, schema, compression)
}
#[cfg(feature = "tls")]
pub async fn connect_tls(
url: &str,
config: std::sync::Arc<tokio_rustls::rustls::ClientConfig>,
schema: Vec<Field>,
compression: Option<Compression>,
) -> io::Result<Self> {
let ws_write = Self::tls_ws_write(url, config).await?;
let sink = TableSink64::new(ws_write, schema, IPCMessageProtocol::Stream, compression)?;
Ok(Self { sink })
}
#[cfg(feature = "tls")]
async fn tls_ws_write(
url: &str,
config: std::sync::Arc<tokio_rustls::rustls::ClientConfig>,
) -> io::Result<WsWrite<WsAsyncWrite>> {
let (_read_half, write_half) = WebSocketTransport::connect_tls(url, config).await?;
let boxed: WsAsyncWrite = Box::new(write_half);
let (_shared, ws_write) = WsWrite::new_client(boxed);
Ok(ws_write)
}
pub fn from_halves<R, W>(
_read_half: R,
write_half: W,
schema: Vec<Field>,
compression: Option<Compression>,
) -> io::Result<Self>
where
R: tokio::io::AsyncRead + Unpin + Send + 'static,
W: tokio::io::AsyncWrite + Send + Sync + Unpin + 'static,
{
let boxed: WsAsyncWrite = Box::new(write_half);
let (_shared, ws_write) = WsWrite::new(boxed);
let sink = TableSink64::new(ws_write, schema, IPCMessageProtocol::Stream, compression)?;
Ok(Self { sink })
}
}
impl IPCTransportWriter for WebSocketTableWriter {
fn schema(&self) -> &[Field] {
&self.sink.schema
}
fn register_dictionary(&mut self, dict_id: i64, values: Vec<String>) {
self.sink.codec.register_dictionary(dict_id, values);
}
async fn write_table(&mut self, table: impl Into<TableV> + Send) -> io::Result<()> {
SinkExt::send(&mut self.sink, table.into()).await?;
SinkExt::flush(&mut self.sink).await?;
Ok(())
}
async fn write_all_tables(&mut self, tables: Vec<Table>) -> io::Result<()> {
let mut sink = Pin::new(&mut self.sink);
for table in tables {
SinkExt::send(&mut sink, table.into()).await?;
}
SinkExt::close(&mut sink).await?;
Ok(())
}
async fn finish(&mut self) -> io::Result<()> {
SinkExt::close(&mut self.sink).await
}
}