use futures::future::BoxFuture;
use std::{
fmt::Debug,
future::Future,
marker::PhantomData,
pin::Pin,
task::{Context, Poll},
};
use crate::{Channel, Result, role::Role};
pub(crate) struct FinishControl {
hook: Option<Box<dyn FnOnce() + Send + 'static>>,
}
impl FinishControl {
fn new(hook: impl FnOnce() + Send + 'static) -> Self {
Self {
hook: Some(Box::new(hook)),
}
}
pub(crate) fn request(&mut self) {
if let Some(hook) = self.hook.take() {
hook();
}
}
}
#[must_use = "connection drivers must be polled to make progress"]
pub struct ConnectionDriver {
future: BoxFuture<'static, Result<()>>,
finish: Option<FinishControl>,
}
impl ConnectionDriver {
pub fn new(future: impl Future<Output = Result<()>> + Send + 'static) -> Self {
Self {
future: Box::pin(future),
finish: None,
}
}
pub fn with_finish(
future: impl Future<Output = Result<()>> + Send + 'static,
finish: impl FnOnce() + Send + 'static,
) -> Self {
Self {
future: Box::pin(future),
finish: Some(FinishControl::new(finish)),
}
}
pub fn map_future<F>(self, map: impl FnOnce(BoxFuture<'static, Result<()>>) -> F) -> Self
where
F: Future<Output = Result<()>> + Send + 'static,
{
Self {
future: Box::pin(map(self.future)),
finish: self.finish,
}
}
#[must_use]
pub fn request_finish(&mut self) -> bool {
if let Some(finish) = self.finish.as_mut() {
finish.request();
true
} else {
false
}
}
pub(crate) fn take_finish(&mut self) -> Option<FinishControl> {
self.finish.take()
}
}
impl Future for ConnectionDriver {
type Output = Result<()>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
self.future.as_mut().poll(cx)
}
}
impl Debug for ConnectionDriver {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ConnectionDriver")
.field("finishable", &self.finish.is_some())
.finish_non_exhaustive()
}
}
#[cfg_attr(
all(feature = "process", not(target_family = "wasm")),
doc = "- **[`AcpAgent`]**: An external agent running in a separate process with stdio communication"
)]
#[cfg_attr(
all(feature = "process", not(target_family = "wasm")),
doc = "[`AcpAgent`]: crate::AcpAgent"
)]
pub trait ConnectTo<R: Role>: Send + 'static {
fn connect_to(
self,
client: impl ConnectTo<R::Counterpart>,
) -> impl Future<Output = Result<()>> + Send;
fn into_channel_and_future(self) -> (Channel, Option<ConnectionDriver>)
where
Self: Sized,
{
let (channel_a, channel_b) = Channel::duplex();
let future = ConnectionDriver::new(self.connect_to(channel_b));
(channel_a, Some(future))
}
}
trait ErasedConnectTo<R: Role>: Send {
fn type_name(&self) -> &'static str;
fn connect_to_erased(
self: Box<Self>,
client: Box<dyn ErasedConnectTo<R::Counterpart>>,
) -> BoxFuture<'static, Result<()>>;
fn into_channel_and_future_erased(self: Box<Self>) -> (Channel, Option<ConnectionDriver>);
}
impl<C: ConnectTo<R>, R: Role> ErasedConnectTo<R> for C {
fn type_name(&self) -> &'static str {
std::any::type_name::<C>()
}
fn connect_to_erased(
self: Box<Self>,
client: Box<dyn ErasedConnectTo<R::Counterpart>>,
) -> BoxFuture<'static, Result<()>> {
Box::pin(async move {
(*self)
.connect_to(DynConnectTo {
inner: client,
_marker: PhantomData,
})
.await
})
}
fn into_channel_and_future_erased(self: Box<Self>) -> (Channel, Option<ConnectionDriver>) {
(*self).into_channel_and_future()
}
}
pub struct DynConnectTo<R: Role> {
inner: Box<dyn ErasedConnectTo<R>>,
_marker: PhantomData<R>,
}
impl<R: Role> DynConnectTo<R> {
pub fn new<C: ConnectTo<R>>(component: C) -> Self {
Self {
inner: Box::new(component),
_marker: PhantomData,
}
}
#[must_use]
pub fn type_name(&self) -> &'static str {
self.inner.type_name()
}
}
impl<R: Role> ConnectTo<R> for DynConnectTo<R> {
async fn connect_to(self, client: impl ConnectTo<R::Counterpart>) -> Result<()> {
self.inner
.connect_to_erased(Box::new(client) as Box<dyn ErasedConnectTo<R::Counterpart>>)
.await
}
fn into_channel_and_future(self) -> (Channel, Option<ConnectionDriver>) {
self.inner.into_channel_and_future_erased()
}
}
impl<R: Role> Debug for DynConnectTo<R> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("DynConnectTo")
.field("type_name", &self.type_name())
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::role::UntypedRole;
use futures::FutureExt as _;
struct OwnedWork(BoxFuture<'static, Result<()>>);
impl ConnectTo<UntypedRole> for OwnedWork {
async fn connect_to(self, _client: impl ConnectTo<UntypedRole>) -> Result<()> {
self.0.await
}
}
#[test]
fn raw_channel_has_no_owned_work() {
let (channel, _other) = Channel::duplex();
let (_, driver) = ConnectTo::<UntypedRole>::into_channel_and_future(channel);
assert!(driver.is_none());
}
#[test]
fn owned_driver_preserves_errors_and_polls_unpinned() {
let error = crate::Error::internal_error().data("driver failure");
let mut driver = ConnectionDriver::new(futures::future::ready(Err(error.clone())));
assert_eq!(futures::executor::block_on(&mut driver), Err(error));
}
#[test]
fn finish_request_is_idempotent_and_does_not_mean_completion() {
let (finish_tx, finish_rx) = futures::channel::oneshot::channel();
let (flushed_tx, flushed_rx) = futures::channel::oneshot::channel();
let mut driver = ConnectionDriver::with_finish(
async move {
finish_rx.await.unwrap();
flushed_rx.await.unwrap()
},
move || finish_tx.send(()).unwrap(),
);
assert!((&mut driver).now_or_never().is_none());
assert!(driver.request_finish());
assert!(driver.request_finish());
assert!((&mut driver).now_or_never().is_none());
let error = crate::Error::internal_error().data("custom flush failed");
flushed_tx.send(Err(error.clone())).unwrap();
assert_eq!(futures::executor::block_on(driver), Err(error));
}
#[test]
fn opaque_driver_cannot_be_cooperatively_finished() {
let mut driver = ConnectionDriver::new(futures::future::pending());
assert!(!driver.request_finish());
assert!((&mut driver).now_or_never().is_none());
}
#[test]
fn future_decoration_preserves_finish_and_completion_errors() {
use std::sync::{
Arc,
atomic::{AtomicUsize, Ordering},
};
for request_before_wrapping in [false, true] {
let (finish_tx, finish_rx) = futures::channel::oneshot::channel();
let (flush_tx, flush_rx) = futures::channel::oneshot::channel();
let calls = Arc::new(AtomicUsize::new(0));
let hook_calls = calls.clone();
let mut driver = ConnectionDriver::with_finish(
async move {
finish_rx.await.unwrap();
flush_rx.await.unwrap()
},
move || {
hook_calls.fetch_add(1, Ordering::SeqCst);
finish_tx.send(()).unwrap();
},
);
if request_before_wrapping {
assert!(driver.request_finish());
}
let observed = Arc::new(AtomicUsize::new(0));
let observe_completion = observed.clone();
let mut decorated = driver.map_future(|work| {
work.inspect(move |_| {
observe_completion.fetch_add(1, Ordering::SeqCst);
})
});
assert!(decorated.request_finish());
assert!(decorated.request_finish());
assert_eq!(calls.load(Ordering::SeqCst), 1);
assert!((&mut decorated).now_or_never().is_none());
assert_eq!(observed.load(Ordering::SeqCst), 0);
let error = crate::Error::internal_error().data("decorated flush failed");
flush_tx.send(Err(error.clone())).unwrap();
assert_eq!(futures::executor::block_on(decorated), Err(error));
assert_eq!(observed.load(Ordering::SeqCst), 1);
}
}
#[test]
fn future_decoration_does_not_make_opaque_work_cooperative() {
let driver = ConnectionDriver::new(futures::future::pending());
let mut decorated = driver.map_future(|work| work);
assert!(!decorated.request_finish());
assert!((&mut decorated).now_or_never().is_none());
}
#[test]
fn dropping_driver_does_not_invoke_finish_hook() {
let invoked = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false));
let hook_invoked = invoked.clone();
let driver = ConnectionDriver::with_finish(futures::future::pending(), move || {
hook_invoked.store(true, std::sync::atomic::Ordering::Release);
});
drop(driver);
assert!(!invoked.load(std::sync::atomic::Ordering::Acquire));
}
#[test]
fn default_conversion_owns_real_work_until_completion() {
let (done_tx, done_rx) = futures::channel::oneshot::channel();
let component = OwnedWork(async move { done_rx.await.unwrap() }.boxed());
let (_channel, driver) = component.into_channel_and_future();
let mut driver = driver.expect("default conversion always owns its connect_to work");
assert!((&mut driver).now_or_never().is_none());
let error = crate::Error::internal_error().data("owned work failed");
done_tx.send(Err(error.clone())).unwrap();
assert_eq!(futures::executor::block_on(driver), Err(error));
}
#[test]
fn dropping_optional_owned_driver_cancels_unpolled_work() {
let (done_tx, done_rx) = futures::channel::oneshot::channel::<Result<()>>();
let component = OwnedWork(async move { done_rx.await.unwrap() }.boxed());
let (_channel, driver) = component.into_channel_and_future();
assert!(driver.is_some());
assert!(!done_tx.is_canceled());
drop(driver);
assert!(done_tx.is_canceled());
}
#[test]
fn type_erasure_preserves_owned_work_and_finish_metadata() {
let outgoing = futures::sink::unfold((), |(), _line: String| {
futures::future::ready(Ok::<_, std::io::Error>(()))
});
let incoming = futures::stream::pending::<std::io::Result<String>>();
let component = DynConnectTo::<UntypedRole>::new(crate::Lines::new(outgoing, incoming));
let (_channel, driver) = component.into_channel_and_future();
let mut driver = driver.expect("erasure must retain ownership");
assert!((&mut driver).now_or_never().is_none());
assert!(
driver.request_finish(),
"erasure must retain finish coordination"
);
futures::executor::block_on(driver).unwrap();
}
#[test]
fn type_erasure_preserves_passive_lifetime() {
let (channel, _other) = Channel::duplex();
let (_, driver) = DynConnectTo::<UntypedRole>::new(channel).into_channel_and_future();
assert!(driver.is_none());
}
#[test]
fn dyn_connect_to_reports_static_type_name_and_correct_debug_label() {
let (channel, _other) = Channel::duplex();
let component = DynConnectTo::<UntypedRole>::new(channel);
let type_name: &'static str = component.type_name();
assert_eq!(type_name, std::any::type_name::<Channel>());
assert_eq!(
format!("{component:?}"),
format!("DynConnectTo {{ type_name: {type_name:?} }}")
);
}
}