use std::{any::Any, sync::Arc};
use crate::{DriverError, LibsyError, Result};
use parking_lot::Mutex;
use futures::{Stream, StreamExt};
use switchyard_protocol::Context;
use tokio::sync::{mpsc, oneshot};
use tokio::time::{Duration, timeout};
use tokio_stream::wrappers::ReceiverStream;
type BoxAny = Box<dyn Any + Send>;
type StepResult = Result<DriverStep>;
const FULFILL_REQUEST_TIMEOUT: Duration = Duration::from_mins(10);
pub enum DriverStep {
Request(DriverRequest),
Info(BoxAny),
Done(BoxAny),
}
pub struct DriverRequest {
request: BoxAny,
tx: oneshot::Sender<Result<BoxAny>>,
}
impl DriverRequest {
pub fn request<REQ: Any>(&self) -> Result<&REQ> {
self.request.downcast_ref::<REQ>().ok_or_else(|| {
DriverError::TypeMismatch {
expected: std::any::type_name::<REQ>(),
}
.into()
})
}
pub fn respond<RES: Any + Send>(self, res: Result<RES>) -> Result<()> {
let boxed: Result<BoxAny> = res.map(|r| Box::new(r) as BoxAny);
self.tx
.send(boxed)
.map_err(|_| DriverError::ResponseDropped.into())
}
}
struct DriverInner {
step_tx: mpsc::Sender<StepResult>,
step_rx: Mutex<Option<mpsc::Receiver<StepResult>>>,
}
#[derive(Clone)]
pub struct TypeErasedDriver {
inner: Arc<DriverInner>,
}
impl TypeErasedDriver {
pub fn new() -> Self {
let (step_tx, step_rx) = mpsc::channel(1);
TypeErasedDriver {
inner: Arc::new(DriverInner {
step_tx,
step_rx: Mutex::new(Some(step_rx)),
}),
}
}
fn ensure_started(&self) -> Result<()> {
let guard = self.inner.step_rx.lock();
if guard.is_none() {
Ok(())
} else {
Err(DriverError::NotStarted.into())
}
}
pub async fn fulfill_request<REQ, RES>(&self, _ctx: Context, req: REQ) -> Result<RES>
where
REQ: Any + Send + 'static,
RES: Any + Send + 'static,
{
self.ensure_started()?;
let (tx, rx) = oneshot::channel::<Result<BoxAny>>();
let promise = DriverRequest {
request: Box::new(req),
tx,
};
self.inner
.step_tx
.send(Ok(DriverStep::Request(promise)))
.await
.map_err(|_| DriverError::StreamClosed)?;
let response = timeout(FULFILL_REQUEST_TIMEOUT, rx)
.await
.map_err(|_| DriverError::ResponseTimedOut {
timeout: FULFILL_REQUEST_TIMEOUT,
})?
.map_err(|_| DriverError::ResponseDropped)??;
response.downcast::<RES>().map(|boxed| *boxed).map_err(|_| {
DriverError::TypeMismatch {
expected: std::any::type_name::<RES>(),
}
.into()
})
}
pub async fn info<INFO>(&self, _ctx: Context, info: INFO) -> Result<()>
where
INFO: Any + Send + 'static,
{
self.ensure_started()?;
self.inner
.step_tx
.send(Ok(DriverStep::Info(Box::new(info))))
.await
.map_err(|_| DriverError::StreamClosed.into())
}
pub async fn done<T>(&self, _ctx: Context, payload: T) -> Result<()>
where
T: Any + Send + 'static,
{
self.ensure_started()?;
self.inner
.step_tx
.send(Ok(DriverStep::Done(Box::new(payload))))
.await
.map_err(|_| DriverError::StreamClosed.into())
}
pub async fn fail(&self, _ctx: Context, err: LibsyError) -> Result<()> {
self.ensure_started()?;
self.inner
.step_tx
.send(Err(err))
.await
.map_err(|_| DriverError::StreamClosed.into())
}
pub fn stream(&self) -> impl Stream<Item = Result<DriverStep>> + use<> {
let receiver = self
.inner
.step_rx
.lock()
.take()
.ok_or_else(|| LibsyError::from(DriverError::StreamAlreadyTaken));
match receiver {
Ok(rx) => ReceiverStream::new(rx).left_stream(),
Err(error) => futures::stream::once(async move { Err(error) }).right_stream(),
}
}
}
impl Default for TypeErasedDriver {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use futures::StreamExt;
#[derive(Debug, thiserror::Error)]
#[error("{0}")]
struct TestError(&'static str);
fn test_error(message: &'static str) -> LibsyError {
LibsyError::external("test", TestError(message))
}
#[tokio::test]
async fn fulfill_request_round_trips_typed_values() -> Result<()> {
let driver = TypeErasedDriver::new();
let stream = driver.stream();
let producer = driver.clone();
let handle = tokio::spawn(async move {
producer
.fulfill_request::<u32, String>(Context::default(), 7u32)
.await
});
tokio::pin!(stream);
match stream.next().await.ok_or(DriverError::StreamClosed)?? {
DriverStep::Request(promise) => {
let req = *promise.request::<u32>()?;
assert_eq!(req, 7);
promise.respond::<String>(Ok(format!("got {req}")))?;
}
_ => return Err(test_error("expected a Request step")),
}
assert_eq!(handle.await??, "got 7");
Ok(())
}
#[tokio::test]
async fn info_pushes_a_typed_payload() -> Result<()> {
let driver = TypeErasedDriver::new();
let stream = driver.stream();
driver.info(Context::default(), 42u64).await?;
tokio::pin!(stream);
match stream.next().await.ok_or(DriverError::StreamClosed)?? {
DriverStep::Info(payload) => {
let value = payload
.downcast::<u64>()
.map_err(|_| LibsyError::from(DriverError::TypeMismatch { expected: "u64" }))?;
assert_eq!(*value, 42);
}
_ => return Err(test_error("expected an Info step")),
}
Ok(())
}
#[tokio::test]
async fn done_emits_the_terminal_payload() -> Result<()> {
let driver = TypeErasedDriver::new();
let stream = driver.stream();
driver
.done(Context::default(), "finished".to_string())
.await?;
tokio::pin!(stream);
match stream.next().await.ok_or(DriverError::StreamClosed)?? {
DriverStep::Done(payload) => {
let value = payload.downcast::<String>().map_err(|_| {
LibsyError::from(DriverError::TypeMismatch { expected: "String" })
})?;
assert_eq!(*value, "finished");
}
_ => return Err(test_error("expected a Done step")),
}
Ok(())
}
#[tokio::test]
async fn respond_error_propagates_to_the_producer() -> Result<()> {
let driver = TypeErasedDriver::new();
let stream = driver.stream();
let producer = driver.clone();
let handle = tokio::spawn(async move {
producer
.fulfill_request::<u32, u32>(Context::default(), 1u32)
.await
});
tokio::pin!(stream);
match stream.next().await.ok_or(DriverError::StreamClosed)?? {
DriverStep::Request(promise) => {
promise.respond::<u32>(Err(test_error("upstream failed")))?;
}
_ => return Err(test_error("expected a Request step")),
}
match handle.await? {
Ok(_) => Err(test_error("expected the error to propagate")),
Err(err) => {
assert!(err.to_string().contains("upstream failed"));
Ok(())
}
}
}
#[tokio::test]
async fn response_type_mismatch_errors() -> Result<()> {
let driver = TypeErasedDriver::new();
let stream = driver.stream();
let producer = driver.clone();
let handle = tokio::spawn(async move {
producer
.fulfill_request::<u32, String>(Context::default(), 1u32)
.await
});
tokio::pin!(stream);
match stream.next().await.ok_or(DriverError::StreamClosed)?? {
DriverStep::Request(promise) => {
promise.respond::<u32>(Ok(99u32))?;
}
_ => return Err(test_error("expected a Request step")),
}
match handle.await? {
Ok(_) => Err(test_error("expected a response type mismatch")),
Err(err) => {
assert!(matches!(
err,
LibsyError::Driver(DriverError::TypeMismatch { expected })
if expected == std::any::type_name::<String>()
));
Ok(())
}
}
}
#[tokio::test]
async fn request_downcast_to_wrong_type_errors() -> Result<()> {
let driver = TypeErasedDriver::new();
let stream = driver.stream();
let producer = driver.clone();
let handle = tokio::spawn(async move {
producer
.fulfill_request::<u32, u32>(Context::default(), 5u32)
.await
});
tokio::pin!(stream);
match stream.next().await.ok_or(DriverError::StreamClosed)?? {
DriverStep::Request(promise) => {
assert!(promise.request::<String>().is_err());
promise.respond::<u32>(Ok(5u32))?;
}
_ => return Err(test_error("expected a Request step")),
}
assert_eq!(handle.await??, 5);
Ok(())
}
#[tokio::test]
async fn closed_stream_errors_on_send() -> Result<()> {
let driver = TypeErasedDriver::new();
drop(driver.stream());
assert!(
driver
.fulfill_request::<u32, u32>(Context::default(), 1u32)
.await
.is_err()
);
assert!(driver.info(Context::default(), 1u32).await.is_err());
assert!(driver.done(Context::default(), 1u32).await.is_err());
Ok(())
}
#[tokio::test]
async fn promise_dropped_without_response_errors() -> Result<()> {
let driver = TypeErasedDriver::new();
let stream = driver.stream();
let producer = driver.clone();
let handle = tokio::spawn(async move {
producer
.fulfill_request::<u32, u32>(Context::default(), 1u32)
.await
});
tokio::pin!(stream);
match stream.next().await.ok_or(DriverError::StreamClosed)?? {
DriverStep::Request(_promise) => {}
_ => return Err(test_error("expected a Request step")),
}
assert!(handle.await?.is_err());
Ok(())
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn concurrent_producers_do_not_cross_responses() -> Result<()> {
const N: usize = 8;
let driver = TypeErasedDriver::new();
let stream = driver.stream();
let mut handles = Vec::new();
for i in 0..N {
let producer = driver.clone();
handles.push((
i,
tokio::spawn(async move {
producer
.fulfill_request::<usize, usize>(Context::default(), i)
.await
}),
));
}
tokio::pin!(stream);
let mut served = 0;
while served < N {
match stream.next().await.ok_or(DriverError::StreamClosed)?? {
DriverStep::Request(promise) => {
let req = *promise.request::<usize>()?;
promise.respond::<usize>(Ok(req * 10))?;
served += 1;
}
_ => return Err(test_error("expected a Request step")),
}
}
for (i, handle) in handles {
assert_eq!(handle.await??, i * 10);
}
Ok(())
}
#[tokio::test]
async fn stream_taken_twice_yields_an_error_item() -> Result<()> {
let driver = TypeErasedDriver::new();
let _first = driver.stream();
let second = driver.stream();
tokio::pin!(second);
match second.next().await.ok_or(DriverError::StreamClosed)? {
Err(err) => {
assert!(err.to_string().contains("already taken"));
Ok(())
}
Ok(_) => Err(test_error("expected an error item")),
}
}
#[tokio::test]
async fn fail_surfaces_an_error_item_on_the_stream() -> Result<()> {
let driver = TypeErasedDriver::new();
let stream = driver.stream();
driver
.fail(Context::default(), test_error("kaboom"))
.await?;
tokio::pin!(stream);
match stream.next().await.ok_or(DriverError::StreamClosed)? {
Err(err) => {
assert!(err.to_string().contains("kaboom"));
Ok(())
}
Ok(_) => Err(test_error("expected an error item")),
}
}
#[tokio::test]
async fn dropping_the_stream_terminates_the_producer() -> Result<()> {
let driver = TypeErasedDriver::new();
let mut stream = Box::pin(driver.stream());
let producer = driver.clone();
let handle = tokio::spawn(async move {
let mut sent = 0usize;
while producer.info(Context::default(), sent).await.is_ok() {
sent += 1;
}
sent
});
for _ in 0..2 {
match stream.next().await.ok_or(DriverError::StreamClosed)?? {
DriverStep::Info(_) => {}
_ => return Err(test_error("expected an Info step")),
}
}
drop(stream);
let sent = handle.await?;
assert!(sent >= 2, "producer should have published the paced steps");
Ok(())
}
#[test]
fn ensure_started_requires_the_stream_be_taken_first() -> Result<()> {
let driver = TypeErasedDriver::new();
match driver.ensure_started() {
Ok(()) => {
return Err(test_error(
"expected ensure_started to reject an untaken stream",
));
}
Err(err) => assert!(matches!(err, LibsyError::Driver(DriverError::NotStarted))),
}
let _stream = driver.stream();
assert!(driver.ensure_started().is_ok());
Ok(())
}
}