use super::*;
use crate::{Error, Result, transport::Transport, utils::Utils};
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::{Mutex, RwLock, mpsc};
use uuid::Uuid;
pub struct Session {
id: String,
state: Arc<RwLock<SessionState>>,
transport: Arc<Mutex<Box<dyn Transport>>>,
router: MessageRouter,
pending_requests: Arc<Mutex<HashMap<RequestId, PendingRequest>>>,
event_sender: mpsc::UnboundedSender<SessionEvent>,
}
#[derive(Debug, Clone)]
pub struct SessionState {
pub initialized: bool,
pub client_info: Option<Implementation>,
pub server_info: Option<Implementation>,
pub client_capabilities: Option<ClientCapabilities>,
pub server_capabilities: Option<ServerCapabilities>,
pub protocol_version: Option<String>,
pub connected_at: chrono::DateTime<chrono::Utc>,
pub last_activity: chrono::DateTime<chrono::Utc>,
}
#[derive(Debug)]
struct PendingRequest {
sender: tokio::sync::oneshot::Sender<Result<JsonRpcResponse>>,
created_at: std::time::Instant,
timeout: Option<std::time::Duration>,
}
#[derive(Debug, Clone)]
pub enum SessionEvent {
Connected,
Disconnected,
Initialized {
client_info: Implementation,
},
MessageReceived {
message: String,
},
MessageSent {
message: String,
},
Error {
error: String,
},
}
impl Session {
pub fn new(
transport: Box<dyn Transport>,
handler: Arc<dyn MessageHandler>,
) -> (Self, mpsc::UnboundedReceiver<SessionEvent>) {
let id = Uuid::new_v4().to_string();
let router = MessageRouter::new(handler);
let (event_sender, event_receiver) = mpsc::unbounded_channel();
let session = Self {
id,
state: Arc::new(RwLock::new(SessionState {
initialized: false,
client_info: None,
server_info: None,
client_capabilities: None,
server_capabilities: None,
protocol_version: None,
connected_at: chrono::Utc::now(),
last_activity: chrono::Utc::now(),
})),
transport: Arc::new(Mutex::new(transport)),
router,
pending_requests: Arc::new(Mutex::new(HashMap::new())),
event_sender,
};
(session, event_receiver)
}
pub fn id(&self) -> &str {
&self.id
}
pub async fn state(&self) -> SessionState {
self.state.read().await.clone()
}
pub async fn is_initialized(&self) -> bool {
self.state.read().await.initialized
}
pub async fn send_request(&self, request: JsonRpcRequest) -> Result<JsonRpcResponse> {
let request_id = request
.id
.clone()
.ok_or_else(|| Error::InvalidRequest("Request must have an ID".to_string()))?;
let (tx, rx) = tokio::sync::oneshot::channel();
{
let mut pending = self.pending_requests.lock().await;
pending.insert(
request_id.clone(),
PendingRequest {
sender: tx,
created_at: std::time::Instant::now(),
timeout: Some(std::time::Duration::from_secs(30)),
},
);
}
let message = Protocol::serialize_message(&JsonRpcMessage::Request(request))?;
self.send_message(&message).await?;
match rx.await {
Ok(result) => result,
Err(_) => {
self.pending_requests.lock().await.remove(&request_id);
Err(Error::Timeout)
}
}
}
pub async fn send_notification(&self, notification: JsonRpcNotification) -> Result<()> {
let message = Protocol::serialize_message(&JsonRpcMessage::Notification(notification))?;
self.send_message(&message).await
}
async fn send_message(&self, message: &str) -> Result<()> {
{
let mut transport = self.transport.lock().await;
transport.send(message).await?;
}
{
let mut state = self.state.write().await;
state.last_activity = chrono::Utc::now();
}
let _ = self.event_sender.send(SessionEvent::MessageSent {
message: message.to_string(),
});
Ok(())
}
pub async fn run(&self) -> Result<()> {
let _ = self.event_sender.send(SessionEvent::Connected);
loop {
let message = {
let mut transport = self.transport.lock().await;
transport.receive().await?
};
let message = match message {
Some(msg) => msg,
None => {
let _ = self.event_sender.send(SessionEvent::Disconnected);
break;
}
};
{
let mut state = self.state.write().await;
state.last_activity = chrono::Utc::now();
}
let _ = self.event_sender.send(SessionEvent::MessageReceived {
message: message.clone(),
});
if let Err(e) = self.process_message(&message).await {
let _ = self.event_sender.send(SessionEvent::Error {
error: e.to_string(),
});
}
}
Ok(())
}
async fn process_message(&self, message: &str) -> Result<()> {
let jsonrpc_message = Protocol::parse_message(message)?;
match &jsonrpc_message {
JsonRpcMessage::Response(response) => {
self.handle_response(response).await?;
}
JsonRpcMessage::Request(_) | JsonRpcMessage::Notification(_) => {
if let Some(response_message) = self.router.route_message(jsonrpc_message).await? {
let response_str = Protocol::serialize_message(&response_message)?;
self.send_message(&response_str).await?;
}
}
}
Ok(())
}
async fn handle_response(&self, response: &JsonRpcResponse) -> Result<()> {
if let Some(ref response_id) = response.id {
let mut pending = self.pending_requests.lock().await;
if let Some(pending_request) = pending.remove(response_id) {
let _ = pending_request.sender.send(Ok(response.clone()));
}
}
Ok(())
}
pub async fn initialize(
&self,
client_info: Implementation,
client_capabilities: ClientCapabilities,
) -> Result<InitializeResponse> {
let request = JsonRpcRequest {
jsonrpc: "2.0".to_string(),
id: Some(Protocol::generate_request_id()),
method: "initialize".to_string(),
params: Some(Utils::to_json_value(&InitializeRequest {
protocol_version: Protocol::latest_version().to_string(),
capabilities: client_capabilities.clone(),
client_info: client_info.clone(),
})?),
};
let response = self.send_request(request).await?;
if let Some(error) = response.error {
return Err(Error::Server(format!(
"Initialize failed: {}",
error.message
)));
}
let result = response
.result
.ok_or_else(|| Error::Server("Initialize response missing result".to_string()))?;
let init_response: InitializeResponse = Utils::from_json_value(result)?;
{
let mut state = self.state.write().await;
state.client_info = Some(client_info);
state.client_capabilities = Some(client_capabilities);
state.server_info = Some(init_response.server_info.clone());
state.server_capabilities = Some(init_response.capabilities.clone());
state.protocol_version = Some(init_response.protocol_version.clone());
state.initialized = true;
}
let initialized_notification = JsonRpcNotification {
jsonrpc: "2.0".to_string(),
method: "initialized".to_string(),
params: Some(Utils::to_json_value(&InitializedNotification {})?),
};
self.send_notification(initialized_notification).await?;
let _ = self.event_sender.send(SessionEvent::Initialized {
client_info: init_response.server_info.clone(),
});
Ok(init_response)
}
pub async fn close(&self) -> Result<()> {
let mut transport = self.transport.lock().await;
transport.close().await?;
let _ = self.event_sender.send(SessionEvent::Disconnected);
Ok(())
}
pub async fn is_connected(&self) -> bool {
let transport = self.transport.lock().await;
transport.is_connected()
}
pub async fn cleanup_expired_requests(&self) {
let mut pending = self.pending_requests.lock().await;
let now = std::time::Instant::now();
let mut timed_out_requests = Vec::new();
pending.retain(|id, request| {
if let Some(timeout) = request.timeout {
if now.duration_since(request.created_at) > timeout {
timed_out_requests.push(id.clone());
false
} else {
true
}
} else {
true
}
});
for id in timed_out_requests {
if let Some(request) = pending.remove(&id) {
let _ = request.sender.send(Err(Error::Timeout));
}
}
}
}
impl Default for SessionState {
fn default() -> Self {
let now = chrono::Utc::now();
Self {
initialized: false,
client_info: None,
server_info: None,
client_capabilities: None,
server_capabilities: None,
protocol_version: None,
connected_at: now,
last_activity: now,
}
}
}