use crate::{
Result, QsshError, QsshConfig, PortForward,
transport::{Transport, Message, ChannelMessage, ChannelType, ChannelManager, RekeyMessage},
crypto::{PqKeyExchange, SymmetricCrypto},
handshake::ClientHandshake,
vault::QuantumVault,
x11::{X11Forwarder, setup_x11_forwarding},
port_forward::{RemoteForwardRegistry, ForwardedChannelRouter, handle_forwarded_channel},
};
use tokio::net::TcpStream;
use std::sync::Arc;
use std::time::Duration;
use socket2::{Socket, TcpKeepalive};
const KEEPALIVE_IDLE_SECS: u64 = 30;
const KEEPALIVE_INTERVAL_SECS: u64 = 10;
const KEEPALIVE_RETRIES: u32 = 3;
fn apply_tcp_keepalive(stream: TcpStream) -> Result<TcpStream> {
let std_stream = stream.into_std()
.map_err(|e| QsshError::Connection(format!("into_std failed: {}", e)))?;
std_stream.set_nonblocking(true)
.map_err(|e| QsshError::Connection(format!("set_nonblocking failed: {}", e)))?;
let sock = Socket::from(std_stream);
let ka = TcpKeepalive::new()
.with_time(Duration::from_secs(KEEPALIVE_IDLE_SECS))
.with_interval(Duration::from_secs(KEEPALIVE_INTERVAL_SECS))
.with_retries(KEEPALIVE_RETRIES);
sock.set_tcp_keepalive(&ka)
.map_err(|e| QsshError::Connection(format!("set_tcp_keepalive failed: {}", e)))?;
let std_stream: std::net::TcpStream = sock.into();
TcpStream::from_std(std_stream)
.map_err(|e| QsshError::Connection(format!("from_std failed: {}", e)))
}
#[derive(Debug, Clone)]
pub struct ReconnectConfig {
pub max_attempts: u32,
pub initial_delay: Duration,
pub max_delay: Duration,
pub multiplier: f64,
}
impl Default for ReconnectConfig {
fn default() -> Self {
Self {
max_attempts: 10,
initial_delay: Duration::from_millis(500),
max_delay: Duration::from_secs(30),
multiplier: 2.0,
}
}
}
impl ReconnectConfig {
pub fn persistent() -> Self {
Self {
max_attempts: 0,
..Default::default()
}
}
}
enum ExecEvent {
Accepted,
Data(Vec<u8>),
ExitStatus(u32),
Eof,
Close,
Error(String),
}
pub struct QsshClient {
config: QsshConfig,
transport: Option<Transport>,
channel_manager: Arc<ChannelManager>,
vault: Option<Arc<QuantumVault>>,
x11_forwarder: Option<X11Forwarder>,
agent_forward_enabled: bool,
remote_forward_registry: RemoteForwardRegistry,
forwarded_channel_router: ForwardedChannelRouter,
rekey_task: Option<tokio::task::JoinHandle<()>>,
}
impl QsshClient {
pub fn new(config: QsshConfig) -> Self {
Self {
config,
transport: None,
channel_manager: Arc::new(ChannelManager::new()),
vault: None,
x11_forwarder: None,
agent_forward_enabled: false,
remote_forward_registry: RemoteForwardRegistry::new(),
forwarded_channel_router: ForwardedChannelRouter::new(),
rekey_task: None,
}
}
pub fn set_remote_forward_state(&mut self, registry: RemoteForwardRegistry, router: ForwardedChannelRouter) {
self.remote_forward_registry = registry;
self.forwarded_channel_router = router;
}
pub fn has_remote_forwards(&self) -> bool {
true }
pub async fn run_forward_loop(&self) -> Result<()> {
let transport = self.transport.as_ref()
.ok_or_else(|| QsshError::Protocol("Not connected".into()))?;
log::info!("Entering forward dispatch loop (exec completed, -R active)");
loop {
match transport.receive_message::<Message>().await {
Ok(Message::Channel(ChannelMessage::Open {
channel_id,
channel_type: ChannelType::ForwardedTcpip {
connected_host, connected_port,
originator_host, originator_port,
},
..
})) => {
log::info!("Forwarded channel {} open from {}:{} (originator {}:{})",
channel_id, connected_host, connected_port,
originator_host, originator_port);
let transport = transport.clone();
let registry = self.remote_forward_registry.clone();
let router = self.forwarded_channel_router.clone();
tokio::spawn(async move {
if let Err(e) = handle_forwarded_channel(
channel_id, connected_host, connected_port,
Arc::new(transport), registry, router,
).await {
log::error!("Forwarded channel {} error: {}", channel_id, e);
}
});
}
Ok(Message::Channel(ChannelMessage::Accept { channel_id, .. })) => {
self.forwarded_channel_router.route_data(channel_id, Vec::new()).await;
}
Ok(Message::Channel(ChannelMessage::Data { channel_id, data })) => {
if !self.forwarded_channel_router.route_data(channel_id, data).await {
log::debug!("No handler for channel {} data in forward loop", channel_id);
}
}
Ok(Message::Channel(ChannelMessage::Eof { channel_id })) => {
self.forwarded_channel_router.remove(channel_id).await;
}
Ok(Message::Channel(ChannelMessage::Close { channel_id })) => {
self.forwarded_channel_router.remove(channel_id).await;
}
Ok(Message::Disconnect(_)) => {
log::info!("Server disconnected during forward loop");
break;
}
Ok(_) => {}
Err(e) => {
log::debug!("Forward loop transport error: {}", e);
break;
}
}
}
Ok(())
}
pub async fn with_vault(mut self, master_key: &[u8]) -> Result<Self> {
let mut vault = QuantumVault::new();
vault.init(master_key).await?;
self.vault = Some(Arc::new(vault));
Ok(self)
}
pub async fn connect(&mut self) -> Result<()> {
let addr = self.config.server.clone();
let stream = TcpStream::connect(&addr).await
.map_err(|e| QsshError::Connection(format!("Failed to connect to {}: {}", addr, e)))?;
let stream = match apply_tcp_keepalive(stream) {
Ok(s) => { log::debug!("TCP keepalive enabled ({}s idle / {}s interval / {} retries)",
KEEPALIVE_IDLE_SECS, KEEPALIVE_INTERVAL_SECS, KEEPALIVE_RETRIES); s }
Err(e) => return Err(e),
};
log::info!("Connected to {}", addr);
self.connect_with_stream(stream).await
}
pub async fn connect_via_stream(&mut self, stream: TcpStream) -> Result<()> {
log::info!("Connecting via pre-established stream (ProxyJump)");
self.connect_with_stream(stream).await
}
async fn connect_with_stream(&mut self, stream: TcpStream) -> Result<()> {
let handshake = ClientHandshake::new(&self.config, stream);
let transport = handshake.perform().await?;
log::info!("Handshake completed successfully");
if let Some(ref mut pw) = self.config.password {
use zeroize::Zeroize;
pw.zeroize();
}
self.config.password = None;
self.transport = Some(transport);
for forward in &self.config.port_forwards {
self.setup_port_forward(forward.clone()).await?;
}
let interval_secs = self.config.key_rotation_interval;
if interval_secs > 0 {
if let Some(transport) = self.transport.clone() {
let interval = Duration::from_secs(interval_secs);
log::info!("Automatic rekey scheduled every {}s", interval_secs);
let handle = tokio::spawn(async move {
loop {
tokio::time::sleep(interval).await;
log::info!("Rekey timer fired — initiating key rotation");
if let Err(e) = rekey_with_transport(&transport).await {
log::warn!("Automatic rekey failed: {} (will retry next interval)", e);
}
}
});
self.rekey_task = Some(handle);
}
}
Ok(())
}
pub async fn connect_with_retry(&mut self, reconnect: &ReconnectConfig) -> Result<()> {
let mut attempt = 0u32;
let mut delay = reconnect.initial_delay;
loop {
attempt += 1;
match self.connect().await {
Ok(()) => return Ok(()),
Err(e) => {
if reconnect.max_attempts > 0 && attempt >= reconnect.max_attempts {
log::error!("Connection failed after {} attempts: {}", attempt, e);
return Err(e);
}
log::warn!(
"Connection attempt {}{} failed: {}. Retrying in {:.1}s...",
attempt,
if reconnect.max_attempts > 0 { format!("/{}", reconnect.max_attempts) } else { String::new() },
e,
delay.as_secs_f64()
);
tokio::time::sleep(delay).await;
delay = Duration::from_secs_f64(
(delay.as_secs_f64() * reconnect.multiplier).min(reconnect.max_delay.as_secs_f64())
);
}
}
}
}
pub fn transport(&self) -> Option<&Transport> {
self.transport.as_ref()
}
pub async fn enable_x11(&mut self, trusted: bool) -> Result<()> {
if let Some(ref transport) = self.transport {
let forwarder = setup_x11_forwarding(
Arc::new(transport.clone()),
true,
trusted
).await?;
if let Some(fwd) = forwarder {
log::info!("X11 forwarding enabled: DISPLAY={}", fwd.get_display());
self.x11_forwarder = Some(fwd);
}
}
Ok(())
}
pub async fn enable_agent_forwarding(&mut self) -> Result<()> {
let auth_sock = std::env::var("QSSH_AUTH_SOCK").map_err(|_| {
QsshError::Connection("QSSH_AUTH_SOCK not set — is qssh-agent running?".into())
})?;
if !std::path::Path::new(&auth_sock).exists() {
return Err(QsshError::Connection(format!(
"Agent socket not found: {}", auth_sock
)));
}
let transport = self.transport.as_ref()
.ok_or_else(|| QsshError::Protocol("Not connected".into()))?;
let msg = Message::Channel(ChannelMessage::AgentForwardRequest {
channel_id: 0,
});
transport.send_message(&msg).await?;
log::info!("Agent forwarding requested (local socket: {})", auth_sock);
self.agent_forward_enabled = true;
Ok(())
}
pub async fn rekey(&self) -> Result<()> {
use sha3::Sha3_256;
use hkdf::Hkdf;
let transport = self.transport.as_ref()
.ok_or_else(|| QsshError::Protocol("Not connected".into()))?;
log::info!("Initiating rekey (rekey #{})", transport.rekey_count() + 1);
let client_kex = PqKeyExchange::new()?;
let (client_share, client_sig) = client_kex.create_key_share()?;
let rekey_msg = RekeyMessage {
new_falcon_public_key: client_kex.falcon_pk.clone(),
new_key_share: client_share.clone(),
new_key_share_signature: client_sig,
request_qkd: false,
};
transport.send_message(&Message::Rekey(rekey_msg)).await?;
let server_rekey = loop {
match transport.receive_message::<Message>().await? {
Message::Rekey(rekey) => break rekey,
other => {
log::debug!("Received non-rekey message during rekey: {:?}", other);
continue;
}
}
};
let verified = client_kex.verify_falcon(
&server_rekey.new_key_share,
&server_rekey.new_key_share_signature,
&server_rekey.new_falcon_public_key,
)?;
if !verified {
return Err(QsshError::Crypto("Server rekey signature verification failed".into()));
}
let mut ikm = Vec::new();
ikm.extend_from_slice(&client_share);
ikm.extend_from_slice(&server_rekey.new_key_share);
let hk = Hkdf::<Sha3_256>::new(None, &ikm);
let mut server_write_key = vec![0u8; 32];
let mut client_write_key = vec![0u8; 32];
hk.expand(b"qssh-rekey-server-write", &mut server_write_key)
.map_err(|_| QsshError::Crypto("HKDF expand failed".into()))?;
hk.expand(b"qssh-rekey-client-write", &mut client_write_key)
.map_err(|_| QsshError::Crypto("HKDF expand failed".into()))?;
let new_send = SymmetricCrypto::from_shared_secret(&client_write_key)?;
let new_recv = SymmetricCrypto::from_shared_secret(&server_write_key)?;
transport.update_keys(new_send, new_recv).await?;
log::info!("Rekey complete — new session keys active");
Ok(())
}
pub async fn shell(&self) -> Result<()> {
let transport = self.transport.as_ref()
.ok_or_else(|| QsshError::Protocol("Not connected".into()))?;
let channel_id = self.channel_manager.open_channel(ChannelType::Session).await?;
let open_msg = Message::Channel(ChannelMessage::Open {
channel_id,
channel_type: ChannelType::Session,
window_size: 1024 * 1024,
max_packet_size: 32768,
});
transport.send_message(&open_msg).await?;
let router = self.forwarded_channel_router.clone();
let start = std::time::Instant::now();
loop {
if start.elapsed() > std::time::Duration::from_secs(5) {
return Err(QsshError::Protocol("Timeout waiting for channel accept".into()));
}
match transport.receive_message::<Message>().await {
Ok(Message::Channel(ChannelMessage::Accept { channel_id: ch_id, .. })) if ch_id == channel_id => {
log::debug!("Channel {} accepted by server", channel_id);
break;
}
Ok(Message::Channel(ChannelMessage::Accept { channel_id: ch_id, .. })) => {
router.route_data(ch_id, Vec::new()).await;
}
Ok(Message::Channel(ChannelMessage::Data { channel_id: ch_id, data })) => {
router.route_data(ch_id, data).await;
}
Ok(Message::Channel(ChannelMessage::Eof { channel_id: ch_id })) => {
router.remove(ch_id).await;
}
Ok(Message::Channel(ChannelMessage::Close { channel_id: ch_id })) => {
router.remove(ch_id).await;
}
Ok(msg) => {
log::debug!("Received while waiting for accept: {:?}", msg);
}
Err(e) => {
return Err(e);
}
}
}
self.handle_shell_session(channel_id).await
}
pub async fn exec_with_status(&self, command: &str) -> Result<(String, u32)> {
let transport = self.transport.as_ref()
.ok_or_else(|| QsshError::Protocol("Not connected".into()))?;
let channel_id = self.channel_manager.open_channel(ChannelType::Session).await?;
let open_msg = Message::Channel(ChannelMessage::Open {
channel_id,
channel_type: ChannelType::Session,
window_size: 1024 * 1024,
max_packet_size: 32768,
});
transport.send_message(&open_msg).await?;
let (exec_tx, mut exec_rx) = tokio::sync::mpsc::channel::<ExecEvent>(64);
let transport_reader = transport.clone();
let router = self.forwarded_channel_router.clone();
let registry = self.remote_forward_registry.clone();
let reader_handle = tokio::spawn(async move {
loop {
match transport_reader.receive_message::<Message>().await {
Ok(Message::Channel(ChannelMessage::Accept { channel_id: ch_id, .. })) if ch_id == channel_id => {
let _ = exec_tx.send(ExecEvent::Accepted).await;
}
Ok(Message::Channel(ChannelMessage::Data { channel_id: ch_id, data })) if ch_id == channel_id => {
let _ = exec_tx.send(ExecEvent::Data(data)).await;
}
Ok(Message::Channel(ChannelMessage::ExitStatus { channel_id: ch_id, exit_code: code })) if ch_id == channel_id => {
let _ = exec_tx.send(ExecEvent::ExitStatus(code)).await;
}
Ok(Message::Channel(ChannelMessage::Eof { channel_id: ch_id })) if ch_id == channel_id => {
let _ = exec_tx.send(ExecEvent::Eof).await;
}
Ok(Message::Channel(ChannelMessage::Close { channel_id: ch_id })) if ch_id == channel_id => {
let _ = exec_tx.send(ExecEvent::Close).await;
}
Ok(Message::Channel(ChannelMessage::Open {
channel_id: ch_id,
channel_type: ChannelType::ForwardedTcpip {
connected_host, connected_port, ..
},
..
})) => {
log::info!("Forwarded channel {} open for {}:{} (during exec)",
ch_id, connected_host, connected_port);
let t = transport_reader.clone();
let reg = registry.clone();
let rtr = router.clone();
tokio::spawn(async move {
if let Err(e) = handle_forwarded_channel(
ch_id, connected_host, connected_port,
Arc::new(t), reg, rtr,
).await {
log::error!("Forwarded channel {} error: {}", ch_id, e);
}
});
}
Ok(Message::Channel(ChannelMessage::Accept { channel_id: ch_id, .. })) => {
router.route_data(ch_id, Vec::new()).await;
}
Ok(Message::Channel(ChannelMessage::Data { channel_id: ch_id, data })) => {
router.route_data(ch_id, data).await;
}
Ok(Message::Channel(ChannelMessage::Eof { channel_id: ch_id })) => {
router.remove(ch_id).await;
}
Ok(Message::Channel(ChannelMessage::Close { channel_id: ch_id })) => {
router.remove(ch_id).await;
}
Ok(Message::Disconnect(_)) => {
let _ = exec_tx.send(ExecEvent::Error("Server disconnected".into())).await;
break;
}
Ok(_) => {} Err(e) => {
let _ = exec_tx.send(ExecEvent::Error(format!("{}", e))).await;
break;
}
}
}
});
let accept_timeout = tokio::time::timeout(
std::time::Duration::from_secs(5),
async {
while let Some(event) = exec_rx.recv().await {
match event {
ExecEvent::Accepted => return Ok(()),
ExecEvent::Error(e) => return Err(QsshError::Protocol(e)),
_ => {} }
}
Err(QsshError::Protocol("Reader task exited before accept".into()))
},
).await;
match accept_timeout {
Ok(Ok(())) => {} Ok(Err(e)) => {
reader_handle.abort();
return Err(e);
}
Err(_) => {
reader_handle.abort();
return Err(QsshError::Protocol("Timeout waiting for channel accept".into()));
}
}
let exec_msg = Message::Channel(ChannelMessage::ExecRequest {
channel_id,
command: command.to_string(),
});
transport.send_message(&exec_msg).await?;
let mut output = Vec::new();
let mut exit_code: u32 = 255;
let collect_timeout = tokio::time::timeout(
std::time::Duration::from_secs(30),
async {
while let Some(event) = exec_rx.recv().await {
match event {
ExecEvent::Data(data) => output.extend_from_slice(&data),
ExecEvent::ExitStatus(code) => exit_code = code,
ExecEvent::Eof | ExecEvent::Close => break,
ExecEvent::Error(e) => {
log::warn!("Exec reader error: {}", e);
break;
}
_ => {}
}
}
},
).await;
if collect_timeout.is_err() {
log::warn!("Exec command timed out after 30s");
}
reader_handle.abort();
let close_msg = Message::Channel(ChannelMessage::Close { channel_id });
let _ = transport.send_message(&close_msg).await;
Ok((String::from_utf8_lossy(&output).into_owned(), exit_code))
}
pub async fn exec(&self, command: &str) -> Result<String> {
let (output, _exit_code) = self.exec_with_status(command).await?;
Ok(output)
}
pub async fn disconnect(&mut self) -> Result<()> {
if let Some(task) = self.rekey_task.take() {
task.abort();
}
if let Some(transport) = &self.transport {
let disconnect_msg = Message::Disconnect(crate::transport::DisconnectMessage {
reason_code: crate::transport::protocol::disconnect_reasons::BY_APPLICATION,
description: "Client disconnecting".into(),
});
transport.send_message(&disconnect_msg).await?;
transport.close().await?;
}
self.transport = None;
Ok(())
}
#[allow(dead_code)]
fn start_transport_handler(&self) -> Result<()> {
let transport = self.transport.as_ref()
.ok_or_else(|| QsshError::Protocol("Transport not available".into()))?
.clone();
let channel_manager = self.channel_manager.clone();
tokio::spawn(async move {
loop {
match transport.receive_message::<Message>().await {
Ok(msg) => {
if let Err(e) = handle_message(msg, &channel_manager).await {
log::error!("Error handling message: {}", e);
break;
}
}
Err(e) => {
log::error!("Transport error: {}", e);
break;
}
}
}
});
Ok(())
}
async fn setup_port_forward(&self, forward: PortForward) -> Result<()> {
let channel_manager = self.channel_manager.clone();
tokio::spawn(async move {
let mut port_forward = crate::transport::channel::PortForward::new(
forward.local_port,
forward.remote_host,
forward.remote_port,
);
if let Err(e) = port_forward.start(channel_manager).await {
log::error!("Port forward failed: {}", e);
}
});
Ok(())
}
async fn handle_shell_session(&self, channel_id: u32) -> Result<()> {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let transport = self.transport.as_ref()
.ok_or_else(|| QsshError::Protocol("Transport not available".into()))?
.clone();
log::info!("Shell session started (channel {})", channel_id);
let term = std::env::var("TERM").unwrap_or_else(|_| "xterm-256color".to_string());
let (cols, rows) = terminal_size().unwrap_or((80, 24));
let pty_req = Message::Channel(ChannelMessage::PtyRequest {
channel_id,
term: term.clone(),
width_chars: cols as u32,
height_chars: rows as u32,
width_pixels: 0,
height_pixels: 0,
modes: vec![], });
transport.send_message(&pty_req).await?;
let shell_req = Message::Channel(ChannelMessage::ShellRequest { channel_id });
transport.send_message(&shell_req).await?;
tokio::time::sleep(tokio::time::Duration::from_millis(200)).await;
let stdin = tokio::io::stdin();
let mut stdout = tokio::io::stdout();
let is_tty = atty::is(atty::Stream::Stdin);
log::debug!("Terminal is TTY: {}", is_tty);
let term_settings = if is_tty {
match termios::Termios::from_fd(0) {
Ok(settings) => {
log::debug!("Setting terminal to raw mode");
let mut raw = settings;
termios::cfmakeraw(&mut raw);
match termios::tcsetattr(0, termios::TCSANOW, &raw) {
Ok(_) => {
Some(settings)
},
Err(e) => {
log::warn!("Could not set raw mode: {}", e);
None
}
}
}
Err(e) => {
log::warn!("Could not get terminal settings: {}", e);
None
}
}
} else {
log::debug!("Not a TTY, skipping terminal raw mode");
None
};
let mut stdin = stdin;
let mut buffer = vec![0u8; 1024];
let transport_read = transport.clone();
let transport_write = transport.clone();
let registry = self.remote_forward_registry.clone();
let router = self.forwarded_channel_router.clone();
let (tx, mut rx) = tokio::sync::mpsc::channel::<Vec<u8>>(256);
let read_task = tokio::spawn(async move {
loop {
match transport_read.receive_message::<Message>().await {
Ok(Message::Channel(ChannelMessage::Data { channel_id: ch_id, data })) if ch_id == channel_id => {
if data.len() == 1 && data[0] == 0 {
continue;
}
let preview = String::from_utf8_lossy(&data[..std::cmp::min(50, data.len())]);
log::debug!("Received {} bytes of data from server: {:?}", data.len(), preview);
if tx.send(data).await.is_err() {
break;
}
}
Ok(Message::Channel(ChannelMessage::Data { channel_id: ch_id, data })) => {
if !router.route_data(ch_id, data).await {
log::debug!("No handler for channel {} data, ignoring", ch_id);
}
}
Ok(Message::Channel(ChannelMessage::Open {
channel_id: ch_id,
channel_type: ChannelType::ForwardedTcpip { connected_host, connected_port, .. },
..
})) => {
log::info!("Server opened ForwardedTcpip channel {} for {}:{}",
ch_id, connected_host, connected_port);
let transport_fwd = transport_read.clone();
let registry_fwd = registry.clone();
let router_fwd = router.clone();
tokio::spawn(async move {
if let Err(e) = handle_forwarded_channel(
ch_id,
connected_host.clone(),
connected_port,
Arc::new(transport_fwd),
registry_fwd,
router_fwd,
).await {
log::error!("Forwarded channel {} error: {}", ch_id, e);
}
});
}
Ok(Message::Channel(ChannelMessage::Open {
channel_id: ch_id,
channel_type: ChannelType::AgentForward,
..
})) => {
log::info!("Server opened AgentForward channel {}", ch_id);
let transport_agent = transport_read.clone();
let router_agent = router.clone();
tokio::spawn(async move {
if let Err(e) = handle_agent_forward_channel(ch_id, Arc::new(transport_agent), router_agent).await {
log::error!("Agent forward channel {} error: {}", ch_id, e);
}
});
}
Ok(Message::Channel(ChannelMessage::Accept { channel_id: ch_id, .. })) if ch_id != channel_id => {
log::debug!("Routing Accept for channel {} via router", ch_id);
router.route_data(ch_id, Vec::new()).await;
}
Ok(Message::Channel(ChannelMessage::Close { channel_id: ch_id })) if ch_id == channel_id => {
log::info!("Server closed channel {}", ch_id);
break;
}
Ok(Message::Channel(ChannelMessage::Close { channel_id: ch_id })) => {
log::debug!("Channel {} closed", ch_id);
router.remove(ch_id).await;
}
Ok(Message::Channel(ChannelMessage::Eof { channel_id: ch_id })) if ch_id == channel_id => {
log::info!("Received EOF for channel {}", ch_id);
break;
}
Ok(Message::Channel(ChannelMessage::Eof { channel_id: ch_id })) => {
log::debug!("Channel {} EOF", ch_id);
router.remove(ch_id).await;
}
Ok(msg) => {
log::debug!("Received other message: {:?}", msg);
}
Err(e) => {
log::error!("Transport receive error: {}", e);
break;
}
}
}
});
let result = if is_tty {
let mut sigwinch = tokio::signal::unix::signal(
tokio::signal::unix::SignalKind::window_change()
).map_err(QsshError::Io)?;
loop {
tokio::select! {
result = stdin.read(&mut buffer) => {
match result {
Ok(0) => {
break Ok(())
}, Ok(n) => {
log::debug!("Sending {} bytes to server", n);
let data_msg = Message::Channel(ChannelMessage::Data {
channel_id,
data: buffer[..n].to_vec(),
});
match transport_write.send_message(&data_msg).await {
Ok(_) => {
}
Err(e) => {
break Err(e);
}
}
}
Err(e) => {
log::error!("Stdin error: {}", e);
break Err(e.into());
}
}
}
Some(data) = rx.recv() => {
if let Err(e) = stdout.write_all(&data).await {
log::error!("Stdout error: {}", e);
break Err(e.into());
}
if let Err(e) = stdout.flush().await {
log::error!("Stdout flush error: {}", e);
break Err(e.into());
}
}
_ = sigwinch.recv() => {
if let Some((cols, rows)) = terminal_size() {
log::debug!("Terminal resized to {}x{}", cols, rows);
let resize_msg = Message::Channel(ChannelMessage::WindowChange {
channel_id,
width_chars: cols as u32,
height_chars: rows as u32,
width_pixels: 0,
height_pixels: 0,
});
let _ = transport_write.send_message(&resize_msg).await;
}
}
}
}
} else {
log::debug!("Non-TTY mode: handling interactive session");
loop {
tokio::select! {
result = stdin.read(&mut buffer) => {
match result {
Ok(0) => {
break Ok(())
}, Ok(n) => {
log::debug!("Sending {} bytes to server", n);
let data_msg = Message::Channel(ChannelMessage::Data {
channel_id,
data: buffer[..n].to_vec(),
});
if let Err(e) = transport_write.send_message(&data_msg).await {
break Err(e);
}
}
Err(e) => {
log::error!("Stdin error: {}", e);
break Err(e.into());
}
}
}
Some(data) = rx.recv() => {
if let Err(e) = stdout.write_all(&data).await {
log::error!("Stdout error: {}", e);
break Err(e.into());
}
if let Err(e) = stdout.flush().await {
log::error!("Stdout flush error: {}", e);
break Err(e.into());
}
}
}
}
};
read_task.abort();
if let Some(settings) = term_settings {
termios::tcsetattr(0, termios::TCSANOW, &settings)
.map_err(|e| QsshError::Io(std::io::Error::other(e)))?;
}
let close_msg = Message::Channel(ChannelMessage::Close { channel_id });
let _ = transport.send_message(&close_msg).await;
result
}
}
#[allow(dead_code)]
async fn handle_message(msg: Message, channel_manager: &ChannelManager) -> Result<()> {
match msg {
Message::Channel(channel_msg) => {
channel_manager.handle_message(channel_msg).await?;
}
Message::Disconnect(d) => {
log::info!("Server disconnected: {}", d.description);
return Err(QsshError::Connection("Server disconnected".into()));
}
Message::Ping(nonce) => {
log::debug!("Received ping: {}", nonce);
}
_ => {
log::debug!("Unhandled message type");
}
}
Ok(())
}
async fn rekey_with_transport(transport: &Transport) -> Result<()> {
use sha3::Sha3_256;
use hkdf::Hkdf;
log::info!("Initiating rekey (rekey #{})", transport.rekey_count() + 1);
let client_kex = PqKeyExchange::new()?;
let (client_share, client_sig) = client_kex.create_key_share()?;
let rekey_msg = RekeyMessage {
new_falcon_public_key: client_kex.falcon_pk.clone(),
new_key_share: client_share.clone(),
new_key_share_signature: client_sig,
request_qkd: false,
};
transport.send_message(&Message::Rekey(rekey_msg)).await?;
let server_rekey = tokio::time::timeout(Duration::from_secs(30), async {
loop {
match transport.receive_message::<Message>().await {
Ok(Message::Rekey(rekey)) => return Ok(rekey),
Ok(_) => continue,
Err(e) => return Err(e),
}
}
}).await
.map_err(|_| QsshError::Protocol("Rekey response timed out".into()))??;
let verified = client_kex.verify_falcon(
&server_rekey.new_key_share,
&server_rekey.new_key_share_signature,
&server_rekey.new_falcon_public_key,
)?;
if !verified {
return Err(QsshError::Crypto("Server rekey signature verification failed".into()));
}
let mut ikm = Vec::new();
ikm.extend_from_slice(&client_share);
ikm.extend_from_slice(&server_rekey.new_key_share);
let hk = Hkdf::<Sha3_256>::new(None, &ikm);
let mut server_write_key = vec![0u8; 32];
let mut client_write_key = vec![0u8; 32];
hk.expand(b"qssh-rekey-server-write", &mut server_write_key)
.map_err(|_| QsshError::Crypto("HKDF expand failed".into()))?;
hk.expand(b"qssh-rekey-client-write", &mut client_write_key)
.map_err(|_| QsshError::Crypto("HKDF expand failed".into()))?;
let new_send = SymmetricCrypto::from_shared_secret(&client_write_key)?;
let new_recv = SymmetricCrypto::from_shared_secret(&server_write_key)?;
transport.update_keys(new_send, new_recv).await?;
log::info!("Automatic rekey complete — new session keys active");
Ok(())
}
fn terminal_size() -> Option<(u16, u16)> {
unsafe {
let mut ws: libc::winsize = std::mem::zeroed();
if libc::ioctl(0, libc::TIOCGWINSZ, &mut ws) == 0 && ws.ws_col > 0 && ws.ws_row > 0 {
Some((ws.ws_col, ws.ws_row))
} else {
None
}
}
}
async fn handle_agent_forward_channel(
channel_id: u32,
transport: Arc<Transport>,
router: ForwardedChannelRouter,
) -> Result<()> {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::UnixStream;
let auth_sock = std::env::var("QSSH_AUTH_SOCK")
.map_err(|_| QsshError::Connection("QSSH_AUTH_SOCK not set".into()))?;
let accept = Message::Channel(ChannelMessage::Accept {
channel_id,
sender_channel: channel_id,
window_size: 1048576,
max_packet_size: 32768,
});
transport.send_message(&accept).await?;
let agent_stream = UnixStream::connect(&auth_sock).await
.map_err(|e| QsshError::Connection(format!("Failed to connect to local agent: {}", e)))?;
let (mut agent_read, mut agent_write) = tokio::io::split(agent_stream);
let (data_tx, mut data_rx) = tokio::sync::mpsc::channel::<Vec<u8>>(256);
router.register(channel_id, data_tx).await;
let transport_to_agent = transport.clone();
let chan_to_agent = tokio::spawn(async move {
while let Some(data) = data_rx.recv().await {
if agent_write.write_all(&data).await.is_err() {
break;
}
}
});
let agent_to_chan = tokio::spawn(async move {
let mut buf = vec![0u8; 32768];
loop {
match agent_read.read(&mut buf).await {
Ok(0) => break,
Ok(n) => {
let msg = Message::Channel(ChannelMessage::Data {
channel_id,
data: buf[..n].to_vec(),
});
if transport_to_agent.send_message(&msg).await.is_err() {
break;
}
}
Err(_) => break,
}
}
});
tokio::select! {
_ = chan_to_agent => {}
_ = agent_to_chan => {}
}
router.remove(channel_id).await;
let close = Message::Channel(ChannelMessage::Close { channel_id });
let _ = transport.send_message(&close).await;
log::debug!("Agent forward channel {} closed", channel_id);
Ok(())
}