use bytes::Buf;
use futures::{FutureExt, future::BoxFuture};
use serde::{Deserialize, Serialize};
use std::{
fmt,
future::{self, Future},
ops::Deref,
pin::Pin,
sync::{Arc, Weak},
task::{Context, Poll, ready},
time::Duration,
};
use super::{DEFAULT_MAX_ITEM_SIZE, RemoteSendError, base};
use crate::{
RemoteSend, chmux,
codec::{self, AnySend, ErasedDeserializer, ErasedSerializer},
exec::{
self,
time::{Instant, sleep},
},
rch::{BACKCHANNEL_MSG_ERROR, BACKCHANNEL_MSG_RATE_LIMIT},
};
mod receiver;
mod sender;
pub use receiver::{ChangedError, Receiver, ReceiverStream, RecvError};
pub use sender::{SendError, Sender};
pub struct Ref<'a, T>(tokio::sync::watch::Ref<'a, Result<T, RecvError>>);
impl<T> Deref for Ref<'_, T> {
type Target = T;
fn deref(&self) -> &Self::Target {
self.0.as_ref().unwrap()
}
}
impl<T> fmt::Debug for Ref<'_, T>
where
T: fmt::Debug,
{
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
write!(f, "{:?}", **self)
}
}
pub fn channel<T, Codec>(init: T) -> (Sender<T, Codec>, Receiver<T, Codec>)
where
T: RemoteSend,
{
let (tx, rx) = tokio::sync::watch::channel(Ok(init));
let (remote_send_err_tx, remote_send_err_rx) = tokio::sync::mpsc::unbounded_channel();
let (sender_rate_limit_tx, sender_rate_limit_rx) = tokio::sync::watch::channel(default_rate_limit());
let (receiver_rate_limit_tx, receiver_rate_limit_rx) = rate_limit_channel(default_rate_limit());
let sender = Sender::new(
tx,
remote_send_err_tx.clone(),
remote_send_err_rx,
DEFAULT_MAX_ITEM_SIZE,
sender_rate_limit_tx,
sender_rate_limit_rx.clone(),
receiver_rate_limit_tx.clone(),
receiver_rate_limit_rx,
TransferStrategy::default(),
);
let receiver = Receiver::new(
rx,
remote_send_err_tx,
None,
sender_rate_limit_rx,
receiver_rate_limit_tx,
TransferStrategy::default(),
);
(sender, receiver)
}
pub fn forward<T, Codec>(mut local_rx: tokio::sync::watch::Receiver<T>) -> (Forwarding, Receiver<T, Codec>)
where
T: RemoteSend + Sync + Clone,
Codec: codec::Codec,
{
let init = local_rx.borrow_and_update().clone();
let (mut tx, rx) = channel(init);
let sender_rate_limit_tx = tx.inner.as_ref().unwrap().sender_rate_limit_tx.clone();
let hnd = exec::spawn(async move {
loop {
tokio::select! {
biased;
() = tx.closed() => break,
res = local_rx.changed() => {
match res {
Ok(()) => {
let value = local_rx.borrow_and_update().clone();
match tx.send(value) {
Ok(()) => (),
Err(err) if err.is_closed() => break,
Err(err) => return Err(err),
}
}
Err(_) => break,
}
}
}
}
tx.check()
});
(Forwarding { hnd, sender_rate_limit_tx }, rx)
}
pub struct Forwarding {
hnd: exec::task::JoinHandle<Result<(), SendError>>,
sender_rate_limit_tx: tokio::sync::watch::Sender<Duration>,
}
impl fmt::Debug for Forwarding {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
f.debug_struct("Forwarding").finish()
}
}
impl Future for Forwarding {
type Output = Result<(), SendError>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Self::Output> {
match ready!(self.hnd.poll_unpin(cx)) {
Ok(res) => Poll::Ready(res),
Err(_) => Poll::Ready(Err(SendError::Closed)),
}
}
}
impl Forwarding {
pub fn stop(self) {
self.hnd.abort();
}
pub fn rate_limit(&self) -> Duration {
*self.sender_rate_limit_tx.borrow()
}
pub fn set_rate_limit(&mut self, rate_limit: Duration) {
self.sender_rate_limit_tx.send_replace(rate_limit);
}
}
#[derive(Default, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub enum TransferStrategy {
Single,
GlobalBuffered,
#[default]
#[serde(other)]
ChannelBuffered,
}
pub trait WatchExt<T, Codec, const MAX_ITEM_SIZE: usize> {
fn with_max_item_size<const NEW_MAX_ITEM_SIZE: usize>(
self,
) -> (Sender<T, Codec>, Receiver<T, Codec, NEW_MAX_ITEM_SIZE>);
fn with_transfer_strategy(
self, transfer_strategy: TransferStrategy,
) -> (Sender<T, Codec>, Receiver<T, Codec, MAX_ITEM_SIZE>);
}
impl<T, Codec, const MAX_ITEM_SIZE: usize> WatchExt<T, Codec, MAX_ITEM_SIZE>
for (Sender<T, Codec>, Receiver<T, Codec, MAX_ITEM_SIZE>)
where
T: Send + 'static,
{
fn with_max_item_size<const NEW_MAX_ITEM_SIZE: usize>(
self,
) -> (Sender<T, Codec>, Receiver<T, Codec, NEW_MAX_ITEM_SIZE>) {
let (mut tx, rx) = self;
tx.set_max_item_size(NEW_MAX_ITEM_SIZE);
let rx = rx.set_max_item_size();
(tx, rx)
}
fn with_transfer_strategy(
self, transfer_strategy: TransferStrategy,
) -> (Sender<T, Codec>, Receiver<T, Codec, MAX_ITEM_SIZE>) {
let (mut tx, mut rx) = self;
if let Some(inner) = &mut tx.inner {
inner.transfer_strategy = transfer_strategy.clone();
}
rx.transfer_strategy = transfer_strategy;
(tx, rx)
}
}
trait ErasedWatchRx {
fn borrow_and_update_clone(&mut self) -> AnySend;
fn changed<'a>(&'a mut self) -> BoxFuture<'a, Result<(), tokio::sync::watch::error::RecvError>>;
}
impl<T> ErasedWatchRx for tokio::sync::watch::Receiver<Result<T, RecvError>>
where
T: Clone + Send + Sync + 'static,
{
fn borrow_and_update_clone(&mut self) -> AnySend {
Box::new(self.borrow_and_update().clone())
}
fn changed<'a>(&'a mut self) -> BoxFuture<'a, Result<(), tokio::sync::watch::error::RecvError>> {
self.changed().boxed()
}
}
#[allow(clippy::too_many_arguments)]
async fn send_impl(
erased_serializer: ErasedSerializer, mut rx: Box<dyn ErasedWatchRx + Send>, raw_tx: chmux::Sender,
mut raw_rx: chmux::Receiver, remote_send_err_tx: tokio::sync::mpsc::UnboundedSender<RemoteSendError>,
max_item_size: usize, mut sender_rate_limit_rx: tokio::sync::watch::Receiver<Duration>,
mut receiver_rate_limit_tx: RateLimitSender, strategy: TransferStrategy,
) {
let mut remote_tx = base::ErasedSender::new(erased_serializer, raw_tx);
remote_tx.set_max_item_size(max_item_size);
remote_tx.set_global_credits_use(strategy != TransferStrategy::ChannelBuffered);
let mut last_send: Option<Instant> = None;
let mut send_pending = false;
let mut closed = false;
while !closed {
let rate_limit = sender_rate_limit_rx.borrow_and_update().max(receiver_rate_limit_tx.get());
let pending_send_trigger = async {
if send_pending {
if let Some(last_send) = last_send
&& rate_limit > Duration::ZERO
{
let until = last_send + rate_limit;
let delay = until.duration_since(Instant::now());
if delay > Duration::ZERO {
sleep(delay).await;
}
}
} else {
future::pending().await
}
};
let send = tokio::select! {
biased;
backchannel_msg = raw_rx.recv() => {
match backchannel_msg {
Ok(Some(mut msg)) => {
match msg.try_get_u8() {
Ok(BACKCHANNEL_MSG_ERROR) => {
let _ = remote_send_err_tx.send(RemoteSendError::Forward);
}
Ok(BACKCHANNEL_MSG_RATE_LIMIT) => {
if let Ok(ns) = msg.try_get_u128_le() {
receiver_rate_limit_tx.set(Duration::from_nanos_u128(ns));
}
}
_ => (),
}
}
_ => closed = true,
}
false
}
() = pending_send_trigger => true,
Ok(()) = sender_rate_limit_rx.changed() => false,
changed = rx.changed() => {
match changed {
Ok(()) => send_pending = true,
Err(_) => closed = true,
}
false
}
};
if send || (send_pending && closed) {
let value = rx.borrow_and_update_clone();
if let Err(err) = remote_tx.send_erased(value).await {
let _ = remote_send_err_tx.send(RemoteSendError::Send(err.kind.clone()));
if err.is_item_specific() {
tracing::warn!(%err, "sending over remote channel failed");
break;
}
}
last_send = Some(Instant::now());
send_pending = false;
if strategy == TransferStrategy::Single {
remote_tx.all_received().await;
}
}
}
}
trait ErasedWatchTx {
fn send(&self, value: AnySend) -> Result<(), ()>;
fn send_err(&self, err: RecvError) -> Result<(), ()>;
fn closed(&'_ self) -> BoxFuture<'_, ()>;
}
impl<T> ErasedWatchTx for tokio::sync::watch::Sender<Result<T, RecvError>>
where
T: Clone + Send + Sync + 'static,
{
fn send(&self, value: AnySend) -> Result<(), ()> {
let value: Result<T, RecvError> = *value.downcast().expect("type mismatch in watch receiver");
self.send(value).map_err(|_| ())
}
fn send_err(&self, err: RecvError) -> Result<(), ()> {
let value: Result<T, RecvError> = Err(err);
self.send(value).map_err(|_| ())
}
fn closed(&'_ self) -> BoxFuture<'_, ()> {
self.closed().boxed()
}
}
#[allow(clippy::too_many_arguments)]
async fn recv_impl(
erased_deserializer: ErasedDeserializer, tx: Box<dyn ErasedWatchTx + Send>, mut raw_tx: chmux::Sender,
raw_rx: chmux::Receiver, mut remote_send_err_rx: tokio::sync::mpsc::UnboundedReceiver<RemoteSendError>,
mut current_err: Option<RemoteSendError>, max_item_size: usize, mut rate_limit_rx: RateLimitReceiver,
) {
let mut remote_rx = base::ErasedReceiver::new(erased_deserializer, raw_rx);
remote_rx.set_max_item_size(max_item_size);
let mut rate_limit = None;
loop {
tokio::select! {
biased;
() = tx.closed() => break,
Some(_) = remote_send_err_rx.recv() => {
let _ = raw_tx.send(vec![BACKCHANNEL_MSG_ERROR].into()).await;
}
() = futures::future::ready(()), if current_err.is_some() => {
let _ = raw_tx.send(vec![BACKCHANNEL_MSG_ERROR].into()).await;
current_err = None;
}
Ok(()) = rate_limit_rx.changed() => {
let new_rate_limit = rate_limit_rx.get_and_update();
if rate_limit != Some(new_rate_limit) {
let mut msg = vec![BACKCHANNEL_MSG_RATE_LIMIT];
msg.extend(new_rate_limit.as_nanos().to_le_bytes());
let _ = raw_tx.send(msg.into()).await;
rate_limit = Some(new_rate_limit);
}
}
res = remote_rx.recv_erased() => {
match res {
Ok(Some(value)) => {
if tx.send(value).is_err() {
break;
}
},
Ok(None) => break,
Err(err) => {
let is_final_err = err.is_final();
if tx.send_err(RecvError::RemoteReceive(err)).is_err() {
break
}
if is_final_err {
break;
}
},
}
}
}
}
}
pub(crate) fn rate_limit_channel(rate_limit: Duration) -> (RateLimitSender, RateLimitReceiver) {
let current = Arc::new(rate_limit);
let (tx, rx) = tokio::sync::watch::channel(vec![Arc::downgrade(¤t)]);
(RateLimitSender { tx, current }, RateLimitReceiver(rx))
}
#[derive(Clone)]
pub(crate) struct RateLimitSender {
tx: tokio::sync::watch::Sender<Vec<Weak<Duration>>>,
current: Arc<Duration>,
}
impl RateLimitSender {
pub fn get(&self) -> Duration {
*self.current
}
pub fn set(&mut self, rate_limit: Duration) {
self.current = Arc::new(Duration::ZERO);
let rate_limit = Arc::new(rate_limit);
self.tx.send_modify(|limits| {
limits.retain(|weak| weak.strong_count() > 0);
limits.push(Arc::downgrade(&rate_limit));
});
self.current = rate_limit
}
}
impl Drop for RateLimitSender {
fn drop(&mut self) {
self.current = Arc::new(Duration::ZERO);
self.tx.send_modify(|limits| limits.retain(|weak| weak.strong_count() > 0));
}
}
pub(crate) struct RateLimitReceiver(tokio::sync::watch::Receiver<Vec<Weak<Duration>>>);
impl RateLimitReceiver {
pub async fn changed(&mut self) -> Result<(), tokio::sync::watch::error::RecvError> {
self.0.changed().await
}
fn compute(weaks: &[Weak<Duration>]) -> Duration {
weaks.iter().filter_map(|weak| weak.upgrade()).map(|limit| *limit).min().unwrap_or_default()
}
pub fn get(&self) -> Duration {
Self::compute(&self.0.borrow())
}
pub fn get_and_update(&mut self) -> Duration {
Self::compute(&self.0.borrow_and_update())
}
}
const fn default_max_item_size() -> u64 {
u64::MAX
}
const fn default_rate_limit() -> Duration {
Duration::ZERO
}