use async_trait::async_trait;
use serde_json::json;
use std::sync::Arc;
use crate::middleware::{
DispatcherResult, McpMiddleware, MiddlewareError, MiddlewareStack, RequestContext,
SessionInjection, StorageBackedSessionView,
};
use turul_mcp_protocol::ServerCapabilities;
use turul_mcp_session_storage::{BoxedSessionStorage, InMemorySessionStorage, SessionView};
struct StateWriterMiddleware {
key: String,
value: String,
}
#[async_trait]
impl McpMiddleware for StateWriterMiddleware {
async fn before_dispatch(
&self,
_ctx: &mut RequestContext<'_>,
_session: Option<&dyn SessionView>,
injection: &mut SessionInjection,
) -> Result<(), MiddlewareError> {
injection.set_state(&self.key, json!(self.value));
Ok(())
}
async fn after_dispatch(
&self,
_ctx: &RequestContext<'_>,
_result: &mut DispatcherResult,
) -> Result<(), MiddlewareError> {
Ok(())
}
}
struct BlockingMiddleware;
#[async_trait]
impl McpMiddleware for BlockingMiddleware {
async fn before_dispatch(
&self,
_ctx: &mut RequestContext<'_>,
_session: Option<&dyn SessionView>,
_injection: &mut SessionInjection,
) -> Result<(), MiddlewareError> {
Err(MiddlewareError::Unauthenticated("Missing auth".to_string()))
}
async fn after_dispatch(
&self,
_ctx: &RequestContext<'_>,
_result: &mut DispatcherResult,
) -> Result<(), MiddlewareError> {
Ok(())
}
}
#[tokio::test]
async fn test_handler_pattern_middleware_writes_to_session() {
let storage: Arc<BoxedSessionStorage> = Arc::new(InMemorySessionStorage::new());
let session_info = storage
.create_session(ServerCapabilities::default())
.await
.unwrap();
let mut middleware_stack = MiddlewareStack::new();
middleware_stack.push(Arc::new(StateWriterMiddleware {
key: "middleware_key".to_string(),
value: "middleware_value".to_string(),
}));
let session_view =
StorageBackedSessionView::new(session_info.session_id.clone(), Arc::clone(&storage));
let mut ctx = RequestContext::new("test/method", None);
let injection = middleware_stack
.execute_before(&mut ctx, Some(&session_view))
.await
.unwrap();
for (key, value) in injection.state() {
session_view.set_state(key, value.clone()).await.unwrap();
}
let session = storage
.get_session(&session_info.session_id)
.await
.unwrap()
.unwrap();
assert_eq!(
session.state.get("middleware_key"),
Some(&json!("middleware_value")),
"Middleware should write to session state via injection"
);
}
#[tokio::test]
async fn test_handler_pattern_middleware_error_short_circuits() {
let storage: Arc<BoxedSessionStorage> = Arc::new(InMemorySessionStorage::new());
let session_info = storage
.create_session(ServerCapabilities::default())
.await
.unwrap();
let mut middleware_stack = MiddlewareStack::new();
middleware_stack.push(Arc::new(BlockingMiddleware));
let session_view =
StorageBackedSessionView::new(session_info.session_id.clone(), Arc::clone(&storage));
let mut ctx = RequestContext::new("test/method", None);
let result = middleware_stack
.execute_before(&mut ctx, Some(&session_view))
.await;
assert!(result.is_err(), "Blocking middleware should return error");
match result.unwrap_err() {
MiddlewareError::Unauthenticated(msg) => {
assert_eq!(msg, "Missing auth");
}
other => panic!("Expected Unauthenticated error, got {:?}", other),
}
}
#[tokio::test]
async fn test_both_handlers_use_run_middleware_and_dispatch() {
let storage: Arc<BoxedSessionStorage> = Arc::new(InMemorySessionStorage::new());
let session_info = storage
.create_session(ServerCapabilities::default())
.await
.unwrap();
let middleware_stack = MiddlewareStack::new();
let session_view =
StorageBackedSessionView::new(session_info.session_id.clone(), Arc::clone(&storage));
let mut ctx = RequestContext::new("test/method", None);
let injection = middleware_stack
.execute_before(&mut ctx, Some(&session_view))
.await
.unwrap();
assert!(injection.state().is_empty());
assert!(injection.metadata().is_empty());
}