use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use muxio_core::rpc::rpc_internals::RpcStreamEvent;
use muxio_rpc_service::prebuffered::RpcMethodPrebuffered;
use muxio_rpc_service_caller::prebuffered::RpcCallPrebuffered;
use muxio_rpc_service_endpoint::{RpcServiceEndpointInterface, StreamResponder};
use muxio_tokio_rpc_ipc_server::{RpcIpcConnectionContextHandle, RpcIpcServer, RpcIpcServerEvent};
use portable_pty::PtySize;
use tokio::sync::{Mutex, Notify, RwLock, mpsc, oneshot};
use term_session_muxio_service_definitions::{
Attach, ChannelInfo, ChannelName, ClientInfo, CloseSession, KillChannel, KillClient,
ListChannels, ListChannelsResponse, OnPtyResized, RPC_ERROR_SHUTTING_DOWN,
RPC_ERROR_UNATTACHED, ResizePty, STREAM_INPUT_METHOD_ID, SUBSCRIBE_OUTPUT_METHOD_ID,
SessionInfo, ShutdownGateway, Spawn, WriteInput,
};
use term_wm_pty_engine::PtyStatus;
use crate::session::Session;
const SESSION_ID: u64 = 1;
const INPUT_CHANNEL_CAPACITY: usize = 128;
const SESSION_EXIT_FLUSH_GRACE: std::time::Duration = std::time::Duration::from_millis(100);
const SESSION_EXIT_POLL_INTERVAL: std::time::Duration = std::time::Duration::from_millis(100);
const SIGKILL_GRACE: std::time::Duration = std::time::Duration::from_millis(500);
#[cfg(unix)]
const SIGTERM: i32 = libc::SIGTERM;
#[cfg(not(unix))]
const SIGTERM: i32 = 15;
#[cfg(unix)]
const SIGKILL: i32 = libc::SIGKILL;
#[cfg(not(unix))]
#[allow(dead_code)]
const SIGKILL: i32 = 9;
const SHUTDOWN_FLUSH_GRACE_MS: u64 = 50;
#[derive(Clone)]
enum ConnState {
Unattached,
Attached(ChannelName),
}
#[derive(Clone)]
struct ConnEntry {
handle: RpcIpcConnectionContextHandle,
state: ConnState,
hostname: String,
connected_at_unix: u64,
pid: u64,
}
#[derive(Clone)]
struct ClientEntry {
caller: Option<RpcIpcConnectionContextHandle>,
hostname: String,
connected_at_unix: u64,
pid: u64,
cols: u16,
rows: u16,
}
struct SubscriberEntry {
conn_id: usize,
respond: StreamResponder,
}
struct ChannelState {
session: Option<Session>,
clients: HashMap<usize, ClientEntry>,
subscribers: Vec<SubscriberEntry>,
notify: Arc<Notify>,
created_at_unix: u64,
cmd: Vec<String>,
input_tx: mpsc::Sender<Vec<u8>>,
kill_pending: bool,
is_reaped: bool,
}
struct ServerState {
conns: RwLock<HashMap<usize, ConnEntry>>,
channels: RwLock<HashMap<ChannelName, Arc<Mutex<ChannelState>>>>,
is_shutting_down: AtomicBool,
}
type SharedState = Arc<ServerState>;
fn rpc_err(message: &str) -> Box<dyn std::error::Error + Send + Sync> {
Box::new(std::io::Error::other(message.to_string()))
}
fn boxed_io(e: std::io::Error) -> Box<dyn std::error::Error + Send + Sync> {
Box::new(e)
}
fn now_unix() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0)
}
impl ChannelState {
fn new(cmd: Vec<String>, input_tx: mpsc::Sender<Vec<u8>>, notify: Arc<Notify>) -> Self {
Self {
session: None,
clients: HashMap::new(),
subscribers: Vec::new(),
notify,
created_at_unix: now_unix(),
cmd,
input_tx,
kill_pending: false,
is_reaped: false,
}
}
fn set_session(&mut self, mut session: Session) {
let n = self.notify.clone();
session.set_status_callback(Some(Box::new(move |status| {
if matches!(status, PtyStatus::Wakeup | PtyStatus::Exited) {
n.notify_one();
}
})));
self.session = Some(session);
self.notify.notify_one();
}
fn request_session_kill(&mut self, signal: i32) {
let _ = &signal;
if let Some(session) = self.session.as_mut() {
#[cfg(unix)]
let _ = session.pty.signal_process_group(signal);
#[cfg(not(unix))]
let _ = session.pty.kill_child();
}
self.kill_pending = true;
self.notify.notify_one();
}
fn finalize_subscribers(&mut self) {
if let Some(session) = self.session.as_mut() {
let raw = session.read_output();
if !raw.is_empty() {
for sub in &self.subscribers {
sub.respond.respond(raw.clone(), false);
}
}
}
for sub in &self.subscribers {
sub.respond.respond(Vec::new(), true);
}
self.subscribers.clear();
}
fn recalculate_pty_size(&mut self) {
let Some(session) = self.session.as_mut() else {
return;
};
let Some(min_cols) = self
.clients
.values()
.map(|c| c.cols)
.filter(|&c| c != u16::MAX)
.min()
else {
return;
};
let Some(min_rows) = self
.clients
.values()
.map(|c| c.rows)
.filter(|&r| r != u16::MAX)
.min()
else {
return;
};
let size = PtySize {
rows: min_rows,
cols: min_cols,
pixel_width: 0,
pixel_height: 0,
};
let _ = session.pty.resize(size);
session.cols = min_cols;
session.rows = min_rows;
}
fn notify_clients(&self, clients: &[ClientEntry], cols: u16, rows: u16) {
for client in clients {
let Some(caller) = client.caller.clone() else {
continue;
};
tokio::spawn(async move {
if let Err(e) = OnPtyResized::call(&caller, (cols, rows)).await {
tracing::debug!(error = ?e, "Failed to deliver OnPtyResized notification");
}
});
}
}
fn to_info(&self, name: &ChannelName) -> ChannelInfo {
let session = self.session.as_ref().map(|s| SessionInfo {
id: s.id,
cols: s.cols,
rows: s.rows,
exited: s.exited,
exit_code: s.exit_code,
title: s.title.clone().unwrap_or_default(),
});
let clients = self
.clients
.iter()
.map(|(conn_id, c)| ClientInfo {
conn_id: *conn_id,
pid: c.pid,
hostname: c.hostname.clone(),
connected_at_unix: c.connected_at_unix,
cols: c.cols,
rows: c.rows,
})
.collect();
ChannelInfo {
name: name.to_string(),
created_at_unix: self.created_at_unix,
session,
clients,
}
}
}
async fn bound_channel(state: &ServerState, conn_id: usize) -> Option<ChannelName> {
let conns = state.conns.read().await;
match conns.get(&conn_id)?.state {
ConnState::Attached(ref name) => Some(name.clone()),
ConnState::Unattached => None,
}
}
async fn resolve_channel(
state: &ServerState,
name: &ChannelName,
) -> Option<Arc<Mutex<ChannelState>>> {
let channels = state.channels.read().await;
channels.get(name).cloned()
}
async fn get_or_create_channel(
state: &SharedState,
name: &ChannelName,
) -> Arc<Mutex<ChannelState>> {
{
let channels = state.channels.read().await;
if let Some(existing) = channels.get(name) {
let arc = existing.clone();
drop(channels);
let is_reaped = arc.lock().await.is_reaped;
if !is_reaped {
return arc;
}
}
}
let mut channels = state.channels.write().await;
if let Some(existing) = channels.get(name) {
let arc = existing.clone();
let is_reaped = arc.lock().await.is_reaped;
if !is_reaped {
return arc;
}
}
let (input_tx, input_rx) = mpsc::channel::<Vec<u8>>(INPUT_CHANNEL_CAPACITY);
let notify = Arc::new(Notify::new());
let channel = Arc::new(Mutex::new(ChannelState::new(Vec::new(), input_tx, notify)));
let ch = Arc::clone(&channel);
tokio::spawn(async move {
let mut input_rx = input_rx;
while let Some(data) = input_rx.recv().await {
let writer = {
let guard = ch.lock().await;
guard.session.as_ref().map(|s| s.pty.writer_handle())
};
if let Some(writer) = writer {
let _ = tokio::task::spawn_blocking(move || writer.write_bytes(&data)).await;
}
}
});
{
let st = Arc::clone(state);
let ch = Arc::clone(&channel);
let notify = {
let locked = ch.lock().await;
locked.notify.clone()
};
let name_for_task = name.clone();
tokio::spawn(async move {
loop {
tokio::select! {
_ = notify.notified() => {}
_ = tokio::time::sleep(SESSION_EXIT_POLL_INTERVAL) => {}
}
let mut guard = ch.lock().await;
if guard.is_reaped {
break;
}
if guard.subscribers.is_empty() {
if let Some(session) = guard.session.as_mut() {
session.sync_screen();
if session.check_exited() {
tracing::info!(channel = %name_for_task, "Session exited");
guard.session = None;
guard.kill_pending = false;
}
}
} else {
let (raw, exited, code) = {
let Some(session) = guard.session.as_mut() else {
for sub in &guard.subscribers {
sub.respond.respond(Vec::new(), true);
}
guard.subscribers.clear();
guard.notify.notify_one();
continue;
};
let raw = session.read_output();
let exited = session.check_exited();
let code = session.exit_code;
(raw, exited, code)
};
if !raw.is_empty() {
for sub in &guard.subscribers {
sub.respond.respond(raw.clone(), false);
}
}
if exited {
tracing::info!(channel = %name_for_task, "Session exited with code {:?}", code);
for sub in &guard.subscribers {
sub.respond.respond(Vec::new(), true);
}
guard.subscribers.clear();
guard.session = None;
guard.kill_pending = false;
guard.notify.notify_one();
}
}
let should_reap = guard.session.is_none() && guard.clients.is_empty();
drop(guard);
if should_reap {
let mut channels = st.channels.write().await;
if let Some(arc) = channels.get(&name_for_task) {
let mut locked = arc.lock().await;
if locked.session.is_none() && locked.clients.is_empty() {
locked.is_reaped = true;
drop(locked);
channels.remove(&name_for_task);
tracing::info!(channel = %name_for_task, "Reaped idle channel");
}
}
}
}
});
}
channels.insert(name.clone(), Arc::clone(&channel));
channel
}
async fn evict_conn(state: &ServerState, conn_id: usize) {
let channel = {
let mut conns = state.conns.write().await;
let entry = conns.remove(&conn_id);
entry.and_then(|e| match e.state {
ConnState::Attached(name) => Some(name),
ConnState::Unattached => None,
})
};
let Some(channel) = channel else {
return;
};
let Some(ch) = resolve_channel(state, &channel).await else {
return;
};
let mut guard = ch.lock().await;
guard.clients.remove(&conn_id);
guard.subscribers.retain(|s| s.conn_id != conn_id);
guard.recalculate_pty_size();
let session_size = guard.session.as_ref().map(|s| (s.cols, s.rows));
let targets: Vec<ClientEntry> = guard.clients.values().cloned().collect();
drop(guard);
let Some((ncols, nrows)) = session_size else {
return;
};
if let Some(ch) = resolve_channel(state, &channel).await {
let guard = ch.lock().await;
guard.notify_clients(&targets, ncols, nrows);
}
}
async fn spawn_kill_escalation(
state: &SharedState,
name: &ChannelName,
) -> tokio::task::JoinHandle<()> {
let state = Arc::clone(state);
let name = name.clone();
tokio::spawn(async move {
tokio::time::sleep(SIGKILL_GRACE).await;
let Some(ch) = resolve_channel(&state, &name).await else {
return;
};
let mut guard = ch.lock().await;
if !guard.kill_pending {
return;
}
let alive = guard
.session
.as_ref()
.is_some_and(|s| !s.exited && s.pty.reader_is_alive());
guard.kill_pending = false;
if !alive {
return;
}
if let Some(session) = guard.session.as_mut() {
#[cfg(unix)]
let _ = session.pty.signal_process_group(SIGKILL);
#[cfg(not(unix))]
let _ = session.pty.kill_child();
}
})
}
pub async fn run_gateway(
gateway: ChannelName,
) -> Result<i32, Box<dyn std::error::Error + Send + Sync>> {
let socket_name = gateway.to_string();
let state: SharedState = Arc::new(ServerState {
conns: RwLock::new(HashMap::new()),
channels: RwLock::new(HashMap::new()),
is_shutting_down: AtomicBool::new(false),
});
let (event_tx, mut event_rx) = mpsc::unbounded_channel();
let server = RpcIpcServer::new(Some(event_tx));
let endpoint = server.endpoint();
let st = Arc::clone(&state);
endpoint
.register_prebuffered(Attach::METHOD_ID, move |payload, ctx| {
let state = Arc::clone(&st);
async move {
if state.is_shutting_down.load(Ordering::SeqCst) {
return Err(rpc_err(RPC_ERROR_SHUTTING_DOWN));
}
let (channel_str, hostname, pid) = Attach::decode_request(&payload)?;
let name = ChannelName::parse(&channel_str).map_err(|e| rpc_err(&e))?;
let _channel = get_or_create_channel(&state, &name).await;
let conn_id = ctx.conn_id;
let mut conns = state.conns.write().await;
let entry = conns.entry(conn_id).or_insert_with(|| ConnEntry {
handle: RpcIpcConnectionContextHandle(ctx.clone()),
state: ConnState::Unattached,
hostname: String::new(),
connected_at_unix: now_unix(),
pid: 0,
});
entry.state = ConnState::Attached(name);
entry.hostname = hostname;
entry.connected_at_unix = now_unix();
entry.pid = pid;
Attach::encode_response(conn_id).map_err(boxed_io)
}
})
.await
.map_err(|e| format!("register Attach: {e:?}"))?;
let st = Arc::clone(&state);
endpoint
.register_prebuffered(Spawn::METHOD_ID, move |payload, ctx| {
let state = Arc::clone(&st);
async move {
if state.is_shutting_down.load(Ordering::SeqCst) {
return Err(rpc_err(RPC_ERROR_SHUTTING_DOWN));
}
let (cmd, cols, rows) = Spawn::decode_request(&payload)?;
let channel = bound_channel(state.as_ref(), ctx.conn_id).await;
let Some(channel) = channel else {
return Err(rpc_err(RPC_ERROR_UNATTACHED));
};
let ch = get_or_create_channel(&state, &channel).await;
let conn_meta = {
let conns = state.conns.read().await;
conns.get(&ctx.conn_id).map(|c| {
(
c.handle.clone(),
c.hostname.clone(),
c.connected_at_unix,
c.pid,
)
})
};
let mut guard = ch.lock().await;
let entry = guard
.clients
.entry(ctx.conn_id)
.or_insert_with(|| ClientEntry {
caller: conn_meta.as_ref().map(|(h, _, _, _)| h.clone()),
hostname: conn_meta
.as_ref()
.map(|(_, n, _, _)| n.clone())
.unwrap_or_default(),
connected_at_unix: conn_meta.as_ref().map(|(_, _, t, _)| *t).unwrap_or(0),
pid: conn_meta.as_ref().map(|(_, _, _, p)| *p).unwrap_or(0),
cols,
rows,
});
entry.cols = cols;
entry.rows = rows;
if guard.session.as_ref().is_some_and(|s| !s.exited) {
guard.recalculate_pty_size();
let session = guard.session.as_ref().unwrap();
let (ncols, nrows) = (session.cols, session.rows);
let targets: Vec<ClientEntry> = guard.clients.values().cloned().collect();
let id = session.id;
let cols = session.cols;
let rows = session.rows;
drop(guard);
if let Some(ch) = resolve_channel(state.as_ref(), &channel).await {
let g = ch.lock().await;
g.notify_clients(&targets, ncols, nrows);
}
return Spawn::encode_response((id, cols, rows)).map_err(boxed_io);
}
let effective_cmd = if let Some(c) = cmd
&& !c.is_empty()
{
guard.cmd = c.clone();
Some(c)
} else if !guard.cmd.is_empty() {
Some(guard.cmd.clone())
} else {
None
};
let id = SESSION_ID;
let session = Session::spawn(id, effective_cmd, cols, rows, Some(&channel))?;
guard.set_session(session);
guard.recalculate_pty_size();
let targets: Vec<ClientEntry> = guard.clients.values().cloned().collect();
let session = guard.session.as_ref().unwrap();
let (sid, scol, srow) = (session.id, session.cols, session.rows);
let (ncols, nrows) = (scol, srow);
drop(guard);
if let Some(ch) = resolve_channel(state.as_ref(), &channel).await {
let g = ch.lock().await;
g.notify_clients(&targets, ncols, nrows);
}
Spawn::encode_response((sid, scol, srow)).map_err(boxed_io)
}
})
.await
.map_err(|e| format!("register Spawn: {e:?}"))?;
let st = Arc::clone(&state);
endpoint
.register_prebuffered(ResizePty::METHOD_ID, move |payload, ctx| {
let state = Arc::clone(&st);
async move {
if state.is_shutting_down.load(Ordering::SeqCst) {
return Err(rpc_err(RPC_ERROR_SHUTTING_DOWN));
}
let (_id, cols, rows) = ResizePty::decode_request(&payload)?;
let channel = bound_channel(state.as_ref(), ctx.conn_id).await;
let Some(channel) = channel else {
return Err(rpc_err(RPC_ERROR_UNATTACHED));
};
let ch = resolve_channel(state.as_ref(), &channel)
.await
.ok_or_else(|| rpc_err("channel not found"))?;
let mut guard = ch.lock().await;
if let Some(client) = guard.clients.get_mut(&ctx.conn_id) {
client.cols = cols;
client.rows = rows;
}
guard.recalculate_pty_size();
let (ncols, nrows) = guard
.session
.as_ref()
.map(|s| (s.cols, s.rows))
.unwrap_or((cols, rows));
let targets: Vec<ClientEntry> = guard.clients.values().cloned().collect();
drop(guard);
if let Some(ch) = resolve_channel(state.as_ref(), &channel).await {
let g = ch.lock().await;
g.notify_clients(&targets, ncols, nrows);
}
ResizePty::encode_response((ncols, nrows)).map_err(boxed_io)
}
})
.await
.map_err(|e| format!("register ResizePty: {e:?}"))?;
let st = Arc::clone(&state);
endpoint
.register_prebuffered(CloseSession::METHOD_ID, move |payload, ctx| {
let state = Arc::clone(&st);
async move {
if state.is_shutting_down.load(Ordering::SeqCst) {
return Err(rpc_err(RPC_ERROR_SHUTTING_DOWN));
}
let _id = CloseSession::decode_request(&payload)?;
let channel = bound_channel(state.as_ref(), ctx.conn_id).await;
let Some(channel) = channel else {
return Err(rpc_err(RPC_ERROR_UNATTACHED));
};
let ch = resolve_channel(state.as_ref(), &channel)
.await
.ok_or_else(|| rpc_err("channel not found"))?;
let mut guard = ch.lock().await;
guard.request_session_kill(SIGTERM);
guard.finalize_subscribers();
drop(guard);
spawn_kill_escalation(&state, &channel).await;
CloseSession::encode_response(()).map_err(boxed_io)
}
})
.await
.map_err(|e| format!("register CloseSession: {e:?}"))?;
let st = Arc::clone(&state);
endpoint
.register_prebuffered(WriteInput::METHOD_ID, move |payload, ctx| {
let state = Arc::clone(&st);
async move {
if state.is_shutting_down.load(Ordering::SeqCst) {
return Err(rpc_err(RPC_ERROR_SHUTTING_DOWN));
}
let (id, data) = WriteInput::decode_request(&payload)?;
let channel = bound_channel(state.as_ref(), ctx.conn_id).await;
let Some(channel) = channel else {
return Err(rpc_err(RPC_ERROR_UNATTACHED));
};
let ch = resolve_channel(state.as_ref(), &channel)
.await
.ok_or_else(|| rpc_err("channel not found"))?;
let writer = {
let guard = ch.lock().await;
guard
.session
.as_ref()
.filter(|s| s.id == id)
.map(|s| s.pty.writer_handle())
};
if let Some(writer) = writer {
let _ = tokio::task::spawn_blocking(move || writer.write_bytes(&data)).await;
}
WriteInput::encode_response(()).map_err(boxed_io)
}
})
.await
.map_err(|e| format!("register WriteInput: {e:?}"))?;
let st = Arc::clone(&state);
endpoint
.register_stream_handler(STREAM_INPUT_METHOD_ID, move |event, _responder, ctx| {
if let RpcStreamEvent::PayloadChunk { bytes, .. } = event {
let state = Arc::clone(&st);
let conn_id = ctx.conn_id;
tokio::spawn(async move {
let channel = bound_channel(state.as_ref(), conn_id).await;
let Some(channel) = channel else {
return;
};
let ch = resolve_channel(state.as_ref(), &channel).await;
let Some(ch) = ch else {
return;
};
let tx = {
let guard = ch.lock().await;
guard.input_tx.clone()
};
if let Err(e) = tx.try_send(bytes) {
tracing::warn!(error = %e, "gateway input buffer full; dropping input chunk");
}
});
}
})
.await
.map_err(|e| format!("register stream handler STREAM_INPUT: {e:?}"))?;
let st = Arc::clone(&state);
endpoint
.register_stream_handler(SUBSCRIBE_OUTPUT_METHOD_ID, move |event, respond, ctx| {
let is_new = matches!(&event, RpcStreamEvent::Header { .. });
if is_new {
let st = Arc::clone(&st);
let conn_id = ctx.conn_id;
tokio::spawn(async move {
let channel = bound_channel(&st, conn_id).await;
let Some(channel) = channel else {
return;
};
let ch = resolve_channel(&st, &channel).await;
let Some(ch) = ch else {
return;
};
let mut guard = ch.lock().await;
let early = guard.session.as_mut().and_then(|s| {
let data = s.read_output();
if data.is_empty() { None } else { Some(data) }
});
let snapshot = guard.session.as_mut().map(|s| s.generate_snapshot());
guard.subscribers.push(SubscriberEntry {
conn_id,
respond: respond.clone(),
});
guard.notify.notify_one();
let is_dead = guard.session.is_none();
drop(guard);
if let Some(data) = snapshot
&& !data.is_empty()
{
respond.respond(data, false);
}
if let Some(data) = early {
respond.respond(data, false);
}
if is_dead {
respond.respond(Vec::new(), true);
}
});
}
})
.await
.map_err(|e| format!("register SubscribeOutput: {e:?}"))?;
let st = Arc::clone(&state);
let list_socket = socket_name.clone();
endpoint
.register_prebuffered(ListChannels::METHOD_ID, move |_payload, _ctx| {
let state = Arc::clone(&st);
let socket = list_socket.clone();
async move {
let channels = {
let chans = state.channels.read().await;
let mut v: Vec<(ChannelName, Arc<Mutex<ChannelState>>)> =
chans.iter().map(|(k, v)| (k.clone(), v.clone())).collect();
drop(chans);
v.sort_by_key(|a| a.0.to_string());
v
};
let mut out = Vec::with_capacity(channels.len());
for (name, ch) in channels {
let guard = ch.lock().await;
out.push(guard.to_info(&name));
}
ListChannels::encode_response(ListChannelsResponse {
gateway_pid: std::process::id() as u64,
socket,
channels: out,
})
.map_err(boxed_io)
}
})
.await
.map_err(|e| format!("register ListChannels: {e:?}"))?;
let st = Arc::clone(&state);
endpoint
.register_prebuffered(KillChannel::METHOD_ID, move |payload, _ctx| {
let state = Arc::clone(&st);
async move {
if state.is_shutting_down.load(Ordering::SeqCst) {
return Err(rpc_err(RPC_ERROR_SHUTTING_DOWN));
}
let channel_str = KillChannel::decode_request(&payload)?;
let name = ChannelName::parse(&channel_str).map_err(|e| rpc_err(&e))?;
let target_conns: Vec<usize> = {
let conns = state.conns.read().await;
conns
.iter()
.filter(|(_, entry)| matches!(entry.state, ConnState::Attached(ref n) if n == &name))
.map(|(conn_id, _)| *conn_id)
.collect()
};
if let Some(ch) = resolve_channel(state.as_ref(), &name).await {
let mut guard = ch.lock().await;
guard.request_session_kill(SIGTERM);
guard.finalize_subscribers();
for conn_id in &target_conns {
guard.clients.remove(conn_id);
guard.subscribers.retain(|s| s.conn_id != *conn_id);
}
drop(guard);
spawn_kill_escalation(&state, &name).await;
}
let mut conns = state.conns.write().await;
for conn_id in &target_conns {
conns.remove(conn_id);
}
KillChannel::encode_response(()).map_err(boxed_io)
}
})
.await
.map_err(|e| format!("register KillChannel: {e:?}"))?;
let st = Arc::clone(&state);
endpoint
.register_prebuffered(KillClient::METHOD_ID, move |payload, _ctx| {
let state = Arc::clone(&st);
async move {
if state.is_shutting_down.load(Ordering::SeqCst) {
return Err(rpc_err(RPC_ERROR_SHUTTING_DOWN));
}
let (channel_str, conn_id) = KillClient::decode_request(&payload)?;
let name = ChannelName::parse(&channel_str).map_err(|e| rpc_err(&e))?;
let bound_channel = {
let conns = state.conns.read().await;
conns.get(&conn_id).and_then(|c| match &c.state {
ConnState::Attached(n) if n == &name => Some(n.clone()),
_ => None,
})
};
let Some(bound) = bound_channel else {
return Err(rpc_err(&format!(
"client {conn_id} is not attached to channel '{name}'"
)));
};
{
let mut conns = state.conns.write().await;
conns.remove(&conn_id);
}
if let Some(ch) = resolve_channel(state.as_ref(), &bound).await {
let mut guard = ch.lock().await;
let mut evicted: Vec<StreamResponder> = Vec::new();
let mut keep = Vec::with_capacity(guard.subscribers.len());
for sub in guard.subscribers.drain(..) {
if sub.conn_id == conn_id {
evicted.push(sub.respond);
} else {
keep.push(sub);
}
}
guard.subscribers = keep;
for respond in evicted {
respond.respond(Vec::new(), true);
}
guard.clients.remove(&conn_id);
guard.recalculate_pty_size();
let session_size = guard.session.as_ref().map(|s| (s.cols, s.rows));
let targets: Vec<ClientEntry> = guard.clients.values().cloned().collect();
drop(guard);
if let (Some((ncols, nrows)), Some(ch)) =
(session_size, resolve_channel(state.as_ref(), &bound).await)
{
let g = ch.lock().await;
g.notify_clients(&targets, ncols, nrows);
}
}
KillClient::encode_response(()).map_err(boxed_io)
}
})
.await
.map_err(|e| format!("register KillClient: {e:?}"))?;
let st = Arc::clone(&state);
let (shutdown_tx, mut shutdown_rx) = oneshot::channel::<()>();
let shutdown_tx = Arc::new(Mutex::new(Some(shutdown_tx)));
endpoint
.register_prebuffered(ShutdownGateway::METHOD_ID, move |_payload, _ctx| {
let state = Arc::clone(&st);
let shutdown_tx = Arc::clone(&shutdown_tx);
async move {
state.is_shutting_down.store(true, Ordering::SeqCst);
let channels: Vec<(ChannelName, Arc<Mutex<ChannelState>>)> = {
let chans = state.channels.read().await;
chans.iter().map(|(k, v)| (k.clone(), v.clone())).collect()
};
let mut escalations = Vec::new();
for (name, ch) in channels {
let mut guard = ch.lock().await;
tracing::info!(channel = %name, "Shutdown: signaling session tree");
guard.request_session_kill(SIGTERM);
guard.finalize_subscribers();
drop(guard);
escalations.push(spawn_kill_escalation(&state, &name).await);
}
tokio::spawn(async move {
for handle in escalations {
let _ = handle.await;
}
tokio::time::sleep(std::time::Duration::from_millis(SHUTDOWN_FLUSH_GRACE_MS))
.await;
let mut tx_guard = shutdown_tx.lock().await;
if let Some(tx) = tx_guard.take() {
let _ = tx.send(());
}
});
ShutdownGateway::encode_response(()).map_err(boxed_io)
}
})
.await
.map_err(|e| format!("register ShutdownGateway: {e:?}"))?;
let st = Arc::clone(&state);
tokio::spawn(async move {
while let Some(event) = event_rx.recv().await {
match event {
RpcIpcServerEvent::ClientConnected(handle) => {
tracing::info!("Client {} connected", handle.0.conn_id);
let mut conns = st.conns.write().await;
conns.entry(handle.0.conn_id).or_insert_with(|| ConnEntry {
handle: handle.clone(),
state: ConnState::Unattached,
hostname: String::new(),
connected_at_unix: now_unix(),
pid: 0,
});
}
RpcIpcServerEvent::ClientDisconnected(conn_id) => {
tracing::info!("Client {conn_id} disconnected");
evict_conn(st.as_ref(), conn_id).await;
}
}
}
});
tracing::info!("Gateway listening on channel {gateway}");
let exit_code = tokio::select! {
result = async {
server.serve(&socket_name).await.map_err(|e| format!("serve: {e:?}"))
} => {
result?;
0
}
_ = &mut shutdown_rx => {
tokio::time::sleep(SESSION_EXIT_FLUSH_GRACE).await;
0
}
};
Ok(exit_code)
}