use serde_json::Value;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::RwLock;
use tracing::{debug, info, warn};
use crate::config::ClientConfig;
use crate::error::{McpClientResult, SessionError};
use crate::version::McpVersion;
use turul_mcp_protocol_2025_11_25::{
ClientCapabilities, Implementation, InitializeRequest, ServerCapabilities,
};
#[derive(Debug, Clone, PartialEq)]
pub enum SessionState {
Uninitialized,
Initializing,
Active,
Reconnecting,
Terminated,
Error(String),
}
impl std::fmt::Display for SessionState {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
SessionState::Uninitialized => write!(f, "uninitialized"),
SessionState::Initializing => write!(f, "initializing"),
SessionState::Active => write!(f, "active"),
SessionState::Reconnecting => write!(f, "reconnecting"),
SessionState::Terminated => write!(f, "terminated"),
SessionState::Error(err) => write!(f, "error: {}", err),
}
}
}
#[derive(Debug, Clone)]
pub struct SessionInfo {
pub session_id: Option<String>,
pub state: SessionState,
pub client_capabilities: Option<ClientCapabilities>,
pub server_capabilities: Option<ServerCapabilities>,
pub protocol_version: Option<String>,
pub created_at: Instant,
pub last_activity: Instant,
pub connection_attempts: u32,
pub metadata: Value,
}
impl SessionInfo {
pub fn new() -> Self {
let now = Instant::now();
Self {
session_id: None,
state: SessionState::Uninitialized,
client_capabilities: None,
server_capabilities: None,
protocol_version: None,
created_at: now,
last_activity: now,
connection_attempts: 0,
metadata: Value::Null,
}
}
pub fn update_activity(&mut self) {
self.last_activity = Instant::now();
}
pub fn duration(&self) -> Duration {
self.last_activity.duration_since(self.created_at)
}
pub fn idle_time(&self) -> Duration {
Instant::now().duration_since(self.last_activity)
}
pub fn is_active(&self) -> bool {
self.state == SessionState::Active
}
pub fn is_ready(&self) -> bool {
matches!(self.state, SessionState::Active)
}
pub fn needs_initialization(&self) -> bool {
matches!(self.state, SessionState::Uninitialized)
}
}
impl Default for SessionInfo {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug)]
pub struct SessionManager {
session: Arc<RwLock<SessionInfo>>,
config: ClientConfig,
}
impl SessionManager {
pub fn new(config: ClientConfig) -> Self {
Self {
session: Arc::new(RwLock::new(SessionInfo::new())),
config,
}
}
pub async fn session_info(&self) -> SessionInfo {
self.session.read().await.clone()
}
pub async fn session_id(&self) -> McpClientResult<String> {
let session = self.session.read().await;
session
.session_id
.clone()
.ok_or_else(|| SessionError::NotInitialized.into())
}
pub async fn session_id_optional(&self) -> Option<String> {
self.session.read().await.session_id.clone()
}
pub async fn set_session_id(&self, session_id: String) -> McpClientResult<()> {
let mut session = self.session.write().await;
session.session_id = Some(session_id);
Ok(())
}
pub async fn state(&self) -> SessionState {
self.session.read().await.state.clone()
}
pub async fn set_state(&self, state: SessionState) {
let mut session = self.session.write().await;
debug!("Session state transition: {} -> {}", session.state, state);
session.state = state;
session.update_activity();
}
pub async fn is_ready(&self) -> bool {
self.session.read().await.is_ready()
}
pub async fn initialize(
&self,
client_capabilities: ClientCapabilities,
server_capabilities: ServerCapabilities,
protocol_version: String,
) -> McpClientResult<()> {
let mut session = self.session.write().await;
if !matches!(
session.state,
SessionState::Uninitialized | SessionState::Initializing
) {
return Err(SessionError::AlreadyInitialized.into());
}
session.client_capabilities = Some(client_capabilities);
session.server_capabilities = Some(server_capabilities);
session.protocol_version = Some(protocol_version.clone());
session.state = SessionState::Active;
session.update_activity();
info!(
session_id = %session.session_id.as_deref().unwrap_or("None"),
protocol_version = %protocol_version,
"Session initialized successfully"
);
Ok(())
}
pub async fn mark_initializing(&self) -> McpClientResult<()> {
let mut session = self.session.write().await;
if !session.needs_initialization() {
return Err(SessionError::AlreadyInitialized.into());
}
session.state = SessionState::Initializing;
session.connection_attempts += 1;
session.update_activity();
debug!(
session_id = %session.session_id.as_deref().unwrap_or("None"),
attempt = session.connection_attempts,
"Session initialization started"
);
Ok(())
}
pub async fn terminate(&self, reason: Option<String>) {
let mut session = self.session.write().await;
let previous_state = session.state.clone();
let session_id_for_log = session.session_id.as_deref().unwrap_or("None").to_string();
session.state = SessionState::Terminated;
session.session_id = None;
session.update_activity();
info!(
session_id = %session_id_for_log,
previous_state = %previous_state,
reason = reason.as_deref().unwrap_or("user requested"),
"Session terminated"
);
}
pub async fn handle_error(&self, error: String) {
let mut session = self.session.write().await;
let previous_state = session.state.clone();
session.state = SessionState::Error(error.clone());
session.update_activity();
warn!(
session_id = %session.session_id.as_deref().unwrap_or("None"),
previous_state = %previous_state,
error = %error,
"Session encountered error"
);
}
pub async fn start_reconnection(&self) {
let mut session = self.session.write().await;
if matches!(session.state, SessionState::Terminated) {
return; }
session.state = SessionState::Reconnecting;
session.connection_attempts += 1;
session.update_activity();
info!(
session_id = %session.session_id.as_deref().unwrap_or("None"),
attempt = session.connection_attempts,
"Session reconnection started"
);
}
pub async fn reset(&self) {
let mut session = self.session.write().await;
*session = SessionInfo::new();
debug!(
session_id = %session.session_id.as_deref().unwrap_or("None"),
"Session reset for new connection"
);
}
pub async fn update_activity(&self) {
self.session.write().await.update_activity();
}
pub fn create_client_capabilities(&self) -> ClientCapabilities {
let declared = &self.config.declared_capabilities;
ClientCapabilities {
experimental: None,
sampling: declared.sampling.then(Default::default),
elicitation: declared.elicitation.then(Default::default),
roots: declared.roots.then(Default::default),
tasks: None,
}
}
pub async fn create_initialize_request(&self) -> InitializeRequest {
let client_info = &self.config.client_info;
InitializeRequest {
protocol_version: McpVersion::V2025_11_25.as_str().to_string(),
capabilities: self.create_client_capabilities(),
client_info: Implementation {
name: client_info.name.clone(),
version: client_info.version.clone(),
title: None,
description: None,
website_url: None,
icons: None,
},
}
}
pub fn validate_protocol_version(version: &str) -> McpClientResult<()> {
let expected = McpVersion::V2025_11_25.as_str();
if version == expected {
Ok(())
} else {
Err(crate::error::ProtocolError::UnsupportedVersion(format!(
"Server negotiated '{}', the initialize handshake is '{}' only",
version, expected
))
.into())
}
}
pub async fn validate_server_capabilities(
&self,
server_capabilities: &ServerCapabilities,
) -> McpClientResult<()> {
debug!(
tools = ?server_capabilities.tools,
resources = ?server_capabilities.resources,
prompts = ?server_capabilities.prompts,
"Validating server capabilities"
);
Ok(())
}
pub async fn statistics(&self) -> SessionStatistics {
let session = self.session.read().await;
SessionStatistics {
session_id: session.session_id.clone(),
state: session.state.clone(),
duration: session.duration(),
idle_time: session.idle_time(),
connection_attempts: session.connection_attempts,
protocol_version: session.protocol_version.clone(),
}
}
}
#[derive(Debug, Clone)]
pub struct SessionStatistics {
pub session_id: Option<String>,
pub state: SessionState,
pub duration: Duration,
pub idle_time: Duration,
pub connection_attempts: u32,
pub protocol_version: Option<String>,
}
impl SessionStatistics {
pub fn is_healthy(&self) -> bool {
matches!(self.state, SessionState::Active) && self.idle_time < Duration::from_secs(300)
}
pub fn status_summary(&self) -> String {
let session_display = match &self.session_id {
Some(id) => &id[..id.len().min(8)],
None => "None",
};
format!(
"Session {} ({}) - Duration: {:?}, Idle: {:?}, Attempts: {}",
session_display, self.state, self.duration, self.idle_time, self.connection_attempts
)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::ClientConfig;
#[test]
fn test_validate_protocol_version_accepted() {
assert!(SessionManager::validate_protocol_version("2025-11-25").is_ok());
}
#[test]
fn test_validate_protocol_version_rejected() {
let result = SessionManager::validate_protocol_version("2099-01-01");
assert!(result.is_err());
let err = result.unwrap_err();
assert!(err.to_string().contains("2099-01-01"));
}
#[tokio::test]
async fn init_request_advertises_the_legacy_wire_version() {
let manager = SessionManager::new(ClientConfig::default());
let request = manager.create_initialize_request().await;
assert_eq!(request.protocol_version, "2025-11-25");
assert_eq!(request.protocol_version, McpVersion::V2025_11_25.as_str());
}
#[test]
fn init_reply_naming_the_stateless_spec_is_rejected() {
let result = SessionManager::validate_protocol_version(McpVersion::V2026_07_28.as_str());
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("2026-07-28"));
}
#[tokio::test]
async fn test_session_lifecycle() {
let config = ClientConfig::default();
let manager = SessionManager::new(config);
assert_eq!(manager.state().await, SessionState::Uninitialized);
assert!(!manager.is_ready().await);
manager.mark_initializing().await.unwrap();
assert_eq!(manager.state().await, SessionState::Initializing);
let client_caps = manager.create_client_capabilities();
let server_caps = ServerCapabilities {
experimental: None,
logging: None,
prompts: None,
resources: None,
tools: None,
completions: None,
tasks: None,
};
manager
.initialize(client_caps, server_caps, "2025-11-25".to_string())
.await
.unwrap();
assert_eq!(manager.state().await, SessionState::Active);
assert!(manager.is_ready().await);
manager
.set_session_id("session-XYZ".to_string())
.await
.unwrap();
assert_eq!(
manager.session_id_optional().await.as_deref(),
Some("session-XYZ")
);
manager.terminate(Some("test completed".to_string())).await;
assert_eq!(manager.state().await, SessionState::Terminated);
assert!(!manager.is_ready().await);
assert!(
manager.session_id_optional().await.is_none(),
"terminate() must clear session_id so subsequent Drop is a no-op"
);
}
#[tokio::test]
async fn test_session_error_handling() {
let config = ClientConfig::default();
let manager = SessionManager::new(config);
manager.handle_error("test error".to_string()).await;
let SessionState::Error(msg) = manager.state().await else {
panic!("Expected error state, got: {:?}", manager.state().await);
};
assert_eq!(msg, "test error");
}
#[tokio::test]
async fn test_session_reset() {
let config = ClientConfig::default();
let manager = SessionManager::new(config);
manager
.set_session_id("test-session-id".to_string())
.await
.unwrap();
let _original_id = manager.session_id().await.unwrap();
manager.reset().await;
assert!(manager.session_id().await.is_err());
assert_eq!(manager.state().await, SessionState::Uninitialized);
assert!(manager.session_id_optional().await.is_none());
}
}