use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU32, Ordering};
use tokio::sync::{RwLock, mpsc, oneshot, Notify};
use tokio::net::{UnixListener, UnixStream};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use serde::{Serialize, Deserialize};
use crate::{Result, QsshError};
use crate::transport::Transport;
use crate::transport::protocol::{Message, ChannelMessage, ChannelType};
const DEFAULT_WINDOW_SIZE: u32 = 1024 * 1024;
const DEFAULT_MAX_PACKET_SIZE: u32 = 32768;
type PendingChannelMap = HashMap<u32, oneshot::Sender<Result<ChannelOpenResult>>>;
pub struct ControlMaster {
socket_path: PathBuf,
transport: Arc<Transport>,
sessions: Arc<RwLock<HashMap<u32, SessionHandle>>>,
next_session_id: Arc<RwLock<u32>>,
next_channel_id: Arc<AtomicU32>,
shutdown: Arc<AtomicBool>,
shutdown_notify: Arc<Notify>,
pending_channels: Arc<RwLock<PendingChannelMap>>,
}
#[allow(dead_code)]
struct SessionHandle {
id: u32,
channel_id: u32,
remote_channel_id: u32,
to_transport_tx: mpsc::Sender<Vec<u8>>,
to_client_tx: mpsc::Sender<Vec<u8>>,
window_size: Arc<AtomicU32>,
max_packet_size: u32,
}
#[derive(Debug, Clone)]
struct ChannelOpenResult {
local_channel_id: u32,
remote_channel_id: u32,
window_size: u32,
max_packet_size: u32,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum ControlMessage {
NewSession {
command: Option<String>,
env: HashMap<String, String>,
},
SessionCreated {
session_id: u32,
channel_id: u32,
},
Data {
session_id: u32,
data: Vec<u8>,
},
SessionClosed {
session_id: u32,
},
Ping,
Pong,
Exit,
}
impl ControlMaster {
pub fn new(socket_path: PathBuf, transport: Arc<Transport>) -> Self {
Self {
socket_path,
transport,
sessions: Arc::new(RwLock::new(HashMap::new())),
next_session_id: Arc::new(RwLock::new(1)),
next_channel_id: Arc::new(AtomicU32::new(0)),
shutdown: Arc::new(AtomicBool::new(false)),
shutdown_notify: Arc::new(Notify::new()),
pending_channels: Arc::new(RwLock::new(HashMap::new())),
}
}
pub async fn stop(&self) -> Result<()> {
log::info!("Initiating graceful shutdown of control master");
self.shutdown.store(true, Ordering::SeqCst);
self.shutdown_notify.notify_waiters();
let sessions = self.sessions.read().await;
for (session_id, handle) in sessions.iter() {
log::debug!("Closing session {} (channel {})", session_id, handle.channel_id);
let close_msg = Message::Channel(ChannelMessage::Close {
channel_id: handle.remote_channel_id,
});
if let Err(e) = self.transport.send_message(&close_msg).await {
log::warn!("Failed to send close for channel {}: {}", handle.channel_id, e);
}
}
drop(sessions);
self.sessions.write().await.clear();
let mut pending = self.pending_channels.write().await;
for (channel_id, sender) in pending.drain() {
log::debug!("Canceling pending channel open for {}", channel_id);
let _ = sender.send(Err(QsshError::Connection("Control master shutting down".into())));
}
if let Err(e) = std::fs::remove_file(&self.socket_path) {
log::debug!("Could not remove socket file: {}", e);
}
log::info!("Control master shutdown complete");
Ok(())
}
pub async fn start(&self) -> Result<()> {
let _ = std::fs::remove_file(&self.socket_path);
let listener = UnixListener::bind(&self.socket_path)
.map_err(|e| QsshError::Connection(format!("Failed to bind control socket: {}", e)))?;
log::info!("Control master listening on: {:?}", self.socket_path);
let transport_clone = self.transport.clone();
let sessions_clone = self.sessions.clone();
let pending_clone = self.pending_channels.clone();
let shutdown_clone = self.shutdown.clone();
tokio::spawn(async move {
handle_transport_messages(
transport_clone,
sessions_clone,
pending_clone,
shutdown_clone,
).await;
});
loop {
if self.shutdown.load(Ordering::SeqCst) {
log::info!("Control master received shutdown signal");
break;
}
tokio::select! {
accept_result = listener.accept() => {
match accept_result {
Ok((stream, _)) => {
let sessions = self.sessions.clone();
let transport = self.transport.clone();
let next_id = self.next_session_id.clone();
let next_channel = self.next_channel_id.clone();
let pending = self.pending_channels.clone();
let shutdown = self.shutdown.clone();
tokio::spawn(async move {
if let Err(e) = handle_control_client(
stream,
sessions,
transport,
next_id,
next_channel,
pending,
shutdown,
).await {
log::error!("Control client error: {}", e);
}
});
}
Err(e) => {
log::error!("Failed to accept control connection: {}", e);
}
}
}
_ = self.shutdown_notify.notified() => {
log::info!("Control master notified of shutdown");
break;
}
}
}
Ok(())
}
pub async fn check_master(socket_path: &PathBuf) -> bool {
if let Ok(mut stream) = UnixStream::connect(socket_path).await {
let msg = ControlMessage::Ping;
let msg_bytes = match bincode::serialize(&msg) {
Ok(b) => b,
Err(_) => return false,
};
let len_bytes = (msg_bytes.len() as u32).to_be_bytes();
if stream.write_all(&len_bytes).await.is_err() {
return false;
}
if stream.write_all(&msg_bytes).await.is_err() {
return false;
}
let mut buf = vec![0u8; 4];
if stream.read_exact(&mut buf).await.is_err() {
return false;
}
let msg_len = u32::from_be_bytes([buf[0], buf[1], buf[2], buf[3]]) as usize;
let mut msg_buf = vec![0u8; msg_len];
if stream.read_exact(&mut msg_buf).await.is_err() {
return false;
}
if let Ok(response) = bincode::deserialize::<ControlMessage>(&msg_buf) {
matches!(response, ControlMessage::Pong)
} else {
false
}
} else {
false
}
}
}
async fn handle_transport_messages(
transport: Arc<Transport>,
sessions: Arc<RwLock<HashMap<u32, SessionHandle>>>,
pending_channels: Arc<RwLock<PendingChannelMap>>,
shutdown: Arc<AtomicBool>,
) {
loop {
if shutdown.load(Ordering::SeqCst) {
break;
}
let msg_result: Result<Message> = transport.receive_message().await;
match msg_result {
Ok(Message::Channel(channel_msg)) => {
match channel_msg {
ChannelMessage::Accept {
channel_id,
sender_channel,
window_size,
max_packet_size,
} => {
let mut pending = pending_channels.write().await;
if let Some(sender) = pending.remove(&sender_channel) {
let result = ChannelOpenResult {
local_channel_id: sender_channel,
remote_channel_id: channel_id,
window_size,
max_packet_size,
};
let _ = sender.send(Ok(result));
} else {
log::warn!(
"Received channel accept for unknown channel {}",
sender_channel
);
}
}
ChannelMessage::Data { channel_id, data } => {
let sessions_guard = sessions.read().await;
for handle in sessions_guard.values() {
if handle.remote_channel_id == channel_id {
let _ = handle.to_client_tx.send(data.clone()).await;
break;
}
}
}
ChannelMessage::WindowAdjust { channel_id, bytes_to_add } => {
let sessions_guard = sessions.read().await;
for handle in sessions_guard.values() {
if handle.remote_channel_id == channel_id {
handle.window_size.fetch_add(bytes_to_add, Ordering::SeqCst);
break;
}
}
}
ChannelMessage::Eof { channel_id } => {
log::debug!("Received EOF for channel {}", channel_id);
}
ChannelMessage::Close { channel_id } => {
let mut sessions_guard = sessions.write().await;
let session_id_to_remove = sessions_guard
.iter()
.find(|(_, h)| h.remote_channel_id == channel_id)
.map(|(id, _)| *id);
if let Some(session_id) = session_id_to_remove {
sessions_guard.remove(&session_id);
}
}
_ => {
log::debug!("Received unhandled channel message: {:?}", channel_msg);
}
}
}
Ok(Message::Disconnect(disconnect)) => {
log::info!(
"Received disconnect from server: {} (code {})",
disconnect.description,
disconnect.reason_code
);
shutdown.store(true, Ordering::SeqCst);
break;
}
Ok(Message::Ping(seq)) => {
let _ = transport.send_message(&Message::Pong(seq)).await;
}
Ok(_) => {
}
Err(e) => {
if !shutdown.load(Ordering::SeqCst) {
log::error!("Transport receive error: {}", e);
shutdown.store(true, Ordering::SeqCst);
}
break;
}
}
}
}
async fn handle_control_client(
stream: UnixStream,
sessions: Arc<RwLock<HashMap<u32, SessionHandle>>>,
transport: Arc<Transport>,
next_session_id: Arc<RwLock<u32>>,
next_channel_id: Arc<AtomicU32>,
pending_channels: Arc<RwLock<PendingChannelMap>>,
shutdown: Arc<AtomicBool>,
) -> Result<()> {
let (mut stream_reader, stream_writer) = stream.into_split();
let stream_writer = Arc::new(tokio::sync::Mutex::new(stream_writer));
let mut buffer = vec![0u8; 65536];
loop {
if shutdown.load(Ordering::SeqCst) {
break;
}
let n = stream_reader.read(&mut buffer[..4]).await?;
if n == 0 {
break; }
let msg_len = u32::from_be_bytes([buffer[0], buffer[1], buffer[2], buffer[3]]) as usize;
let n = stream_reader.read(&mut buffer[..msg_len]).await?;
if n != msg_len {
return Err(QsshError::Protocol("Incomplete control message".into()));
}
let request: ControlMessage = bincode::deserialize(&buffer[..msg_len])
.map_err(|e| QsshError::Protocol(format!("Failed to parse control message: {}", e)))?;
let response = match request {
ControlMessage::NewSession { command, env } => {
let mut id_guard = next_session_id.write().await;
let session_id = *id_guard;
*id_guard += 1;
drop(id_guard);
match create_channel(
&transport,
&next_channel_id,
&pending_channels,
command,
env,
).await {
Ok(result) => {
let (to_transport_tx, mut to_transport_rx) = mpsc::channel::<Vec<u8>>(256);
let (to_client_tx, mut to_client_rx) = mpsc::channel::<Vec<u8>>(256);
let window_size = Arc::new(AtomicU32::new(result.window_size));
let handle = SessionHandle {
id: session_id,
channel_id: result.local_channel_id,
remote_channel_id: result.remote_channel_id,
to_transport_tx,
to_client_tx,
window_size: window_size.clone(),
max_packet_size: result.max_packet_size,
};
let channel_id = result.local_channel_id;
let remote_channel_id = result.remote_channel_id;
let max_packet_size = result.max_packet_size;
sessions.write().await.insert(session_id, handle);
let transport_clone = transport.clone();
tokio::spawn(async move {
while let Some(data) = to_transport_rx.recv().await {
if let Err(e) = forward_to_channel(
&transport_clone,
remote_channel_id,
data,
max_packet_size,
).await {
log::error!("Failed to forward data to channel {}: {}", channel_id, e);
break;
}
}
});
let client_writer = stream_writer.clone();
let sid = session_id;
tokio::spawn(async move {
while let Some(data) = to_client_rx.recv().await {
let msg = ControlMessage::Data {
session_id: sid,
data,
};
let msg_bytes = match bincode::serialize(&msg) {
Ok(b) => b,
Err(e) => {
log::error!("Failed to serialize server data: {}", e);
break;
}
};
let len_bytes = (msg_bytes.len() as u32).to_be_bytes();
let mut writer = client_writer.lock().await;
if writer.write_all(&len_bytes).await.is_err() {
break;
}
if writer.write_all(&msg_bytes).await.is_err() {
break;
}
}
});
ControlMessage::SessionCreated { session_id, channel_id }
}
Err(e) => {
log::error!("Failed to create channel: {}", e);
let response_bytes = bincode::serialize(&ControlMessage::SessionClosed {
session_id,
}).map_err(|e| QsshError::Protocol(format!("Serialization error: {}", e)))?;
let len_bytes = (response_bytes.len() as u32).to_be_bytes();
let mut writer = stream_writer.lock().await;
writer.write_all(&len_bytes).await?;
writer.write_all(&response_bytes).await?;
continue;
}
}
}
ControlMessage::Data { session_id, data } => {
if let Some(handle) = sessions.read().await.get(&session_id) {
let _ = handle.to_transport_tx.send(data).await;
}
continue; }
ControlMessage::SessionClosed { session_id } => {
if let Some(handle) = sessions.write().await.remove(&session_id) {
let close_msg = Message::Channel(ChannelMessage::Close {
channel_id: handle.remote_channel_id,
});
let _ = transport.send_message(&close_msg).await;
}
continue; }
ControlMessage::Ping => ControlMessage::Pong,
ControlMessage::Exit => {
shutdown.store(true, Ordering::SeqCst);
break;
}
_ => continue,
};
let response_bytes = bincode::serialize(&response)
.map_err(|e| QsshError::Protocol(format!("Failed to serialize response: {}", e)))?;
let len_bytes = (response_bytes.len() as u32).to_be_bytes();
let mut writer = stream_writer.lock().await;
writer.write_all(&len_bytes).await?;
writer.write_all(&response_bytes).await?;
}
Ok(())
}
async fn create_channel(
transport: &Arc<Transport>,
next_channel_id: &Arc<AtomicU32>,
pending_channels: &Arc<RwLock<PendingChannelMap>>,
command: Option<String>,
env: HashMap<String, String>,
) -> Result<ChannelOpenResult> {
let local_channel_id = next_channel_id.fetch_add(1, Ordering::SeqCst);
log::debug!(
"Creating channel {} (command: {:?}, env keys: {:?})",
local_channel_id,
command,
env.keys().collect::<Vec<_>>()
);
let (response_tx, response_rx) = oneshot::channel();
{
let mut pending = pending_channels.write().await;
pending.insert(local_channel_id, response_tx);
}
let channel_open_msg = Message::Channel(ChannelMessage::Open {
channel_id: local_channel_id,
channel_type: ChannelType::Session,
window_size: DEFAULT_WINDOW_SIZE,
max_packet_size: DEFAULT_MAX_PACKET_SIZE,
});
if let Err(e) = transport.send_message(&channel_open_msg).await {
pending_channels.write().await.remove(&local_channel_id);
return Err(QsshError::Connection(format!(
"Failed to send channel open: {}",
e
)));
}
let timeout_duration = tokio::time::Duration::from_secs(30);
let result = tokio::time::timeout(timeout_duration, response_rx).await;
match result {
Ok(Ok(channel_result)) => {
let result = channel_result?;
log::info!(
"Channel {} opened successfully (remote: {}, window: {}, max_packet: {})",
result.local_channel_id,
result.remote_channel_id,
result.window_size,
result.max_packet_size
);
if let Some(cmd) = command {
let exec_msg = Message::Channel(ChannelMessage::ExecRequest {
channel_id: result.remote_channel_id,
command: cmd.clone(),
});
transport.send_message(&exec_msg).await?;
log::debug!("Sent exec request for command: {}", cmd);
} else {
let shell_msg = Message::Channel(ChannelMessage::ShellRequest {
channel_id: result.remote_channel_id,
});
transport.send_message(&shell_msg).await?;
log::debug!("Sent shell request");
}
for (key, value) in env {
log::debug!("Setting env {}={}", key, value);
}
Ok(result)
}
Ok(Err(_)) => {
pending_channels.write().await.remove(&local_channel_id);
Err(QsshError::Connection(
"Channel open cancelled - receiver dropped".into(),
))
}
Err(_) => {
pending_channels.write().await.remove(&local_channel_id);
Err(QsshError::Connection(format!(
"Channel open timed out after {} seconds",
timeout_duration.as_secs()
)))
}
}
}
async fn forward_to_channel(
transport: &Arc<Transport>,
remote_channel_id: u32,
data: Vec<u8>,
max_packet_size: u32,
) -> Result<()> {
if data.is_empty() {
return Ok(());
}
let max_size = max_packet_size as usize;
for chunk in data.chunks(max_size) {
let data_msg = Message::Channel(ChannelMessage::Data {
channel_id: remote_channel_id,
data: chunk.to_vec(),
});
transport.send_message(&data_msg).await.map_err(|e| {
QsshError::Connection(format!(
"Failed to forward data to channel {}: {}",
remote_channel_id, e
))
})?;
}
log::trace!(
"Forwarded {} bytes to channel {}",
data.len(),
remote_channel_id
);
Ok(())
}
pub struct ControlClient {
socket_path: PathBuf,
stream: Option<UnixStream>,
}
impl ControlClient {
pub fn new(socket_path: PathBuf) -> Self {
Self {
socket_path,
stream: None,
}
}
pub async fn connect(&mut self) -> Result<()> {
let stream = UnixStream::connect(&self.socket_path).await
.map_err(|e| QsshError::Connection(format!("Failed to connect to control master: {}", e)))?;
self.stream = Some(stream);
Ok(())
}
pub async fn new_session(&mut self, command: Option<String>) -> Result<u32> {
let stream = self.stream.as_mut()
.ok_or_else(|| QsshError::Protocol("Not connected to control master".into()))?;
let msg = ControlMessage::NewSession {
command,
env: HashMap::new(),
};
let msg_bytes = bincode::serialize(&msg)
.map_err(|e| QsshError::Protocol(format!("Failed to serialize message: {}", e)))?;
let len_bytes = (msg_bytes.len() as u32).to_be_bytes();
stream.write_all(&len_bytes).await?;
stream.write_all(&msg_bytes).await?;
let mut buf = vec![0u8; 4];
stream.read_exact(&mut buf).await?;
let msg_len = u32::from_be_bytes([buf[0], buf[1], buf[2], buf[3]]) as usize;
let mut msg_buf = vec![0u8; msg_len];
stream.read_exact(&mut msg_buf).await?;
let response: ControlMessage = bincode::deserialize(&msg_buf)
.map_err(|e| QsshError::Protocol(format!("Failed to parse response: {}", e)))?;
match response {
ControlMessage::SessionCreated { session_id, .. } => Ok(session_id),
_ => Err(QsshError::Protocol("Unexpected response from control master".into())),
}
}
pub async fn send_data(&mut self, session_id: u32, data: Vec<u8>) -> Result<()> {
let stream = self.stream.as_mut()
.ok_or_else(|| QsshError::Protocol("Not connected to control master".into()))?;
let msg = ControlMessage::Data { session_id, data };
let msg_bytes = bincode::serialize(&msg)
.map_err(|e| QsshError::Protocol(format!("Failed to serialize message: {}", e)))?;
let len_bytes = (msg_bytes.len() as u32).to_be_bytes();
stream.write_all(&len_bytes).await?;
stream.write_all(&msg_bytes).await?;
Ok(())
}
pub async fn close_session(&mut self, session_id: u32) -> Result<()> {
let stream = self.stream.as_mut()
.ok_or_else(|| QsshError::Protocol("Not connected to control master".into()))?;
let msg = ControlMessage::SessionClosed { session_id };
let msg_bytes = bincode::serialize(&msg)
.map_err(|e| QsshError::Protocol(format!("Failed to serialize message: {}", e)))?;
let len_bytes = (msg_bytes.len() as u32).to_be_bytes();
stream.write_all(&len_bytes).await?;
stream.write_all(&msg_bytes).await?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::TempDir;
#[tokio::test]
async fn test_control_master_ping() {
let temp_dir = TempDir::new().unwrap();
let socket_path = temp_dir.path().join("control.sock");
assert!(!ControlMaster::check_master(&socket_path).await);
}
#[tokio::test]
async fn test_control_client_creation() {
let temp_dir = TempDir::new().unwrap();
let socket_path = temp_dir.path().join("control.sock");
let client = ControlClient::new(socket_path);
assert!(client.stream.is_none());
}
#[tokio::test]
async fn test_control_message_serialization() {
let messages = vec![
ControlMessage::NewSession {
command: Some("ls -la".to_string()),
env: {
let mut env = HashMap::new();
env.insert("TERM".to_string(), "xterm-256color".to_string());
env
},
},
ControlMessage::SessionCreated {
session_id: 42,
channel_id: 100,
},
ControlMessage::Data {
session_id: 42,
data: vec![0x48, 0x65, 0x6c, 0x6c, 0x6f], },
ControlMessage::SessionClosed { session_id: 42 },
ControlMessage::Ping,
ControlMessage::Pong,
ControlMessage::Exit,
];
for msg in messages {
let serialized = bincode::serialize(&msg).expect("Serialization should succeed");
let deserialized: ControlMessage =
bincode::deserialize(&serialized).expect("Deserialization should succeed");
assert_eq!(format!("{:?}", msg), format!("{:?}", deserialized));
}
}
#[tokio::test]
async fn test_channel_open_result() {
let result = ChannelOpenResult {
local_channel_id: 5,
remote_channel_id: 10,
window_size: DEFAULT_WINDOW_SIZE,
max_packet_size: DEFAULT_MAX_PACKET_SIZE,
};
assert_eq!(result.local_channel_id, 5);
assert_eq!(result.remote_channel_id, 10);
assert_eq!(result.window_size, 1024 * 1024);
assert_eq!(result.max_packet_size, 32768);
}
#[tokio::test]
async fn test_session_handle_window_size() {
let (to_transport_tx, _rx1) = mpsc::channel(256);
let (to_client_tx, _rx2) = mpsc::channel(256);
let window_size = Arc::new(AtomicU32::new(DEFAULT_WINDOW_SIZE));
let handle = SessionHandle {
id: 1,
channel_id: 0,
remote_channel_id: 100,
to_transport_tx,
to_client_tx,
window_size: window_size.clone(),
max_packet_size: DEFAULT_MAX_PACKET_SIZE,
};
assert_eq!(handle.window_size.load(Ordering::SeqCst), DEFAULT_WINDOW_SIZE);
handle.window_size.fetch_add(4096, Ordering::SeqCst);
assert_eq!(
handle.window_size.load(Ordering::SeqCst),
DEFAULT_WINDOW_SIZE + 4096
);
handle.window_size.fetch_sub(1024, Ordering::SeqCst);
assert_eq!(
handle.window_size.load(Ordering::SeqCst),
DEFAULT_WINDOW_SIZE + 4096 - 1024
);
}
#[tokio::test]
async fn test_control_client_not_connected() {
let temp_dir = TempDir::new().unwrap();
let socket_path = temp_dir.path().join("control.sock");
let mut client = ControlClient::new(socket_path);
let result = client.new_session(None).await;
assert!(result.is_err());
assert!(result
.unwrap_err()
.to_string()
.contains("Not connected to control master"));
let result = client.send_data(1, vec![1, 2, 3]).await;
assert!(result.is_err());
let result = client.close_session(1).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_control_client_connect_nonexistent() {
let temp_dir = TempDir::new().unwrap();
let socket_path = temp_dir.path().join("nonexistent.sock");
let mut client = ControlClient::new(socket_path);
let result = client.connect().await;
assert!(result.is_err());
assert!(result
.unwrap_err()
.to_string()
.contains("Failed to connect to control master"));
}
#[test]
fn test_default_constants() {
assert_eq!(DEFAULT_WINDOW_SIZE, 1024 * 1024); assert_eq!(DEFAULT_MAX_PACKET_SIZE, 32768);
let max_packet = DEFAULT_MAX_PACKET_SIZE;
let window = DEFAULT_WINDOW_SIZE;
assert!(max_packet < window, "max packet size must be less than window size");
}
#[tokio::test]
async fn test_data_chunking_logic() {
let large_data = vec![0u8; 100_000]; let max_packet_size = 32768u32;
let chunks: Vec<&[u8]> = large_data.chunks(max_packet_size as usize).collect();
assert_eq!(chunks.len(), 4);
assert_eq!(chunks[0].len(), 32768);
assert_eq!(chunks[1].len(), 32768);
assert_eq!(chunks[2].len(), 32768);
assert_eq!(chunks[3].len(), 100_000 - 3 * 32768);
}
#[tokio::test]
async fn test_atomic_channel_id_generation() {
let next_channel_id = Arc::new(AtomicU32::new(0));
let mut handles = vec![];
for _ in 0..100 {
let id_clone = next_channel_id.clone();
handles.push(tokio::spawn(async move {
id_clone.fetch_add(1, Ordering::SeqCst)
}));
}
let mut ids = vec![];
for handle in handles {
ids.push(handle.await.unwrap());
}
ids.sort();
ids.dedup();
assert_eq!(ids.len(), 100);
assert_eq!(next_channel_id.load(Ordering::SeqCst), 100);
}
#[tokio::test]
async fn test_pending_channels_cleanup() {
let pending_channels: Arc<RwLock<PendingChannelMap>> =
Arc::new(RwLock::new(HashMap::new()));
for i in 0..5 {
let (tx, _rx) = oneshot::channel();
pending_channels.write().await.insert(i, tx);
}
assert_eq!(pending_channels.read().await.len(), 5);
{
let mut pending = pending_channels.write().await;
for (channel_id, sender) in pending.drain() {
let _ = sender.send(Err(QsshError::Connection(format!(
"Channel {} cancelled",
channel_id
))));
}
}
assert_eq!(pending_channels.read().await.len(), 0);
}
}