use std::fmt;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
#[derive(Clone)]
pub struct Partition {
id: PartitionId,
spawner: DriverSpawner,
interface: Option<Arc<str>>,
}
impl Partition {
pub fn new(id: PartitionId, spawner: DriverSpawner) -> Self {
Self {
id,
spawner,
interface: None,
}
}
#[cfg(any(
target_os = "android",
target_os = "fuchsia",
target_os = "illumos",
target_os = "ios",
target_os = "linux",
target_os = "macos",
target_os = "solaris",
target_os = "tvos",
target_os = "visionos",
target_os = "watchos",
))]
#[cfg_attr(
docsrs,
doc(cfg(any(
target_os = "android",
target_os = "fuchsia",
target_os = "illumos",
target_os = "ios",
target_os = "linux",
target_os = "macos",
target_os = "solaris",
target_os = "tvos",
target_os = "visionos",
target_os = "watchos",
)))
)]
pub fn interface(mut self, interface: impl Into<String>) -> Self {
self.interface = Some(Arc::from(interface.into()));
self
}
#[cfg(any(
target_os = "android",
target_os = "fuchsia",
target_os = "illumos",
target_os = "ios",
target_os = "linux",
target_os = "macos",
target_os = "solaris",
target_os = "tvos",
target_os = "visionos",
target_os = "watchos",
))]
pub(super) fn interface_name(&self) -> Option<&str> {
self.interface.as_deref()
}
pub(super) fn id(&self) -> PartitionId {
self.id
}
pub(super) fn into_parts(self) -> (PartitionId, DriverSpawner, Option<Arc<str>>) {
(self.id, self.spawner, self.interface)
}
}
impl fmt::Debug for Partition {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Partition")
.field("id", &self.id)
.field("spawner", &self.spawner)
.field("interface", &self.interface)
.finish()
}
}
#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
pub struct PartitionId(usize);
impl PartitionId {
pub const ANONYMOUS: Self = Self(usize::MAX);
pub const fn from_index(index: usize) -> Self {
Self(index)
}
pub const fn is_anonymous(self) -> bool {
self.0 == Self::ANONYMOUS.0
}
}
#[derive(Clone)]
pub struct DriverSpawner {
inner: Arc<dyn Spawn>,
}
impl DriverSpawner {
#[cfg(feature = "rt-tokio")]
#[cfg_attr(docsrs, doc(cfg(feature = "rt-tokio")))]
pub fn tokio(handle: tokio::runtime::Handle) -> Self {
Self::new(TokioSpawn { handle })
}
#[cfg(feature = "test-util")]
#[doc(hidden)]
pub fn from_fn<F>(spawn: F) -> Self
where
F: Fn(Pin<Box<dyn Future<Output = ()> + Send + 'static>>) + Send + Sync + 'static,
{
Self::new(FnSpawn(spawn))
}
#[cfg_attr(
not(any(feature = "rt-tokio", feature = "test-util", test)),
allow(
dead_code,
reason = "no public constructor exists without a runtime feature"
)
)]
pub(in crate::client::pool) fn new(spawn: impl Spawn) -> Self {
Self {
inner: Arc::new(spawn),
}
}
pub(in crate::client::pool) fn spawn(
&self,
driver: Pin<Box<dyn Future<Output = ()> + Send + 'static>>,
) {
self.inner.spawn(driver);
}
#[cfg(test)]
pub(in crate::client::pool) fn ptr_eq(&self, other: &Self) -> bool {
Arc::ptr_eq(&self.inner, &other.inner)
}
}
impl fmt::Debug for DriverSpawner {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_tuple("DriverSpawner").field(&self.inner).finish()
}
}
pub(in crate::client::pool) trait Spawn:
fmt::Debug + Send + Sync + 'static
{
fn spawn(&self, driver: Pin<Box<dyn Future<Output = ()> + Send + 'static>>);
}
#[cfg(feature = "rt-tokio")]
#[derive(Debug)]
struct TokioSpawn {
handle: tokio::runtime::Handle,
}
#[cfg(feature = "rt-tokio")]
impl Spawn for TokioSpawn {
fn spawn(&self, driver: Pin<Box<dyn Future<Output = ()> + Send + 'static>>) {
drop(self.handle.spawn(driver));
}
}
#[cfg(feature = "test-util")]
struct FnSpawn<F>(F);
#[cfg(feature = "test-util")]
impl<F> fmt::Debug for FnSpawn<F> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("FnSpawn")
}
}
#[cfg(feature = "test-util")]
impl<F> Spawn for FnSpawn<F>
where
F: Fn(Pin<Box<dyn Future<Output = ()> + Send + 'static>>) + Send + Sync + 'static,
{
fn spawn(&self, driver: Pin<Box<dyn Future<Output = ()> + Send + 'static>>) {
(self.0)(driver);
}
}
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
#[non_exhaustive]
pub enum ConnectionReuseScope {
Partition,
#[default]
NetworkInterface,
Pool,
}
#[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
pub(in crate::client::pool) enum EligibilityGroup {
Partition(PartitionId),
NetworkInterface(Option<Arc<str>>),
Pool,
}
#[cfg(test)]
mod tests {
use super::*;
#[cfg(feature = "rt-tokio")]
use std::sync::atomic::{AtomicBool, Ordering};
#[derive(Debug)]
struct TestSpawner;
impl Spawn for TestSpawner {
fn spawn(&self, driver: Pin<Box<dyn Future<Output = ()> + Send + 'static>>) {
drop(driver);
}
}
#[test]
fn anonymous_identity_is_reserved() {
assert!(PartitionId::ANONYMOUS.is_anonymous());
assert!(PartitionId::from_index(usize::MAX).is_anonymous());
assert!(!PartitionId::from_index(0).is_anonymous());
}
#[test]
fn partition_retains_placement() {
let spawner = DriverSpawner::new(TestSpawner);
let partition = Partition::new(PartitionId::from_index(7), spawner.clone());
let (id, retained, interface) = partition.into_parts();
assert_eq!(PartitionId::from_index(7), id);
assert_eq!(None, interface);
assert!(spawner.ptr_eq(&retained));
let driver = Box::pin(async {});
retained.spawn(driver);
}
#[cfg(feature = "test-util")]
#[test]
fn from_fn_receives_each_driver() {
let submitted = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let observed = submitted.clone();
let spawner = DriverSpawner::from_fn(move |driver| {
observed.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
drop(driver);
});
spawner.spawn(Box::pin(async {}));
spawner.spawn(Box::pin(async {}));
assert_eq!(2, submitted.load(std::sync::atomic::Ordering::SeqCst));
assert_eq!("DriverSpawner(FnSpawn)", format!("{spawner:?}"));
}
#[cfg(feature = "rt-tokio")]
#[test]
fn tokio_spawner_uses_captured_runtime() {
let owner = tokio::runtime::Builder::new_current_thread()
.build()
.unwrap();
let owner_id = owner.handle().id();
let ran = Arc::new(AtomicBool::new(false));
let task_ran = ran.clone();
let spawner = DriverSpawner::tokio(owner.handle().clone());
owner.block_on(async {
spawner.spawn(Box::pin(async move {
assert_eq!(owner_id, tokio::runtime::Handle::current().id());
task_ran.store(true, Ordering::SeqCst);
}));
for _ in 0..10 {
if ran.load(Ordering::SeqCst) {
break;
}
tokio::task::yield_now().await;
}
});
assert!(ran.load(Ordering::SeqCst));
}
#[cfg(feature = "rt-tokio")]
#[test]
fn tokio_spawner_accepts_work_from_a_foreign_runtime() {
let owner = tokio::runtime::Builder::new_current_thread()
.build()
.unwrap();
let foreign = tokio::runtime::Builder::new_current_thread()
.build()
.unwrap();
let owner_id = owner.handle().id();
let ran = Arc::new(AtomicBool::new(false));
let task_ran = ran.clone();
let spawner = DriverSpawner::tokio(owner.handle().clone());
foreign.block_on(async {
spawner.spawn(Box::pin(async move {
assert_eq!(owner_id, tokio::runtime::Handle::current().id());
task_ran.store(true, Ordering::SeqCst);
}));
});
owner.block_on(async {
for _ in 0..10 {
if ran.load(Ordering::SeqCst) {
break;
}
tokio::task::yield_now().await;
}
});
assert!(ran.load(Ordering::SeqCst));
}
}