use std::io;
use http::{Method, Request, Uri};
use minarrow::{Field, Table, TableV};
use tokio::net::TcpStream;
use tokio::sync::mpsc;
use tokio::task::JoinHandle;
use crate::compression::Compression;
use crate::models::writers::http::HttpTableWriter;
use crate::traits::parallel_transport_writer::{ParallelTransportWriter, SEQ_ID_META_KEY};
use crate::traits::transport_writer::IPCTransportWriter;
const STREAM_CHANNEL_DEPTH: usize = 8;
pub struct HttpParallelTableWriter {
schema: Vec<Field>,
senders: Vec<mpsc::Sender<(TableV, Option<u64>)>>,
tasks: Vec<JoinHandle<io::Result<()>>>,
next: usize,
ordered: bool,
}
impl HttpParallelTableWriter {
pub async fn connect(
url: &str,
stream_count: usize,
schema: Vec<Field>,
dictionaries: Vec<(i64, Vec<String>)>,
compression: Option<Compression>,
) -> io::Result<Self> {
assert!(stream_count >= 1, "stream_count must be at least 1");
let uri: Uri = url
.parse()
.map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, e))?;
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(80);
let tcp = TcpStream::connect((host.as_str(), port)).await?;
let (mut send_request, conn) = h2::client::handshake(tcp)
.await
.map_err(io::Error::other)?;
tokio::spawn(async move {
let _ = conn.await;
});
let mut senders = Vec::with_capacity(stream_count);
let mut tasks = Vec::with_capacity(stream_count);
for _ in 0..stream_count {
let req = Request::builder()
.method(Method::POST)
.uri(uri.clone())
.body(())
.map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, e))?;
let (response, send_stream) = send_request
.send_request(req, false)
.map_err(io::Error::other)?;
let mut writer = HttpTableWriter::from_send_stream(
send_stream,
response,
schema.clone(),
compression,
)?;
for (dict_id, values) in &dictionaries {
writer.register_dictionary(*dict_id, values.clone());
}
let (tx, mut rx) = mpsc::channel::<(TableV, Option<u64>)>(STREAM_CHANNEL_DEPTH);
let task = tokio::spawn(async move {
while let Some((table, seq)) = rx.recv().await {
match seq {
Some(seq) => {
writer
.write_table_with_metadata(
table,
vec![(SEQ_ID_META_KEY.to_string(), seq.to_string())],
)
.await?
}
None => writer.write_table(table).await?,
}
}
writer.finish().await
});
senders.push(tx);
tasks.push(task);
}
Ok(Self { schema, senders, tasks, next: 0, ordered: false })
}
pub async fn connect_ordered(
url: &str,
stream_count: usize,
schema: Vec<Field>,
dictionaries: Vec<(i64, Vec<String>)>,
compression: Option<Compression>,
) -> io::Result<Self> {
let mut writer =
Self::connect(url, stream_count, schema, dictionaries, compression).await?;
writer.ordered = true;
Ok(writer)
}
}
impl ParallelTransportWriter for HttpParallelTableWriter {
fn schema(&self) -> &[Field] {
&self.schema
}
fn stream_count(&self) -> usize {
self.senders.len()
}
async fn write_table(&mut self, table: impl Into<TableV> + Send) -> io::Result<()> {
let seq = if self.ordered { Some(self.next as u64) } else { None };
let idx = self.next % self.senders.len();
self.next = self.next.wrapping_add(1);
self.senders[idx]
.send((table.into(), seq))
.await
.map_err(|_| io::Error::new(io::ErrorKind::BrokenPipe, "HTTP/2 stream task closed"))
}
async fn write_all_tables(&mut self, tables: Vec<Table>) -> io::Result<()> {
for table in tables {
self.write_table(table).await?;
}
Ok(())
}
async fn finish(mut self) -> io::Result<()> {
self.senders.clear();
let mut first_err: Option<io::Error> = None;
for task in self.tasks.drain(..) {
match task.await {
Ok(Ok(())) => {}
Ok(Err(e)) => {
if first_err.is_none() {
first_err = Some(e);
}
}
Err(join_err) => {
if first_err.is_none() {
first_err = Some(io::Error::other(join_err));
}
}
}
}
match first_err {
Some(e) => Err(e),
None => Ok(()),
}
}
}