use crate::context::RequestContext;
use crate::middleware::{Middleware, MiddlewareStack};
use crate::observability::{MetricsCollector, MonitoringConfig};
use async_trait::async_trait;
use pulseengine_auth::{AuthConfig, AuthenticationManager, config::StorageConfig};
use pulseengine_mcp_protocol::*;
use pulseengine_mcp_security::{SecurityConfig, SecurityMiddleware};
use std::sync::Arc;
use uuid::Uuid;
#[test]
fn test_middleware_stack_new() {
let stack = MiddlewareStack::new();
let context = RequestContext::new();
let request = Request {
jsonrpc: "2.0".to_string(),
id: Some(pulseengine_mcp_protocol::NumberOrString::String(Arc::from(
"test",
))),
method: "test".to_string(),
params: serde_json::Value::Null,
};
tokio_test::block_on(async {
let result = stack.process_request(request.clone(), &context).await;
assert!(result.is_ok());
});
}
#[test]
fn test_middleware_stack_default() {
let stack = MiddlewareStack::default();
let context = RequestContext::new();
let request = Request {
jsonrpc: "2.0".to_string(),
id: Some(pulseengine_mcp_protocol::NumberOrString::String(Arc::from(
"test",
))),
method: "test".to_string(),
params: serde_json::Value::Null,
};
tokio_test::block_on(async {
let result = stack.process_request(request, &context).await;
assert!(result.is_ok());
});
}
#[test]
fn test_middleware_stack_builder_pattern() {
let security_config = SecurityConfig::default();
let security_middleware = SecurityMiddleware::new(security_config);
let monitoring_config = MonitoringConfig::default();
let monitoring = Arc::new(MetricsCollector::new(monitoring_config));
tokio_test::block_on(async {
let auth_config = AuthConfig {
storage: StorageConfig::Memory,
enabled: false,
cache_size: 100,
session_timeout_secs: 3600,
max_failed_attempts: 5,
rate_limit_window_secs: 900,
};
let auth_manager = Arc::new(AuthenticationManager::new(auth_config).await.unwrap());
let stack = MiddlewareStack::new()
.with_security(security_middleware)
.with_monitoring(monitoring)
.with_auth(auth_manager);
let context = RequestContext::new();
let request = Request {
jsonrpc: "2.0".to_string(),
id: Some(pulseengine_mcp_protocol::NumberOrString::String(Arc::from(
"test",
))),
method: "test".to_string(),
params: serde_json::Value::Null,
};
let result = stack.process_request(request, &context).await;
assert!(result.is_ok());
});
}
#[tokio::test]
async fn test_middleware_stack_process_request() {
let context = RequestContext::new()
.with_user("test_user")
.with_role("admin");
let request = Request {
jsonrpc: "2.0".to_string(),
id: Some(pulseengine_mcp_protocol::NumberOrString::String(Arc::from(
"test_request",
))),
method: "tools/list".to_string(),
params: serde_json::json!({"cursor": null}),
};
let security_config = SecurityConfig::default();
let security_middleware = SecurityMiddleware::new(security_config);
let stack = MiddlewareStack::new().with_security(security_middleware);
let result = stack.process_request(request.clone(), &context).await;
assert!(result.is_ok());
let processed_request = result.unwrap();
assert_eq!(processed_request.method, "tools/list");
assert_eq!(processed_request.jsonrpc, "2.0");
}
#[tokio::test]
async fn test_middleware_stack_process_response() {
let context = RequestContext::new()
.with_user("test_user")
.with_role("admin");
let response = Response {
jsonrpc: "2.0".to_string(),
id: Some(pulseengine_mcp_protocol::NumberOrString::String(Arc::from(
"test_response",
))),
result: Some(serde_json::json!({"tools": []})),
error: None,
};
let monitoring_config = MonitoringConfig::default();
let monitoring = Arc::new(MetricsCollector::new(monitoring_config));
let stack = MiddlewareStack::new().with_monitoring(monitoring);
let result = stack.process_response(response.clone(), &context).await;
assert!(result.is_ok());
let processed_response = result.unwrap();
assert_eq!(processed_response.jsonrpc, "2.0");
assert!(processed_response.result.is_some());
}
#[tokio::test]
async fn test_middleware_stack_with_auth() {
let auth_config = AuthConfig {
storage: StorageConfig::Memory,
enabled: false,
cache_size: 100,
session_timeout_secs: 3600,
max_failed_attempts: 5,
rate_limit_window_secs: 900,
};
let auth_manager = Arc::new(AuthenticationManager::new(auth_config).await.unwrap());
let stack = MiddlewareStack::new().with_auth(auth_manager);
let context = RequestContext::new().with_user("authenticated_user");
let request = Request {
jsonrpc: "2.0".to_string(),
id: Some(pulseengine_mcp_protocol::NumberOrString::String(Arc::from(
"auth_test",
))),
method: "tools/call".to_string(),
params: serde_json::json!({
"name": "test_tool",
"arguments": {}
}),
};
let result = stack.process_request(request, &context).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_middleware_stack_full_pipeline() {
let security_config = SecurityConfig::default();
let security_middleware = SecurityMiddleware::new(security_config);
let auth_config = AuthConfig {
storage: StorageConfig::Memory,
enabled: false,
cache_size: 100,
session_timeout_secs: 3600,
max_failed_attempts: 5,
rate_limit_window_secs: 900,
};
let auth_manager = Arc::new(AuthenticationManager::new(auth_config).await.unwrap());
let monitoring_config = MonitoringConfig::default();
let monitoring = Arc::new(MetricsCollector::new(monitoring_config));
let stack = MiddlewareStack::new()
.with_security(security_middleware)
.with_auth(auth_manager)
.with_monitoring(monitoring);
let context = RequestContext::new()
.with_user("full_pipeline_user")
.with_role("admin")
.with_metadata("request_source", "test");
let request = Request {
jsonrpc: "2.0".to_string(),
id: Some(pulseengine_mcp_protocol::NumberOrString::String(Arc::from(
"full_pipeline_test",
))),
method: "resources/list".to_string(),
params: serde_json::json!({"cursor": null}),
};
let processed_request = stack.process_request(request, &context).await.unwrap();
assert_eq!(processed_request.method, "resources/list");
let response = Response {
jsonrpc: "2.0".to_string(),
id: Some(pulseengine_mcp_protocol::NumberOrString::String(Arc::from(
"full_pipeline_test",
))),
result: Some(serde_json::json!({"resources": []})),
error: None,
};
let processed_response = stack.process_response(response, &context).await.unwrap();
assert!(processed_response.result.is_some());
}
#[tokio::test]
async fn test_middleware_stack_error_handling() {
let stack = MiddlewareStack::new();
let context = RequestContext::new();
let malformed_request = Request {
jsonrpc: "2.0".to_string(),
id: Some(pulseengine_mcp_protocol::NumberOrString::String(Arc::from(
"error_test",
))),
method: "".to_string(), params: serde_json::Value::Null,
};
let result = stack.process_request(malformed_request, &context).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_middleware_stack_request_context_usage() {
let monitoring_config = MonitoringConfig::default();
let monitoring = Arc::new(MetricsCollector::new(monitoring_config));
let stack = MiddlewareStack::new().with_monitoring(monitoring);
let request_id = Uuid::new_v4();
let context = RequestContext::with_id(request_id)
.with_user("context_test_user")
.with_metadata("test_key", "test_value");
let request = Request {
jsonrpc: "2.0".to_string(),
id: Some(pulseengine_mcp_protocol::NumberOrString::String(Arc::from(
"context_test",
))),
method: "ping".to_string(),
params: serde_json::Value::Null,
};
let result = stack.process_request(request, &context).await;
assert!(result.is_ok());
assert_eq!(context.request_id, request_id);
assert_eq!(
context.authenticated_user.as_ref().unwrap(),
"context_test_user"
);
assert_eq!(context.get_metadata("test_key").unwrap(), "test_value");
}
struct MockMiddleware {
should_fail: bool,
}
#[async_trait]
impl Middleware for MockMiddleware {
async fn process_request(
&self,
request: Request,
_context: &RequestContext,
) -> std::result::Result<Request, Error> {
if self.should_fail {
Err(Error::internal_error("Mock middleware failed"))
} else {
Ok(request)
}
}
async fn process_response(
&self,
response: Response,
_context: &RequestContext,
) -> std::result::Result<Response, Error> {
if self.should_fail {
Err(Error::internal_error("Mock middleware failed"))
} else {
Ok(response)
}
}
}
#[tokio::test]
async fn test_custom_middleware_implementation() {
let mock_middleware = MockMiddleware { should_fail: false };
let context = RequestContext::new();
let request = Request {
jsonrpc: "2.0".to_string(),
id: Some(pulseengine_mcp_protocol::NumberOrString::String(Arc::from(
"mock_test",
))),
method: "test".to_string(),
params: serde_json::Value::Null,
};
let result = mock_middleware.process_request(request, &context).await;
assert!(result.is_ok());
let response = Response {
jsonrpc: "2.0".to_string(),
id: Some(pulseengine_mcp_protocol::NumberOrString::String(Arc::from(
"mock_test",
))),
result: Some(serde_json::Value::Null),
error: None,
};
let result = mock_middleware.process_response(response, &context).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_custom_middleware_failure() {
let mock_middleware = MockMiddleware { should_fail: true };
let context = RequestContext::new();
let request = Request {
jsonrpc: "2.0".to_string(),
id: Some(pulseengine_mcp_protocol::NumberOrString::String(Arc::from(
"fail_test",
))),
method: "test".to_string(),
params: serde_json::Value::Null,
};
let result = mock_middleware.process_request(request, &context).await;
assert!(result.is_err());
let response = Response {
jsonrpc: "2.0".to_string(),
id: Some(pulseengine_mcp_protocol::NumberOrString::String(Arc::from(
"fail_test",
))),
result: Some(serde_json::Value::Null),
error: None,
};
let result = mock_middleware.process_response(response, &context).await;
assert!(result.is_err());
}
#[test]
fn test_middleware_types_send_sync() {
fn assert_send<T: Send>() {}
fn assert_sync<T: Sync>() {}
assert_send::<MiddlewareStack>();
assert_sync::<MiddlewareStack>();
}
#[test]
fn test_middleware_stack_clone() {
let security_config = SecurityConfig::default();
let security_middleware = SecurityMiddleware::new(security_config);
let stack = MiddlewareStack::new().with_security(security_middleware);
let cloned_stack = stack.clone();
let context = RequestContext::new();
let request = Request {
jsonrpc: "2.0".to_string(),
id: Some(pulseengine_mcp_protocol::NumberOrString::String(Arc::from(
"clone_test",
))),
method: "test".to_string(),
params: serde_json::Value::Null,
};
tokio_test::block_on(async {
let result1 = stack.process_request(request.clone(), &context).await;
let result2 = cloned_stack.process_request(request, &context).await;
assert!(result1.is_ok());
assert!(result2.is_ok());
});
}