use std::io;
use std::pin::Pin;
use futures_util::sink::SinkExt;
use minarrow::{Field, Table, TableV};
use tokio::net::tcp::OwnedWriteHalf;
#[cfg(feature = "tls")]
use tokio::net::TcpStream;
use tokio::net::{TcpListener, ToSocketAddrs};
use crate::compression::Compression;
use crate::enums::IPCMessageProtocol;
use crate::models::sinks::table_sink::TableSink64;
use crate::models::streams::tcp::TcpWriteHalf;
use crate::models::transports::tcp::TcpTransport;
use crate::traits::transport_writer::IPCTransportWriter;
pub struct TcpTableWriter {
sink: TableSink64<TcpWriteHalf>,
}
impl TcpTableWriter {
pub async fn connect(
addr: impl ToSocketAddrs,
schema: Vec<Field>,
compression: Option<Compression>,
) -> io::Result<Self> {
let (_read, write) = TcpTransport::connect(addr).await?;
let sink = TableSink64::new(
TcpWriteHalf::Plain(write),
schema,
IPCMessageProtocol::Stream,
compression,
)?;
Ok(Self { sink })
}
pub async fn accept(
listener: &TcpListener,
schema: Vec<Field>,
compression: Option<Compression>,
) -> io::Result<Self> {
let (_read, write) = TcpTransport::accept(listener).await?;
Self::from_write_half(write, schema, compression)
}
pub fn from_write_half(
write_half: OwnedWriteHalf,
schema: Vec<Field>,
compression: Option<Compression>,
) -> io::Result<Self> {
let sink = TableSink64::new(
TcpWriteHalf::Plain(write_half),
schema,
IPCMessageProtocol::Stream,
compression,
)?;
Ok(Self { sink })
}
pub async fn write_table_with_metadata(
&mut self,
table: impl Into<TableV> + Send,
metadata: Vec<(String, String)>,
) -> io::Result<()> {
self.sink.encode_frame(&table.into(), Some(metadata.as_slice()))?;
SinkExt::flush(&mut self.sink).await?;
Ok(())
}
#[cfg(feature = "tls")]
pub async fn connect_tls(
addr: impl ToSocketAddrs,
server_name: rustls_pki_types::ServerName<'static>,
config: std::sync::Arc<tokio_rustls::rustls::ClientConfig>,
schema: Vec<Field>,
compression: Option<Compression>,
) -> io::Result<Self> {
let half = Self::tls_write_half(addr, server_name, config).await?;
let sink = TableSink64::new(half, schema, IPCMessageProtocol::Stream, compression)?;
Ok(Self { sink })
}
#[cfg(feature = "tls")]
async fn tls_write_half(
addr: impl ToSocketAddrs,
server_name: rustls_pki_types::ServerName<'static>,
config: std::sync::Arc<tokio_rustls::rustls::ClientConfig>,
) -> io::Result<TcpWriteHalf> {
let tcp = TcpStream::connect(addr).await?;
let connector = tokio_rustls::TlsConnector::from(config);
let tls = connector.connect(server_name, tcp).await?;
let (_read_half, write_half) = tokio::io::split(tls);
Ok(TcpWriteHalf::Tls(Box::new(write_half)))
}
}
impl IPCTransportWriter for TcpTableWriter {
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
}
}