use futures::{Future, future, pin_mut};
use serde::{Deserialize, Serialize};
use std::{fmt, sync::Arc};
use tokio::sync::Semaphore;
use tracing::Instrument;
use super::{CallError, msg::RFnRequest};
use crate::{
RemoteSend, codec, exec,
rch::{mpsc, oneshot},
};
pub struct RFnProvider {
keep_tx: Option<tokio::sync::oneshot::Sender<()>>,
max_concurrency_tx: tokio::sync::watch::Sender<usize>,
}
impl fmt::Debug for RFnProvider {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
f.debug_struct("RFnProvider").finish()
}
}
impl RFnProvider {
pub const DEFAULT_MAX_CONCURRENCY: usize = 32;
pub fn keep(mut self) {
let _ = self.keep_tx.take().unwrap().send(());
}
pub async fn done(&mut self) {
self.keep_tx.as_mut().unwrap().closed().await
}
pub fn max_concurrency(&self) -> usize {
*self.max_concurrency_tx.borrow()
}
pub fn set_max_concurrency(&self, limit: usize) {
self.max_concurrency_tx.send_if_modified(|current| {
if *current != limit {
*current = limit;
true
} else {
false
}
});
}
}
impl Drop for RFnProvider {
fn drop(&mut self) {
}
}
#[derive(Serialize, Deserialize)]
#[serde(bound(serialize = "A: RemoteSend, R: RemoteSend, Codec: codec::Codec"))]
#[serde(bound(deserialize = "A: RemoteSend, R: RemoteSend, Codec: codec::Codec"))]
pub struct RFn<A, R, Codec = codec::Default> {
request_tx: mpsc::Sender<RFnRequest<A, R, Codec>, Codec, 1>,
}
impl<A, R, Codec> Clone for RFn<A, R, Codec> {
fn clone(&self) -> Self {
Self { request_tx: self.request_tx.clone() }
}
}
impl<A, R, Codec> fmt::Debug for RFn<A, R, Codec> {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
f.debug_struct("RFn").finish()
}
}
impl<A, R, Codec> RFn<A, R, Codec>
where
A: RemoteSend,
R: RemoteSend,
Codec: codec::Codec,
{
fn new_int<F, Fut>(fun: F) -> Self
where
F: Fn(A) -> Fut + Send + Sync + 'static,
Fut: Future<Output = R> + Send,
{
let (rfn, provider) = Self::provided_int(fun);
provider.keep();
rfn
}
fn provided_int<F, Fut>(fun: F) -> (Self, RFnProvider)
where
F: Fn(A) -> Fut + Send + Sync + 'static,
Fut: Future<Output = R> + Send,
{
let (request_tx, request_rx) = mpsc::channel(1);
let request_tx = request_tx.set_buffer();
let mut request_rx = request_rx.set_buffer::<1>();
let (keep_tx, keep_rx) = tokio::sync::oneshot::channel();
let (max_concurrency_tx, mut max_concurrency_rx) =
tokio::sync::watch::channel(RFnProvider::DEFAULT_MAX_CONCURRENCY);
let fun = Arc::new(fun);
exec::spawn(
async move {
let mut semaphore = Arc::new(Semaphore::new(*max_concurrency_rx.borrow_and_update()));
let term = async move {
if let Ok(()) = keep_rx.await {
future::pending().await
}
};
pin_mut!(term);
loop {
tokio::select! {
biased;
() = &mut term => break,
Ok(()) = max_concurrency_rx.changed() => {
semaphore = Arc::new(Semaphore::new(*max_concurrency_rx.borrow_and_update()));
}
req_res = request_rx.recv() => {
match req_res {
Ok(Some(RFnRequest {argument, result_tx})) => {
let fun_task = fun.clone();
let semaphore = semaphore.clone();
exec::spawn(async move {
let _permit = semaphore.acquire().await.ok();
let result = fun_task(argument).await;
let _ = result_tx.send(result);
}.in_current_span());
}
Ok(None) => break,
Err(err) if err.is_final() => break,
Err(_) => (),
}
}
}
}
}
.in_current_span(),
);
(Self { request_tx }, RFnProvider { keep_tx: Some(keep_tx), max_concurrency_tx })
}
async fn try_call_int(&self, argument: A) -> Result<R, CallError> {
let (result_tx, result_rx) = oneshot::channel();
let _ = self.request_tx.send(RFnRequest { argument, result_tx }).await;
let result = result_rx.await?;
Ok(result)
}
}
impl<A, RT, RE, Codec> RFn<A, Result<RT, RE>, Codec>
where
A: RemoteSend,
RT: RemoteSend,
RE: RemoteSend + From<CallError>,
Codec: codec::Codec,
{
async fn call_int(&self, argument: A) -> Result<RT, RE> {
self.try_call_int(argument).await?
}
}
#[rustfmt::skip] arg_stub!(RFn, Fn, RFnProvider, new_0, provided_0, (&), );
#[rustfmt::skip] arg_stub!(RFn, Fn, RFnProvider, new_1, provided_1, (&), arg1: A1);
#[rustfmt::skip] arg_stub!(RFn, Fn, RFnProvider, new_2, provided_2, (&), arg1: A1, arg2: A2);
#[rustfmt::skip] arg_stub!(RFn, Fn, RFnProvider, new_3, provided_3, (&), arg1: A1, arg2: A2, arg3: A3);
#[rustfmt::skip] arg_stub!(RFn, Fn, RFnProvider, new_4, provided_4, (&), arg1: A1, arg2: A2, arg3: A3, arg4: A4);
#[rustfmt::skip] arg_stub!(RFn, Fn, RFnProvider, new_5, provided_5, (&), arg1: A1, arg2: A2, arg3: A3, arg4: A4, arg5: A5);
#[rustfmt::skip] arg_stub!(RFn, Fn, RFnProvider, new_6, provided_6, (&), arg1: A1, arg2: A2, arg3: A3, arg4: A4, arg5: A5, arg6: A6);
#[rustfmt::skip] arg_stub!(RFn, Fn, RFnProvider, new_7, provided_7, (&), arg1: A1, arg2: A2, arg3: A3, arg4: A4, arg5: A5, arg6: A6, arg7: A7);
#[rustfmt::skip] arg_stub!(RFn, Fn, RFnProvider, new_8, provided_8, (&), arg1: A1, arg2: A2, arg3: A3, arg4: A4, arg5: A5, arg6: A6, arg7: A7, arg8: A8);
#[rustfmt::skip] arg_stub!(RFn, Fn, RFnProvider, new_9, provided_9, (&), arg1: A1, arg2: A2, arg3: A3, arg4: A4, arg5: A5, arg6: A6, arg7: A7, arg8: A8, arg9: A9);
#[rustfmt::skip] arg_stub!(RFn, Fn, RFnProvider, new_10, provided_10, (&), arg1: A1, arg2: A2, arg3: A3, arg4: A4, arg5: A5, arg6: A6, arg7: A7, arg8: A8, arg9: A9, arg10: A10);
impl<A, R, Codec> Drop for RFn<A, R, Codec> {
fn drop(&mut self) {
}
}