use std::io;
use minarrow::{Field, TableV};
use quinn::Connection;
use tokio::sync::mpsc;
use tokio::task::JoinHandle;
use crate::compression::Compression;
use crate::models::writers::quic::QuicTableWriter;
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 QuicParallelTableWriter {
schema: Vec<Field>,
senders: Vec<mpsc::Sender<(TableV, Option<u64>)>>,
tasks: Vec<JoinHandle<io::Result<()>>>,
next: usize,
ordered: bool,
}
impl QuicParallelTableWriter {
pub async fn open(
conn: &Connection,
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 mut senders = Vec::with_capacity(stream_count);
let mut tasks = Vec::with_capacity(stream_count);
for _ in 0..stream_count {
let send = conn.open_uni().await.map_err(io::Error::other)?;
let mut writer = QuicTableWriter::new(send, 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 open_ordered(
conn: &Connection,
stream_count: usize,
schema: Vec<Field>,
dictionaries: Vec<(i64, Vec<String>)>,
compression: Option<Compression>,
) -> io::Result<Self> {
let mut writer =
Self::open(conn, stream_count, schema, dictionaries, compression).await?;
writer.ordered = true;
Ok(writer)
}
}
impl ParallelTransportWriter for QuicParallelTableWriter {
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, "QUIC stream task closed"))
}
async fn write_all_tables(
&mut self,
tables: Vec<impl Into<TableV> + Send>,
) -> 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(()),
}
}
}