use bytes::{Buf, Bytes};
use futures::{
Future,
future::{BoxFuture, FutureExt},
};
use serde::{Deserialize, Serialize, de::DeserializeOwned};
use std::{
any::Any,
cell::RefCell,
collections::HashMap,
error::Error,
fmt,
ops::{Deref, DerefMut},
panic,
rc::{Rc, Weak},
};
use tracing::Instrument;
use super::{super::DEFAULT_MAX_ITEM_SIZE, BIG_DATA_CHUNK_QUEUE, io::ChannelBytesReader};
use crate::{
chmux::{self, AnyStorage, Received, RecvChunkError},
codec::{self, AnySend, DeserializationError, ErasedDeserializer, StreamingUnavailable},
exec::{
self,
task::{self, JoinHandle},
},
};
#[derive(Clone, Debug, Serialize, Deserialize)]
pub enum RecvError {
Receive(chmux::RecvError),
Deserialize(DeserializationError),
MissingPorts(Vec<u32>),
MaxItemSizeExceeded,
}
impl From<chmux::RecvError> for RecvError {
fn from(err: chmux::RecvError) -> Self {
Self::Receive(err)
}
}
impl From<DeserializationError> for RecvError {
fn from(err: DeserializationError) -> Self {
Self::Deserialize(err)
}
}
impl fmt::Display for RecvError {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
match self {
Self::Receive(err) => write!(f, "receive error: {err}"),
Self::Deserialize(err) => write!(f, "deserialization error: {err}"),
Self::MissingPorts(ports) => write!(
f,
"missing chmux ports: {}",
ports.iter().map(|p| p.to_string()).collect::<Vec<_>>().join(", ")
),
Self::MaxItemSizeExceeded => write!(f, "maximum item size exceeded"),
}
}
}
impl Error for RecvError {}
impl RecvError {
pub fn is_final(&self) -> bool {
match self {
Self::Receive(err) => err.is_final(),
Self::Deserialize(_) | Self::MissingPorts(_) | Self::MaxItemSizeExceeded => false,
}
}
}
pub struct PortDeserializer {
allocator: chmux::PortAllocator,
#[allow(clippy::type_complexity)]
expected: HashMap<
u32,
(
chmux::PortNumber,
Box<dyn FnOnce(chmux::PortNumber, chmux::Request) -> BoxFuture<'static, ()> + Send + 'static>,
),
>,
storage: AnyStorage,
tasks: Vec<BoxFuture<'static, ()>>,
}
impl PortDeserializer {
thread_local! {
static INSTANCE: RefCell<Weak<RefCell<PortDeserializer>>> = const { RefCell::new(Weak::new()) };
}
fn start(allocator: chmux::PortAllocator, storage: AnyStorage) -> Rc<RefCell<PortDeserializer>> {
let this =
Rc::new(RefCell::new(Self { allocator, expected: HashMap::new(), storage, tasks: Vec::new() }));
let weak = Rc::downgrade(&this);
Self::INSTANCE.with(move |i| i.replace(weak));
this
}
fn instance<E>() -> Result<Rc<RefCell<Self>>, E>
where
E: serde::de::Error,
{
match Self::INSTANCE.with(|i| i.borrow().upgrade()) {
Some(this) => Ok(this),
None => Err(serde::de::Error::custom("this remoc object can only be deserialized during receiving")),
}
}
fn finish(this: Rc<RefCell<PortDeserializer>>) -> Self {
match Rc::try_unwrap(this) {
Ok(i) => i.into_inner(),
Err(_) => panic!("PortDeserializer is referenced after deserialization finished"),
}
}
pub fn accept<E>(
remote_port: u32,
callback: impl FnOnce(chmux::PortNumber, chmux::Request) -> BoxFuture<'static, ()> + Send + 'static,
) -> Result<u32, E>
where
E: serde::de::Error,
{
let this = Self::instance()?;
let mut this =
this.try_borrow_mut().expect("PortDeserializer is referenced multiple times during deserialization");
let local_port =
this.allocator.try_allocate().ok_or_else(|| serde::de::Error::custom("ports exhausted"))?;
let local_port_num = *local_port;
this.expected.insert(remote_port, (local_port, Box::new(callback)));
Ok(local_port_num)
}
pub fn storage<E>() -> Result<AnyStorage, E>
where
E: serde::de::Error,
{
let this = Self::instance()?;
let this =
this.try_borrow().expect("PortDeserializer is referenced multiple times during deserialization");
Ok(this.storage.clone())
}
pub fn spawn<E>(task: impl Future<Output = ()> + Send + 'static) -> Result<(), E>
where
E: serde::de::Error,
{
let this = Self::instance()?;
let mut this =
this.try_borrow_mut().expect("PortDeserializer is referenced multiple times during deserialization");
this.tasks.push(task.boxed());
Ok(())
}
}
pub struct Receiver<T, Codec = codec::Default> {
erased: ErasedReceiver,
_phantom: fn(T, Codec),
}
impl<T, Codec> fmt::Debug for Receiver<T, Codec> {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
f.debug_tuple("Receiver").field(&self.erased).finish()
}
}
impl<T, Codec> Deref for Receiver<T, Codec>
where
T: DeserializeOwned + Send + 'static,
Codec: codec::Codec,
{
type Target = ErasedReceiver;
fn deref(&self) -> &Self::Target {
&self.erased
}
}
impl<T, Codec> DerefMut for Receiver<T, Codec>
where
T: DeserializeOwned + Send + 'static,
Codec: codec::Codec,
{
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.erased
}
}
impl<T, Codec> Receiver<T, Codec>
where
T: DeserializeOwned + Send + 'static,
Codec: codec::Codec,
{
pub fn new(receiver: chmux::Receiver) -> Self {
Self { erased: ErasedReceiver::typed::<T, Codec>(receiver), _phantom: |_, _| () }
}
pub fn into_inner(self) -> chmux::Receiver {
self.erased.into_inner()
}
fn from_any(any_item: AnySend) -> T {
let Ok(item) = any_item.downcast::<T>() else { panic!("mismatched type for Receiver") };
*item
}
pub async fn recv(&mut self) -> Result<Option<T>, RecvError> {
self.erased.recv_erased().await.map(|opt| opt.map(Self::from_any))
}
}
pub struct ErasedReceiver {
deserializer: ErasedDeserializer,
receiver: chmux::Receiver,
recved: Option<Option<Received>>,
data: DataSource,
item: Option<Box<dyn Any + Send>>,
port_deser: Option<PortDeserializer>,
default_max_ports: Option<usize>,
max_item_size: usize,
}
impl fmt::Debug for ErasedReceiver {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
f.debug_struct("ErasedReceiver")
.field("deserializer", &self.deserializer)
.field("receiver", &self.receiver)
.finish()
}
}
enum DataSource {
None,
Buffered(Option<chmux::DataBuf>),
Streamed {
tx: Option<tokio::sync::mpsc::Sender<Result<Bytes, ()>>>,
#[allow(clippy::type_complexity)]
task: JoinHandle<Result<(Box<dyn Any + Send>, PortDeserializer), DeserializationError>>,
total: usize,
},
}
impl ErasedReceiver {
pub fn new(deserializer: ErasedDeserializer, receiver: chmux::Receiver) -> Self {
Self {
deserializer,
receiver,
recved: None,
data: DataSource::None,
item: None,
port_deser: None,
default_max_ports: None,
max_item_size: DEFAULT_MAX_ITEM_SIZE,
}
}
pub fn typed<T, Codec>(receiver: chmux::Receiver) -> Self
where
T: DeserializeOwned + Send + 'static,
Codec: codec::Codec,
{
Self::new(ErasedDeserializer::new::<T, Codec>(), receiver)
}
pub fn into_inner(self) -> chmux::Receiver {
self.receiver
}
pub async fn recv_erased(&mut self) -> Result<Option<AnySend>, RecvError> {
if self.default_max_ports.is_none() {
self.default_max_ports = Some(self.receiver.max_ports());
}
'restart: loop {
if self.item.is_none() {
if let DataSource::None = &self.data {
if self.recved.is_none() {
self.recved = Some(self.receiver.recv_any().await?);
}
if let Some(Some(Received::Chunks)) = &self.recved
&& !exec::are_threads_available().await
{
return Err(RecvError::Deserialize(DeserializationError::new(StreamingUnavailable)));
}
self.data = match self.recved.take().unwrap() {
Some(Received::Data(data)) => DataSource::Buffered(Some(data)),
Some(Received::Chunks) => {
let deserializer = self.deserializer.clone();
let allocator = self.receiver.port_allocator();
let handle_storage = self.receiver.storage();
let (tx, rx) = tokio::sync::mpsc::channel(BIG_DATA_CHUNK_QUEUE);
let task = task::spawn_blocking(move || {
let mut cbr = ChannelBytesReader::new(rx);
let pds_ref = PortDeserializer::start(allocator, handle_storage);
let item = deserializer.deserialize(&mut cbr)?;
let pds = PortDeserializer::finish(pds_ref);
Ok((item, pds))
});
DataSource::Streamed { tx: Some(tx), task, total: 0 }
}
Some(Received::Requests(_)) => continue 'restart,
None => return Ok(None),
};
}
match &mut self.data {
DataSource::None => unreachable!(),
DataSource::Buffered(None) => {
self.data = DataSource::None;
continue 'restart;
}
DataSource::Buffered(Some(data)) => {
if data.remaining() > self.max_item_size {
self.data = DataSource::None;
return Err(RecvError::MaxItemSizeExceeded);
}
let pdf_ref =
PortDeserializer::start(self.receiver.port_allocator(), self.receiver.storage());
let item_res = self.deserializer.deserialize(&mut data.reader());
self.data = DataSource::None;
self.item = Some(item_res?);
self.port_deser = Some(PortDeserializer::finish(pdf_ref));
}
DataSource::Streamed { tx, task, total } => {
enum FeedError {
RecvChunkError(RecvChunkError),
MaxItemSizeExceeded,
}
if let Some(tx) = &tx {
let res = loop {
let tx_permit = match tx.reserve().await {
Ok(tx_permit) => tx_permit,
_ => {
break Ok(());
}
};
match self.receiver.recv_chunk().await {
Ok(Some(chunk)) => {
*total += chunk.remaining();
if *total > self.max_item_size {
break Err(FeedError::MaxItemSizeExceeded);
}
tx_permit.send(Ok(chunk));
}
Ok(None) => break Ok(()),
Err(err) => break Err(FeedError::RecvChunkError(err)),
}
};
match res {
Ok(()) => (),
Err(FeedError::RecvChunkError(RecvChunkError::Cancelled)) => {
self.data = DataSource::None;
continue 'restart;
}
Err(FeedError::RecvChunkError(RecvChunkError::ChMux)) => {
self.data = DataSource::None;
return Err(RecvError::Receive(chmux::RecvError::ChMux));
}
Err(FeedError::MaxItemSizeExceeded) => {
self.data = DataSource::None;
return Err(RecvError::MaxItemSizeExceeded);
}
}
}
*tx = None;
match task.await {
Ok(Ok((item, pds))) => {
self.item = Some(item);
self.port_deser = Some(pds);
self.data = DataSource::None;
}
Ok(Err(err)) => {
self.data = DataSource::None;
return Err(RecvError::Deserialize(err));
}
Err(err) => {
self.data = DataSource::None;
match err.try_into_panic() {
Ok(payload) => panic::resume_unwind(payload),
Err(err) => {
return Err(RecvError::Deserialize(DeserializationError::new(err)));
}
}
}
}
}
}
}
let pds = self.port_deser.as_mut().unwrap();
if !pds.expected.is_empty() {
self.receiver.set_max_ports(pds.expected.len() + self.default_max_ports.unwrap());
let requests = match self.receiver.recv_any().await? {
Some(chmux::Received::Requests(requests)) => requests,
other => {
self.recved = Some(other);
self.data = DataSource::None;
self.item = None;
self.port_deser = None;
continue 'restart;
}
};
for request in requests {
if let Some((local_port, callback)) = pds.expected.remove(&request.id()) {
exec::spawn(callback(local_port, request).in_current_span());
}
}
if !pds.expected.is_empty() {
return Err(RecvError::MissingPorts(pds.expected.keys().copied().collect()));
}
}
for task in pds.tasks.drain(..) {
exec::spawn(task.in_current_span());
}
return Ok(Some(self.item.take().unwrap()));
}
}
pub async fn close(&mut self) {
self.receiver.close().await
}
pub fn max_item_size(&self) -> usize {
self.max_item_size
}
pub fn set_max_item_size(&mut self, max_item_size: usize) {
self.max_item_size = max_item_size;
}
}