use std::{fmt, sync::Arc};
use crate::{BoxFuture, Command, CommandHandler, Query, QueryHandler, SoapResult};
pub trait CommandDispatcher<C>: Send + Sync
where
C: Command,
{
fn dispatch(&self, command: C) -> BoxFuture<'_, SoapResult<C::Output>>;
}
impl<C, H> CommandDispatcher<C> for H
where
C: Command,
H: CommandHandler<C> + ?Sized,
{
fn dispatch(&self, command: C) -> BoxFuture<'_, SoapResult<C::Output>> {
self.command(command)
}
}
pub trait QueryDispatcher<Q>: Send + Sync
where
Q: Query,
{
fn dispatch(&self, query: Q) -> BoxFuture<'_, SoapResult<Q::Output>>;
}
impl<Q, H> QueryDispatcher<Q> for H
where
Q: Query,
H: QueryHandler<Q> + ?Sized,
{
fn dispatch(&self, query: Q) -> BoxFuture<'_, SoapResult<Q::Output>> {
self.query(query)
}
}
pub struct CommandNext<'a, C>
where
C: Command + 'static,
{
middleware: &'a [Arc<dyn CommandMiddleware<C>>],
handler: &'a dyn CommandHandler<C>,
}
impl<C> fmt::Debug for CommandNext<'_, C>
where
C: Command + 'static,
{
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("CommandNext")
.field("remaining_middleware", &self.middleware.len())
.finish_non_exhaustive()
}
}
impl<'a, C> CommandNext<'a, C>
where
C: Command + 'static,
{
pub fn run(self, command: C) -> BoxFuture<'a, SoapResult<C::Output>> {
if let Some((middleware, remaining)) = self.middleware.split_first() {
middleware.handle(
command,
Self {
middleware: remaining,
handler: self.handler,
},
)
} else {
self.handler.command(command)
}
}
}
pub trait CommandMiddleware<C>: Send + Sync
where
C: Command + 'static,
{
fn handle<'a>(
&'a self,
command: C,
next: CommandNext<'a, C>,
) -> BoxFuture<'a, SoapResult<C::Output>>;
}
pub struct CommandPipeline<C>
where
C: Command + 'static,
{
handler: Arc<dyn CommandHandler<C>>,
middleware: Vec<Arc<dyn CommandMiddleware<C>>>,
}
impl<C> fmt::Debug for CommandPipeline<C>
where
C: Command + 'static,
{
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("CommandPipeline")
.field("middleware", &self.middleware.len())
.finish_non_exhaustive()
}
}
impl<C> CommandPipeline<C>
where
C: Command + 'static,
{
pub fn new(handler: Arc<dyn CommandHandler<C>>) -> Self {
Self {
handler,
middleware: Vec::new(),
}
}
#[must_use]
pub fn with_middleware(mut self, middleware: Arc<dyn CommandMiddleware<C>>) -> Self {
self.middleware.push(middleware);
self
}
pub fn push_middleware(&mut self, middleware: Arc<dyn CommandMiddleware<C>>) {
self.middleware.push(middleware);
}
pub fn middleware_count(&self) -> usize {
self.middleware.len()
}
}
impl<C> CommandHandler<C> for CommandPipeline<C>
where
C: Command + 'static,
{
fn command(&self, command: C) -> BoxFuture<'_, SoapResult<C::Output>> {
CommandNext {
middleware: &self.middleware,
handler: self.handler.as_ref(),
}
.run(command)
}
}
pub struct QueryNext<'a, Q>
where
Q: Query + 'static,
{
middleware: &'a [Arc<dyn QueryMiddleware<Q>>],
handler: &'a dyn QueryHandler<Q>,
}
impl<Q> fmt::Debug for QueryNext<'_, Q>
where
Q: Query + 'static,
{
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("QueryNext")
.field("remaining_middleware", &self.middleware.len())
.finish_non_exhaustive()
}
}
impl<'a, Q> QueryNext<'a, Q>
where
Q: Query + 'static,
{
pub fn run(self, query: Q) -> BoxFuture<'a, SoapResult<Q::Output>> {
if let Some((middleware, remaining)) = self.middleware.split_first() {
middleware.handle(
query,
Self {
middleware: remaining,
handler: self.handler,
},
)
} else {
self.handler.query(query)
}
}
}
pub trait QueryMiddleware<Q>: Send + Sync
where
Q: Query + 'static,
{
fn handle<'a>(
&'a self,
query: Q,
next: QueryNext<'a, Q>,
) -> BoxFuture<'a, SoapResult<Q::Output>>;
}
pub struct QueryPipeline<Q>
where
Q: Query + 'static,
{
handler: Arc<dyn QueryHandler<Q>>,
middleware: Vec<Arc<dyn QueryMiddleware<Q>>>,
}
impl<Q> fmt::Debug for QueryPipeline<Q>
where
Q: Query + 'static,
{
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("QueryPipeline")
.field("middleware", &self.middleware.len())
.finish_non_exhaustive()
}
}
impl<Q> QueryPipeline<Q>
where
Q: Query + 'static,
{
pub fn new(handler: Arc<dyn QueryHandler<Q>>) -> Self {
Self {
handler,
middleware: Vec::new(),
}
}
#[must_use]
pub fn with_middleware(mut self, middleware: Arc<dyn QueryMiddleware<Q>>) -> Self {
self.middleware.push(middleware);
self
}
pub fn push_middleware(&mut self, middleware: Arc<dyn QueryMiddleware<Q>>) {
self.middleware.push(middleware);
}
pub fn middleware_count(&self) -> usize {
self.middleware.len()
}
}
impl<Q> QueryHandler<Q> for QueryPipeline<Q>
where
Q: Query + 'static,
{
fn query(&self, query: Q) -> BoxFuture<'_, SoapResult<Q::Output>> {
QueryNext {
middleware: &self.middleware,
handler: self.handler.as_ref(),
}
.run(query)
}
}
#[cfg(test)]
mod tests {
use std::{
future::Future,
sync::{Arc, Mutex},
task::{Context, Poll, Waker},
};
use crate::{
BoxFuture, Command, CommandDispatcher, CommandHandler, CommandMiddleware, CommandNext,
CommandPipeline, Query, QueryDispatcher, QueryHandler, QueryMiddleware, QueryNext,
QueryPipeline, SoapError, SoapErrorKind, SoapResult,
};
fn block_on<F>(future: F) -> F::Output
where
F: Future,
{
let mut context = Context::from_waker(Waker::noop());
let mut future = Box::pin(future);
loop {
match future.as_mut().poll(&mut context) {
Poll::Ready(output) => return output,
Poll::Pending => std::thread::yield_now(),
}
}
}
#[derive(Debug)]
struct Add(i32);
impl Command for Add {
type Output = i32;
}
struct AddHandler {
trace: Arc<Mutex<Vec<&'static str>>>,
}
impl CommandHandler<Add> for AddHandler {
fn command(&self, command: Add) -> BoxFuture<'_, SoapResult<i32>> {
Box::pin(async move {
self.trace
.lock()
.map_err(|_| SoapError::infrastructure("command test trace lock poisoned"))?
.push("handler");
Ok(command.0)
})
}
}
struct AroundCommand {
before: &'static str,
after: &'static str,
add: i32,
trace: Arc<Mutex<Vec<&'static str>>>,
}
impl CommandMiddleware<Add> for AroundCommand {
fn handle<'a>(
&'a self,
command: Add,
next: CommandNext<'a, Add>,
) -> BoxFuture<'a, SoapResult<i32>> {
Box::pin(async move {
self.trace
.lock()
.map_err(|_| SoapError::infrastructure("command test trace lock poisoned"))?
.push(self.before);
let output = next.run(command).await?;
self.trace
.lock()
.map_err(|_| SoapError::infrastructure("command test trace lock poisoned"))?
.push(self.after);
Ok(output + self.add)
})
}
}
#[test]
fn command_pipeline_preserves_types_and_around_order() {
let trace = Arc::new(Mutex::new(Vec::new()));
let handler: Arc<dyn CommandHandler<Add>> = Arc::new(AddHandler {
trace: Arc::clone(&trace),
});
let pipeline = CommandPipeline::new(handler)
.with_middleware(Arc::new(AroundCommand {
before: "outer-before",
after: "outer-after",
add: 1,
trace: Arc::clone(&trace),
}))
.with_middleware(Arc::new(AroundCommand {
before: "inner-before",
after: "inner-after",
add: 10,
trace: Arc::clone(&trace),
}));
assert_eq!(pipeline.middleware_count(), 2);
assert_eq!(block_on(pipeline.dispatch(Add(5))).ok(), Some(16));
assert_eq!(
trace.lock().ok().map(|items| items.clone()),
Some(vec![
"outer-before",
"inner-before",
"handler",
"inner-after",
"outer-after"
])
);
}
struct RejectCommand;
impl CommandMiddleware<Add> for RejectCommand {
fn handle<'a>(
&'a self,
_command: Add,
_next: CommandNext<'a, Add>,
) -> BoxFuture<'a, SoapResult<i32>> {
Box::pin(async { Err(SoapError::validation("command rejected by middleware")) })
}
}
#[test]
fn command_middleware_can_short_circuit_the_handler() {
let trace = Arc::new(Mutex::new(Vec::new()));
let handler: Arc<dyn CommandHandler<Add>> = Arc::new(AddHandler {
trace: Arc::clone(&trace),
});
let pipeline = CommandPipeline::new(handler).with_middleware(Arc::new(RejectCommand));
let result = block_on(pipeline.dispatch(Add(5)));
assert_eq!(
result.as_ref().map_err(SoapError::kind),
Err(SoapErrorKind::Validation)
);
assert_eq!(trace.lock().ok().map(|items| items.is_empty()), Some(true));
}
struct Double(i32);
impl Query for Double {
type Output = i32;
}
struct DoubleHandler;
impl QueryHandler<Double> for DoubleHandler {
fn query(&self, query: Double) -> BoxFuture<'_, SoapResult<i32>> {
Box::pin(async move { Ok(query.0 * 2) })
}
}
struct AddToQuery(i32);
impl QueryMiddleware<Double> for AddToQuery {
fn handle<'a>(
&'a self,
query: Double,
next: QueryNext<'a, Double>,
) -> BoxFuture<'a, SoapResult<i32>> {
Box::pin(async move { Ok(next.run(query).await? + self.0) })
}
}
#[test]
fn query_pipeline_and_direct_handlers_share_typed_dispatch() {
assert_eq!(block_on(DoubleHandler.dispatch(Double(4))).ok(), Some(8));
let handler: Arc<dyn QueryHandler<Double>> = Arc::new(DoubleHandler);
let mut pipeline = QueryPipeline::new(handler);
pipeline.push_middleware(Arc::new(AddToQuery(3)));
assert_eq!(pipeline.middleware_count(), 1);
assert_eq!(block_on(pipeline.dispatch(Double(4))).ok(), Some(11));
}
}