use std::io;
use minarrow::{Field, Vec64};
use tokio_uring::buf::BoundedBuf;
use crate::models::codecs::lightstream::LightstreamCodec;
use crate::models::decoders::limits::DecodeLimits;
use crate::models::frames::lightstream_message::{FRAME_HEADER_SIZE, LightstreamMessage};
use super::buf::UringBuf;
use super::stream::UringStream;
pub type IoUringUdsConnection = IoUringConnection<tokio_uring::net::UnixStream>;
pub type IoUringTcpConnection = IoUringConnection<tokio_uring::net::TcpStream>;
async fn read_exact_vec64<S: UringStream>(
stream: &S,
mut buf: UringBuf,
mut offset: usize,
len: usize,
) -> io::Result<UringBuf> {
let target = offset + len;
while offset < target {
let slice = buf.slice(offset..target);
let (result, slice) = stream.read(slice).await;
buf = slice.into_inner();
let n = result?;
if n == 0 {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"connection closed during read",
));
}
offset += n;
}
Ok(buf)
}
async fn read_exact_vec<S: UringStream>(
stream: &S,
mut buf: Vec<u8>,
mut offset: usize,
len: usize,
) -> io::Result<Vec<u8>> {
let target = offset + len;
while offset < target {
let slice = buf.slice(offset..target);
let (result, slice) = stream.read(slice).await;
buf = slice.into_inner();
let n = result?;
if n == 0 {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"connection closed during read",
));
}
offset += n;
}
Ok(buf)
}
pub struct IoUringConnection<S: UringStream> {
stream: S,
read_codec: LightstreamCodec<Vec64<u8>>,
write_codec: LightstreamCodec<Vec64<u8>>,
encode_buf: Vec64<u8>,
header_buf: Vec<u8>,
eof: bool,
limits: DecodeLimits,
}
impl<S: UringStream> IoUringConnection<S> {
pub fn new(stream: S, limits: Option<DecodeLimits>) -> Self {
let limits = limits.unwrap_or_default();
let mut header_buf = Vec::with_capacity(FRAME_HEADER_SIZE);
header_buf.resize(FRAME_HEADER_SIZE, 0);
Self {
stream,
read_codec: LightstreamCodec::new(Some(limits)),
write_codec: LightstreamCodec::new(Some(limits)),
encode_buf: Vec64::with_capacity(0),
header_buf,
eof: false,
limits,
}
}
pub fn register_message(&mut self, name: impl Into<String>) -> u8 {
let name = name.into();
let tag = self.write_codec.register_message(name.clone());
let _ = self.read_codec.register_message(name);
tag
}
pub fn register_table(&mut self, name: impl Into<String>, schema: Vec<Field>) -> u8 {
let name = name.into();
let tag = self
.write_codec
.register_table(name.clone(), schema.clone());
let _ = self.read_codec.register_table(name, schema);
tag
}
pub async fn send(&mut self, name: &str, payload: &[u8]) -> io::Result<()> {
let tag = self.write_codec.tag_by_name(name).ok_or_else(|| {
io::Error::new(
io::ErrorKind::InvalidInput,
format!("unknown type name '{}'", name),
)
})?;
let frame = self.write_codec.encode_message(tag, payload)?;
self.stream.write_all(UringBuf(frame)).await.0?;
Ok(())
}
pub async fn send_table(
&mut self,
name: &str,
table: impl Into<minarrow::TableV>,
) -> io::Result<()> {
let view: minarrow::TableV = table.into();
let tag = self.write_codec.tag_by_name(name).ok_or_else(|| {
io::Error::new(
io::ErrorKind::InvalidInput,
format!("unknown type name '{}'", name),
)
})?;
self.write_codec
.encode_table(tag, &view, &mut self.encode_buf)?;
let wire_buf = std::mem::replace(&mut self.encode_buf, Vec64::with_capacity(0));
let (result, UringBuf(returned)) = self.stream.write_all(UringBuf(wire_buf)).await;
result?;
self.encode_buf = returned;
self.encode_buf.clear();
Ok(())
}
pub async fn recv(&mut self) -> Option<io::Result<LightstreamMessage>> {
if self.eof {
return None;
}
let header_buf = std::mem::take(&mut self.header_buf);
self.header_buf = match read_exact_vec(&self.stream, header_buf, 0, FRAME_HEADER_SIZE).await
{
Ok(buf) => buf,
Err(e) if e.kind() == io::ErrorKind::UnexpectedEof => {
self.eof = true;
return None;
}
Err(e) => {
self.eof = true;
return Some(Err(e));
}
};
let tag = self.header_buf[0];
let payload_len = u32::from_le_bytes(self.header_buf[1..5].try_into().unwrap()) as usize;
if let Err(e) = self
.limits
.check(payload_len, self.limits.max_frame_bytes, "TLV frame bytes")
{
self.eof = true;
return Some(Err(e));
}
let payload_buf = UringBuf(Vec64::with_capacity(payload_len));
if payload_len > 0 {
let payload_buf =
match read_exact_vec64(&self.stream, payload_buf, 0, payload_len).await {
Ok(buf) => buf,
Err(e) => {
self.eof = true;
return Some(Err(e));
}
};
Some(self.read_codec.decode_frame(tag, payload_buf.0))
} else {
Some(self.read_codec.decode_frame(tag, payload_buf.0))
}
}
pub async fn flush(&mut self) -> io::Result<()> {
Ok(())
}
pub async fn shutdown(&mut self) -> io::Result<()> {
self.stream.shutdown(std::net::Shutdown::Write)
}
#[cfg(feature = "protobuf")]
pub async fn send_proto<M: prost::Message>(&mut self, name: &str, msg: &M) -> io::Result<()> {
let bytes = msg.encode_to_vec();
self.send(name, &bytes).await
}
#[cfg(feature = "msgpack")]
pub async fn send_msgpack<M: serde::Serialize>(
&mut self,
name: &str,
msg: &M,
) -> io::Result<()> {
let bytes = encode_msgpack(msg)?;
self.send(name, &bytes).await
}
}
impl IoUringUdsConnection {
pub fn from_unix_stream(
stream: std::os::unix::net::UnixStream,
limits: Option<DecodeLimits>,
) -> Self {
Self::new(tokio_uring::net::UnixStream::from_std(stream), limits)
}
pub fn from_tokio_unix_stream(
stream: tokio::net::UnixStream,
limits: Option<DecodeLimits>,
) -> io::Result<Self> {
Ok(Self::from_unix_stream(stream.into_std()?, limits))
}
pub fn socketpair(
limits: Option<DecodeLimits>,
) -> io::Result<(Self, std::os::unix::net::UnixStream)> {
let (parent, child) = std::os::unix::net::UnixStream::pair()?;
parent.set_nonblocking(true)?;
let conn = Self::new(tokio_uring::net::UnixStream::from_std(parent), limits);
Ok((conn, child))
}
}
impl IoUringTcpConnection {
pub fn from_tcp_stream(
stream: std::net::TcpStream,
limits: Option<DecodeLimits>,
) -> Self {
Self::new(tokio_uring::net::TcpStream::from_std(stream), limits)
}
pub fn from_tokio_tcp_stream(
stream: tokio::net::TcpStream,
limits: Option<DecodeLimits>,
) -> io::Result<Self> {
Ok(Self::from_tcp_stream(stream.into_std()?, limits))
}
}
#[cfg(feature = "msgpack")]
pub(super) fn encode_msgpack<M: serde::Serialize>(msg: &M) -> io::Result<Vec<u8>> {
let mut buf = Vec::new();
let mut serializer =
rmp_serde::Serializer::new(&mut buf).with_bytes(rmp_serde::config::BytesMode::ForceAll);
msg.serialize(&mut serializer)
.map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
Ok(buf)
}