use crate::observability::{MetricsCollector, MonitoringConfig};
use crate::{backend::McpBackend, handler::GenericServerHandler, middleware::MiddlewareStack};
use async_trait::async_trait;
use pulseengine_auth::{AuthConfig, AuthenticationManager};
use pulseengine_logging::{
AlertConfig, AlertManager, DashboardConfig, DashboardManager, PerformanceProfiler,
PersistenceConfig, ProfilingConfig, SanitizationConfig, StructuredLogger,
};
use pulseengine_mcp_protocol::*;
use pulseengine_mcp_security::{SecurityConfig, SecurityMiddleware};
use pulseengine_mcp_transport::{RequestHandler, Transport, TransportConfig, TransportError};
use std::sync::Arc;
use std::time::Duration;
use thiserror::Error;
use tokio::signal;
use tokio::sync::RwLock;
use tracing::{error, info, warn};
struct TransportHandle {
transport: Arc<RwLock<Box<dyn Transport>>>,
}
#[async_trait]
impl Transport for TransportHandle {
async fn start(&mut self, _handler: RequestHandler) -> std::result::Result<(), TransportError> {
Err(TransportError::NotSupported(
"Cannot start transport through handle".to_string(),
))
}
async fn stop(&mut self) -> std::result::Result<(), TransportError> {
Err(TransportError::NotSupported(
"Cannot stop transport through handle".to_string(),
))
}
async fn health_check(&self) -> std::result::Result<(), TransportError> {
let transport = self.transport.read().await;
transport.health_check().await
}
fn supports_bidirectional(&self) -> bool {
self.transport
.try_read()
.map(|t| t.supports_bidirectional())
.unwrap_or(false)
}
async fn send_notification(
&self,
session_id: Option<&str>,
method: &str,
params: serde_json::Value,
) -> std::result::Result<(), TransportError> {
let transport = self.transport.read().await;
transport
.send_notification(session_id, method, params)
.await
}
async fn send_request(
&self,
session_id: Option<&str>,
method: &str,
params: serde_json::Value,
timeout: Duration,
) -> std::result::Result<serde_json::Value, TransportError> {
let transport = self.transport.read().await;
transport
.send_request(session_id, method, params, timeout)
.await
}
fn register_pending_request(
&self,
request_id: &str,
) -> Option<tokio::sync::oneshot::Receiver<serde_json::Value>> {
self.transport
.try_read()
.ok()
.and_then(|t| t.register_pending_request(request_id))
}
}
#[derive(Debug, Error)]
pub enum ServerError {
#[error("Server configuration error: {0}")]
Configuration(String),
#[error("Transport error: {0}")]
Transport(String),
#[error("Authentication error: {0}")]
Authentication(String),
#[error("Backend error: {0}")]
Backend(String),
#[error("Server already running")]
AlreadyRunning,
#[error("Server not running")]
NotRunning,
#[error("Shutdown timeout")]
ShutdownTimeout,
}
#[derive(Debug, Clone)]
pub struct ServerConfig {
pub server_info: ServerInfo,
pub auth_config: AuthConfig,
pub transport_config: TransportConfig,
pub security_config: SecurityConfig,
pub monitoring_config: MonitoringConfig,
pub sanitization_config: SanitizationConfig,
pub persistence_config: Option<PersistenceConfig>,
pub alert_config: AlertConfig,
pub dashboard_config: DashboardConfig,
pub profiling_config: ProfilingConfig,
pub graceful_shutdown: bool,
pub shutdown_timeout_secs: u64,
}
impl Default for ServerConfig {
fn default() -> Self {
Self {
server_info: ServerInfo {
protocol_version: ProtocolVersion::default(),
capabilities: ServerCapabilities::default(),
server_info: Implementation::new("MCP Server", "1.0.0"),
instructions: None,
},
auth_config: pulseengine_auth::default_config(),
transport_config: pulseengine_mcp_transport::TransportConfig::default(),
security_config: pulseengine_mcp_security::default_config(),
monitoring_config: crate::observability::default_config(),
sanitization_config: SanitizationConfig::default(),
persistence_config: None,
alert_config: AlertConfig::default(),
dashboard_config: DashboardConfig::default(),
profiling_config: ProfilingConfig::default(),
graceful_shutdown: true,
shutdown_timeout_secs: 30,
}
}
}
pub struct McpServer<B: McpBackend> {
backend: Arc<B>,
handler: GenericServerHandler<B>,
auth_manager: Arc<AuthenticationManager>,
transport: Arc<tokio::sync::RwLock<Box<dyn Transport>>>,
#[allow(dead_code)]
middleware_stack: MiddlewareStack,
monitoring_metrics: Arc<MetricsCollector>,
#[allow(dead_code)]
logging_metrics: Arc<pulseengine_logging::MetricsCollector>,
#[allow(dead_code)]
logger: StructuredLogger,
alert_manager: Arc<AlertManager>,
dashboard_manager: Arc<DashboardManager>,
profiler: Option<Arc<PerformanceProfiler>>,
config: ServerConfig,
running: Arc<tokio::sync::RwLock<bool>>,
}
impl<B: McpBackend + 'static> McpServer<B> {
pub async fn new(backend: B, config: ServerConfig) -> std::result::Result<Self, ServerError> {
let logger = StructuredLogger::new();
info!("Initializing MCP server with backend");
let auth_manager = if config.auth_config.enabled {
Arc::new(
AuthenticationManager::new(config.auth_config.clone())
.await
.map_err(|e| ServerError::Authentication(e.to_string()))?,
)
} else {
Arc::new(AuthenticationManager::new_disabled())
};
let transport = Arc::new(tokio::sync::RwLock::new(
pulseengine_mcp_transport::create_transport(config.transport_config.clone())
.map_err(|e| ServerError::Transport(e.to_string()))?,
));
let security_middleware = SecurityMiddleware::new(config.security_config.clone());
let monitoring_metrics = Arc::new(MetricsCollector::new(config.monitoring_config.clone()));
let logging_metrics = Arc::new(pulseengine_logging::MetricsCollector::new());
if let Some(persistence_config) = config.persistence_config.clone() {
logging_metrics
.enable_persistence(persistence_config.clone())
.await
.map_err(|e| {
ServerError::Configuration(format!(
"Failed to initialize metrics persistence: {e}"
))
})?;
}
let middleware_stack = MiddlewareStack::new()
.with_security(security_middleware)
.with_monitoring(monitoring_metrics.clone())
.with_auth(auth_manager.clone());
let backend = Arc::new(backend);
let alert_manager = Arc::new(AlertManager::new(config.alert_config.clone()));
let dashboard_manager = Arc::new(DashboardManager::new(config.dashboard_config.clone()));
let profiler = if config.profiling_config.enabled {
Some(Arc::new(PerformanceProfiler::new(
config.profiling_config.clone(),
)))
} else {
None
};
let handler = GenericServerHandler::new(
backend.clone(),
auth_manager.clone(),
middleware_stack.clone(),
);
Ok(Self {
backend,
handler,
auth_manager,
transport,
middleware_stack,
monitoring_metrics,
logging_metrics,
logger,
alert_manager,
dashboard_manager,
profiler,
config,
running: Arc::new(tokio::sync::RwLock::new(false)),
})
}
#[tracing::instrument(skip(self))]
pub async fn start(&mut self) -> std::result::Result<(), ServerError> {
{
let mut running = self.running.write().await;
if *running {
return Err(ServerError::AlreadyRunning);
}
*running = true;
}
info!("Starting MCP server");
self.backend
.on_startup()
.await
.map_err(|e| ServerError::Backend(e.to_string()))?;
self.auth_manager
.start_background_tasks()
.await
.map_err(|e| ServerError::Authentication(e.to_string()))?;
self.alert_manager.start().await;
self.start_dashboard_metrics_update().await;
if let Some(profiler) = &self.profiler {
profiler
.start_session(
format!("server_session_{}", chrono::Utc::now().timestamp()),
pulseengine_logging::ProfilingSessionType::Continuous,
)
.await
.map_err(|e| {
ServerError::Configuration(format!("Failed to start profiling session: {e}"))
})?;
}
let transport_handle: Arc<dyn Transport> = Arc::new(TransportHandle {
transport: self.transport.clone(),
});
self.handler.set_transport(transport_handle);
let handler = self.handler.clone();
{
let mut transport_guard = self.transport.write().await;
transport_guard
.start(Box::new(move |request| {
let handler = handler.clone();
Box::pin(async move {
match handler.handle_request(request).await {
Ok(response) => response,
Err(error) => Response {
jsonrpc: "2.0".to_string(),
id: None,
result: None,
error: Some(error.into()),
},
}
})
}))
.await
.map_err(|e| ServerError::Transport(e.to_string()))?;
}
info!("MCP server started successfully");
if self.config.graceful_shutdown {
let running = self.running.clone();
tokio::spawn(async move {
signal::ctrl_c().await.expect("Failed to listen for Ctrl+C");
warn!("Shutdown signal received");
let mut running = running.write().await;
*running = false;
});
}
Ok(())
}
pub async fn stop(&mut self) -> std::result::Result<(), ServerError> {
{
let mut running = self.running.write().await;
if !*running {
return Err(ServerError::NotRunning);
}
*running = false;
}
info!("Stopping MCP server");
{
let mut transport_guard = self.transport.write().await;
transport_guard
.stop()
.await
.map_err(|e| ServerError::Transport(e.to_string()))?;
}
self.monitoring_metrics.stop_collection().await;
self.auth_manager
.stop_background_tasks()
.await
.map_err(|e| ServerError::Authentication(e.to_string()))?;
if let Some(profiler) = &self.profiler {
profiler.stop_session().await.map_err(|e| {
ServerError::Configuration(format!("Failed to stop profiling session: {e}"))
})?;
}
self.backend
.on_shutdown()
.await
.map_err(|e| ServerError::Backend(e.to_string()))?;
info!("MCP server stopped");
Ok(())
}
pub async fn run(&mut self) -> std::result::Result<(), ServerError> {
self.start().await?;
loop {
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let running = self.running.read().await;
if !*running {
break;
}
}
self.stop().await?;
Ok(())
}
pub async fn health_check(&self) -> std::result::Result<HealthStatus, ServerError> {
let backend_healthy = self.backend.health_check().await.is_ok();
let transport_healthy = {
let transport_guard = self.transport.read().await;
transport_guard.health_check().await.is_ok()
};
let auth_healthy = self.auth_manager.health_check().await.is_ok();
let overall_healthy = backend_healthy && transport_healthy && auth_healthy;
Ok(HealthStatus {
status: if overall_healthy {
"healthy".to_string()
} else {
"unhealthy".to_string()
},
components: vec![
("backend".to_string(), backend_healthy),
("transport".to_string(), transport_healthy),
("auth".to_string(), auth_healthy),
]
.into_iter()
.collect(),
uptime_seconds: self.monitoring_metrics.get_uptime_seconds(),
})
}
pub async fn get_metrics(&self) -> ServerMetrics {
self.monitoring_metrics.get_current_metrics().await
}
pub fn get_server_info(&self) -> &ServerInfo {
&self.config.server_info
}
pub async fn is_running(&self) -> bool {
*self.running.read().await
}
pub fn get_alert_manager(&self) -> Arc<AlertManager> {
self.alert_manager.clone()
}
pub fn get_dashboard_manager(&self) -> Arc<DashboardManager> {
self.dashboard_manager.clone()
}
pub fn get_profiler(&self) -> Option<Arc<PerformanceProfiler>> {
self.profiler.clone()
}
async fn start_dashboard_metrics_update(&self) {
if !self.config.dashboard_config.enabled {
return;
}
let logging_metrics = self.logging_metrics.clone();
let dashboard_manager = self.dashboard_manager.clone();
let refresh_interval = self.config.dashboard_config.refresh_interval_secs;
tokio::spawn(async move {
let mut interval =
tokio::time::interval(std::time::Duration::from_secs(refresh_interval));
loop {
interval.tick().await;
let metrics_snapshot = logging_metrics.get_metrics_snapshot().await;
dashboard_manager.update_metrics(metrics_snapshot).await;
}
});
}
}
#[derive(Debug, serde::Serialize, serde::Deserialize)]
pub struct HealthStatus {
pub status: String,
pub components: std::collections::HashMap<String, bool>,
pub uptime_seconds: u64,
}
pub use crate::observability::ServerMetrics;