//! Define a [`CompoundRuntime`] part that can be built from several component //! pieces. use std::{net, sync::Arc, time::Duration}; use crate::traits::*; use crate::{CoarseInstant, CoarseTimeProvider}; use async_trait::async_trait; use educe::Educe; use futures::{future::FutureObj, task::Spawn}; use std::future::Future; use std::io::Result as IoResult; use tor_general_addr::unix; use tracing::instrument; use web_time_compat::{Instant, SystemTime}; /// A runtime made of several parts, each of which implements one trait-group. /// /// The `TaskR` component should implement [`Spawn`], [`Blocking`] and maybe [`ToplevelBlockOn`]; /// the `SleepR` component should implement [`SleepProvider`]; /// the `CoarseTimeR` component should implement [`CoarseTimeProvider`]; /// the `TcpR` component should implement [`NetStreamProvider`] for [`net::SocketAddr`]; /// the `UnixR` component should implement [`NetStreamProvider`] for [`unix::SocketAddr`]; /// and /// the `TlsR` component should implement [`TlsProvider`]. /// /// You can use this structure to create new runtimes in two ways: either by /// overriding a single part of an existing runtime, or by building an entirely /// new runtime from pieces. #[derive(Educe)] #[educe(Clone)] // #[derive(Clone)] wrongly infers Clone bounds on the generic parameters pub struct CompoundRuntime { /// The actual collection of Runtime objects. /// /// We wrap this in an Arc rather than requiring that each item implement /// Clone, though we could change our minds later on. inner: Arc>, } /// A collection of objects implementing that traits that make up a [`Runtime`] struct Inner { /// A `Spawn` and `BlockOn` implementation. spawn: TaskR, /// A `SleepProvider` implementation. sleep: SleepR, /// A `CoarseTimeProvider`` implementation. coarse_time: CoarseTimeR, /// A `NetStreamProvider` implementation tcp: TcpR, /// A `NetStreamProvider` implementation. unix: UnixR, /// A `TlsProvider` implementation. tls: TlsR, /// A `UdpProvider` implementation udp: UdpR, } impl CompoundRuntime { /// Construct a new CompoundRuntime from its components. pub fn new( spawn: TaskR, sleep: SleepR, coarse_time: CoarseTimeR, tcp: TcpR, unix: UnixR, tls: TlsR, udp: UdpR, ) -> Self { #[allow(clippy::arc_with_non_send_sync)] CompoundRuntime { inner: Arc::new(Inner { spawn, sleep, coarse_time, tcp, unix, tls, udp, }), } } } impl Spawn for CompoundRuntime where TaskR: Spawn, { #[inline] #[track_caller] fn spawn_obj(&self, future: FutureObj<'static, ()>) -> Result<(), futures::task::SpawnError> { self.inner.spawn.spawn_obj(future) } } impl Blocking for CompoundRuntime where TaskR: Blocking, SleepR: Clone + Send + Sync + 'static, CoarseTimeR: Clone + Send + Sync + 'static, TcpR: Clone + Send + Sync + 'static, UnixR: Clone + Send + Sync + 'static, TlsR: Clone + Send + Sync + 'static, UdpR: Clone + Send + Sync + 'static, { type ThreadHandle = TaskR::ThreadHandle; #[inline] #[track_caller] fn spawn_blocking(&self, f: F) -> TaskR::ThreadHandle where F: FnOnce() -> T + Send + 'static, T: Send + 'static, { self.inner.spawn.spawn_blocking(f) } #[inline] #[track_caller] fn reenter_block_on(&self, future: F) -> F::Output where F: Future, F::Output: Send + 'static, { self.inner.spawn.reenter_block_on(future) } #[track_caller] fn blocking_io(&self, f: F) -> impl futures::Future where F: FnOnce() -> T + Send + 'static, T: Send + 'static, { self.inner.spawn.blocking_io(f) } } impl ToplevelBlockOn for CompoundRuntime where TaskR: ToplevelBlockOn, SleepR: Clone + Send + Sync + 'static, CoarseTimeR: Clone + Send + Sync + 'static, TcpR: Clone + Send + Sync + 'static, UnixR: Clone + Send + Sync + 'static, TlsR: Clone + Send + Sync + 'static, UdpR: Clone + Send + Sync + 'static, { #[inline] #[track_caller] fn block_on(&self, future: F) -> F::Output { self.inner.spawn.block_on(future) } } impl SleepProvider for CompoundRuntime where SleepR: SleepProvider, TaskR: Clone + Send + Sync + 'static, CoarseTimeR: Clone + Send + Sync + 'static, TcpR: Clone + Send + Sync + 'static, UnixR: Clone + Send + Sync + 'static, TlsR: Clone + Send + Sync + 'static, UdpR: Clone + Send + Sync + 'static, { type SleepFuture = SleepR::SleepFuture; #[inline] fn sleep(&self, duration: Duration) -> Self::SleepFuture { self.inner.sleep.sleep(duration) } #[inline] fn now(&self) -> Instant { self.inner.sleep.now() } #[inline] fn wallclock(&self) -> SystemTime { self.inner.sleep.wallclock() } } impl CoarseTimeProvider for CompoundRuntime where CoarseTimeR: CoarseTimeProvider, SleepR: Clone + Send + Sync + 'static, TaskR: Clone + Send + Sync + 'static, CoarseTimeR: Clone + Send + Sync + 'static, TcpR: Clone + Send + Sync + 'static, UnixR: Clone + Send + Sync + 'static, TlsR: Clone + Send + Sync + 'static, UdpR: Clone + Send + Sync + 'static, { #[inline] fn now_coarse(&self) -> CoarseInstant { self.inner.coarse_time.now_coarse() } } #[async_trait] impl NetStreamProvider for CompoundRuntime where TcpR: NetStreamProvider, TaskR: Send + Sync + 'static, SleepR: Send + Sync + 'static, CoarseTimeR: Send + Sync + 'static, TcpR: Send + Sync + 'static, UnixR: Clone + Send + Sync + 'static, TlsR: Send + Sync + 'static, UdpR: Send + Sync + 'static, { type Stream = TcpR::Stream; type Listener = TcpR::Listener; type ConnectOptions = TcpR::ConnectOptions; type ListenOptions = TcpR::ListenOptions; #[inline] #[instrument(skip_all, level = "trace")] async fn connect( &self, addr: &net::SocketAddr, options: &Self::ConnectOptions, ) -> IoResult { self.inner.tcp.connect(addr, options).await } #[inline] async fn listen( &self, addr: &net::SocketAddr, options: &Self::ListenOptions, ) -> IoResult { self.inner.tcp.listen(addr, options).await } } #[async_trait] impl NetStreamProvider for CompoundRuntime where UnixR: NetStreamProvider, TaskR: Send + Sync + 'static, SleepR: Send + Sync + 'static, CoarseTimeR: Send + Sync + 'static, TcpR: Send + Sync + 'static, UnixR: Clone + Send + Sync + 'static, TlsR: Send + Sync + 'static, UdpR: Send + Sync + 'static, { type Stream = UnixR::Stream; type Listener = UnixR::Listener; type ConnectOptions = UnixR::ConnectOptions; type ListenOptions = UnixR::ListenOptions; #[inline] #[instrument(skip_all, level = "trace")] async fn connect( &self, addr: &unix::SocketAddr, options: &Self::ConnectOptions, ) -> IoResult { self.inner.unix.connect(addr, options).await } #[inline] async fn listen( &self, addr: &unix::SocketAddr, options: &Self::ListenOptions, ) -> IoResult { self.inner.unix.listen(addr, options).await } } impl TlsProvider for CompoundRuntime where TcpR: NetStreamProvider, TlsR: TlsProvider, UnixR: Clone + Send + Sync + 'static, SleepR: Clone + Send + Sync + 'static, CoarseTimeR: Clone + Send + Sync + 'static, TaskR: Clone + Send + Sync + 'static, UdpR: Clone + Send + Sync + 'static, S: StreamOps, { type Connector = TlsR::Connector; type TlsStream = TlsR::TlsStream; type Acceptor = TlsR::Acceptor; type TlsServerStream = TlsR::TlsServerStream; #[inline] fn tls_connector(&self) -> Self::Connector { self.inner.tls.tls_connector() } #[inline] fn tls_acceptor(&self, settings: TlsAcceptorSettings) -> IoResult { self.inner.tls.tls_acceptor(settings) } #[inline] fn supports_keying_material_export(&self) -> bool { self.inner.tls.supports_keying_material_export() } } impl std::fmt::Debug for CompoundRuntime { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("CompoundRuntime").finish_non_exhaustive() } } #[async_trait] impl UdpProvider for CompoundRuntime where UdpR: UdpProvider, TaskR: Send + Sync + 'static, SleepR: Send + Sync + 'static, CoarseTimeR: Send + Sync + 'static, TcpR: Send + Sync + 'static, UnixR: Clone + Send + Sync + 'static, TlsR: Send + Sync + 'static, UdpR: Send + Sync + 'static, { type UdpSocket = UdpR::UdpSocket; #[inline] async fn bind(&self, addr: &net::SocketAddr) -> IoResult { self.inner.udp.bind(addr).await } } /// Module to seal RuntimeSubstExt mod sealed { /// Helper for sealing RuntimeSubstExt #[allow(unreachable_pub)] pub trait Sealed {} } /// Extension trait on Runtime: /// Construct new Runtimes that replace part of an original runtime. /// /// (If you need to do more complicated versions of this, you should likely construct /// CompoundRuntime directly.) pub trait RuntimeSubstExt: sealed::Sealed + Sized { /// Return a new runtime wrapping this runtime, but replacing its TCP NetStreamProvider. fn with_tcp_provider( &self, new_tcp: T, ) -> CompoundRuntime; /// Return a new runtime wrapping this runtime, but replacing its SleepProvider. fn with_sleep_provider( &self, new_sleep: T, ) -> CompoundRuntime; /// Return a new runtime wrapping this runtime, but replacing its CoarseTimeProvider. fn with_coarse_time_provider( &self, new_coarse_time: T, ) -> CompoundRuntime; } impl sealed::Sealed for R {} impl RuntimeSubstExt for R { fn with_tcp_provider( &self, new_tcp: T, ) -> CompoundRuntime { CompoundRuntime::new( self.clone(), self.clone(), self.clone(), new_tcp, self.clone(), self.clone(), self.clone(), ) } fn with_sleep_provider( &self, new_sleep: T, ) -> CompoundRuntime { CompoundRuntime::new( self.clone(), new_sleep, self.clone(), self.clone(), self.clone(), self.clone(), self.clone(), ) } fn with_coarse_time_provider( &self, new_coarse_time: T, ) -> CompoundRuntime { CompoundRuntime::new( self.clone(), self.clone(), new_coarse_time, self.clone(), self.clone(), self.clone(), self.clone(), ) } }