use std::io;
use std::pin::Pin;
use bytes::Bytes;
use futures_util::sink::SinkExt;
use http::{Method, Request, Uri};
use minarrow::{Field, Table, TableV};
use tokio::net::{TcpListener, TcpStream};
use crate::compression::Compression;
use crate::enums::IPCMessageProtocol;
use crate::models::sinks::table_sink::TableSink64;
use crate::models::streams::http::{H2RecvRead, H2SendWrite};
use crate::models::transports::http::HttpTransport;
use crate::traits::transport_writer::IPCTransportWriter;
pub struct HttpTableWriter {
sink: TableSink64<H2SendWrite>,
}
impl HttpTableWriter {
pub async fn post(
url: &str,
schema: Vec<Field>,
compression: Option<Compression>,
) -> io::Result<Self> {
let req = parse_post(url)?;
Self::from_request(req, schema, compression).await
}
#[cfg(feature = "tls")]
pub async fn post_tls(
url: &str,
config: std::sync::Arc<tokio_rustls::rustls::ClientConfig>,
schema: Vec<Field>,
compression: Option<Compression>,
) -> io::Result<Self> {
let req = parse_post(url)?;
Self::from_request_tls(req, config, schema, compression).await
}
pub async fn from_request(
req: Request<()>,
schema: Vec<Field>,
compression: Option<Compression>,
) -> io::Result<Self> {
let (host, port) = host_port(req.uri(), "http", 80)?;
let tcp = TcpStream::connect((host.as_str(), port)).await?;
let (send_stream, response_fut) = h2_send_post(tcp, req).await?;
Self::from_send_stream(send_stream, response_fut, schema, compression)
}
#[cfg(feature = "tls")]
pub async fn from_request_tls(
req: Request<()>,
config: std::sync::Arc<tokio_rustls::rustls::ClientConfig>,
schema: Vec<Field>,
compression: Option<Compression>,
) -> io::Result<Self> {
let (host, port) = host_port(req.uri(), "https", 443)?;
let server_name = rustls_pki_types::ServerName::try_from(host.clone())
.map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, e))?;
let tcp = TcpStream::connect((host.as_str(), port)).await?;
let connector = tokio_rustls::TlsConnector::from(config);
let tls = connector.connect(server_name, tcp).await?;
if tls.get_ref().1.alpn_protocol() != Some(b"h2") {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"TLS ALPN did not negotiate h2; set \
config.alpn_protocols = vec![b\"h2\".to_vec()] on ClientConfig",
));
}
let (send_stream, response_fut) = h2_send_post(tls, req).await?;
Self::from_send_stream(send_stream, response_fut, schema, compression)
}
pub async fn accept(
listener: &TcpListener,
schema: Vec<Field>,
compression: Option<Compression>,
) -> io::Result<Self> {
let (recv_read, send_write) = HttpTransport::accept(listener).await?;
Self::from_exchange(recv_read, send_write, schema, compression)
}
pub fn from_exchange(
mut recv_read: H2RecvRead,
send_write: H2SendWrite,
schema: Vec<Field>,
compression: Option<Compression>,
) -> io::Result<Self> {
tokio::spawn(async move {
let _ = tokio::io::copy(&mut recv_read, &mut tokio::io::sink()).await;
});
let sink = TableSink64::new(send_write, schema, IPCMessageProtocol::Stream, compression)?;
Ok(Self { sink })
}
pub fn from_send_stream(
send_stream: h2::SendStream<Bytes>,
response_fut: h2::client::ResponseFuture,
schema: Vec<Field>,
compression: Option<Compression>,
) -> io::Result<Self> {
tokio::spawn(async move {
let _ = response_fut.await;
});
let write = H2SendWrite::new(send_stream);
let sink = TableSink64::new(write, 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(())
}
}
impl IPCTransportWriter for HttpTableWriter {
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
}
}
fn parse_post(url: &str) -> io::Result<Request<()>> {
let uri: Uri = url
.parse()
.map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, e))?;
Request::builder()
.method(Method::POST)
.uri(uri)
.body(())
.map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, e))
}
fn host_port(uri: &Uri, expected_scheme: &str, default_port: u16) -> io::Result<(String, u16)> {
let scheme = uri
.scheme_str()
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidInput, "uri missing scheme"))?;
if scheme != expected_scheme {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("expected scheme {expected_scheme}, got {scheme}"),
));
}
let host = uri
.host()
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidInput, "uri missing host"))?
.to_string();
let port = uri.port_u16().unwrap_or(default_port);
Ok((host, port))
}
async fn h2_send_post<T>(
io: T,
req: Request<()>,
) -> io::Result<(h2::SendStream<Bytes>, h2::client::ResponseFuture)>
where
T: tokio::io::AsyncRead + tokio::io::AsyncWrite + Send + Unpin + 'static,
{
let (mut send_request, connection) =
h2::client::handshake(io).await.map_err(io::Error::other)?;
tokio::spawn(async move {
if let Err(e) = connection.await {
tracing::debug!("h2 connection driver exited: {e}");
}
});
let (response_fut, send_stream) = send_request
.send_request(req, false)
.map_err(io::Error::other)?;
Ok((send_stream, response_fut))
}