use std::collections::HashMap;
use std::sync::Arc;
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, mpsc, oneshot};
use term_session_muxio_service_definitions::{
ChannelName, CloseSession, ListSessions, OnPtyResized, ResizePty, STREAM_INPUT_METHOD_ID,
SUBSCRIBE_OUTPUT_METHOD_ID, Spawn, WriteInput,
};
use term_wm_pty_engine::PtyStatus;
use crate::session::Session;
const FALLBACK_COLS: u16 = 80;
const FALLBACK_ROWS: u16 = 24;
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);
pub struct SessionServerConfig {
pub channel: ChannelName,
pub cmd: Vec<String>,
pub cols: u16,
pub rows: u16,
}
#[derive(Clone)]
struct ClientEntry {
caller: Option<RpcIpcConnectionContextHandle>,
cols: u16,
rows: u16,
}
struct SubscriberEntry {
conn_id: usize,
respond: StreamResponder,
}
struct ServerState {
session: Option<Session>,
clients: HashMap<usize, ClientEntry>,
subscribers: Vec<SubscriberEntry>,
notify: Arc<Notify>,
}
impl ServerState {
fn new(notify: Arc<Notify>) -> Self {
Self {
session: None,
clients: HashMap::new(),
subscribers: Vec::new(),
notify,
}
}
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 clear_session(&mut self) {
if let Some(mut session) = self.session.take() {
let _ = session.pty.kill_child();
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();
self.notify.notify_one();
}
fn recalculate_pty_size(&mut self) {
let Some(session) = self.session.as_mut() else {
return;
};
if self.clients.is_empty() {
return;
}
let min_cols = self
.clients
.values()
.map(|c| c.cols)
.filter(|&c| c != u16::MAX)
.min()
.unwrap_or(FALLBACK_COLS);
let min_rows = self
.clients
.values()
.map(|c| c.rows)
.filter(|&r| r != u16::MAX)
.min()
.unwrap_or(FALLBACK_ROWS);
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(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");
}
});
}
}
}
type SharedState = Arc<Mutex<ServerState>>;
pub async fn run_server(
config: SessionServerConfig,
) -> Result<i32, Box<dyn std::error::Error + Send + Sync>> {
let socket_name = config.channel.to_string();
let notify = Arc::new(Notify::new());
let state: SharedState = Arc::new(Mutex::new(ServerState::new(notify.clone())));
{
let mut st = state.lock().await;
let cmd = if config.cmd.is_empty() {
None
} else {
Some(config.cmd.clone())
};
let session = Session::spawn(
SESSION_ID,
cmd,
config.cols,
config.rows,
Some(&config.channel),
)?;
st.set_session(session);
}
let channel_id = config.channel.clone();
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);
let ch = channel_id.clone();
endpoint
.register_prebuffered(Spawn::METHOD_ID, move |payload, ctx| {
let state = Arc::clone(&st);
let ch = ch.clone();
async move {
let mut guard = state.lock().await;
let (cmd, cols, rows) = Spawn::decode_request(&payload)
.map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)?;
let entry = guard
.clients
.entry(ctx.conn_id)
.or_insert_with(|| ClientEntry {
caller: None,
cols,
rows,
});
entry.cols = cols;
entry.rows = rows;
if guard.session.as_ref().is_some_and(|s| !s.exited) {
guard.recalculate_pty_size();
let (ncols, nrows) = guard
.session
.as_ref()
.map(|s| (s.cols, s.rows))
.unwrap_or((FALLBACK_COLS, FALLBACK_ROWS));
let targets: Vec<ClientEntry> = guard.clients.values().cloned().collect();
let id = guard.session.as_ref().map(|s| s.id).unwrap_or(SESSION_ID);
let cols = guard.session.as_ref().map(|s| s.cols).unwrap_or(cols);
let rows = guard.session.as_ref().map(|s| s.rows).unwrap_or(rows);
drop(guard);
ServerState::notify_clients(&targets, ncols, nrows);
return Spawn::encode_response((id, cols, rows))
.map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>);
}
let id = SESSION_ID;
let session = Session::spawn(id, cmd, cols, rows, Some(&ch))?;
guard.set_session(session);
guard.recalculate_pty_size();
let (ncols, nrows) = guard
.session
.as_ref()
.map(|s| (s.cols, s.rows))
.unwrap_or((FALLBACK_COLS, FALLBACK_ROWS));
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);
drop(guard);
ServerState::notify_clients(&targets, ncols, nrows);
Spawn::encode_response((sid, scol, srow))
.map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)
}
})
.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 {
let (_id, cols, rows) = ResizePty::decode_request(&payload)
.map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)?;
let mut guard = state.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);
ServerState::notify_clients(&targets, ncols, nrows);
ResizePty::encode_response((ncols, nrows))
.map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)
}
})
.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 {
let _id = CloseSession::decode_request(&payload)
.map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)?;
let mut guard = state.lock().await;
guard.clear_session();
CloseSession::encode_response(())
.map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)
}
})
.await
.map_err(|e| format!("register CloseSession: {e:?}"))?;
let st = Arc::clone(&state);
endpoint
.register_prebuffered(ListSessions::METHOD_ID, move |_payload, _ctx| {
let state = Arc::clone(&st);
async move {
let guard = state.lock().await;
let sessions = match &guard.session {
Some(s) => vec![(s.id, String::new(), s.exited)],
None => vec![],
};
ListSessions::encode_response(sessions)
.map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)
}
})
.await
.map_err(|e| format!("register ListSessions: {e:?}"))?;
let st = Arc::clone(&state);
endpoint
.register_prebuffered(WriteInput::METHOD_ID, move |payload, _ctx| {
let state = Arc::clone(&st);
async move {
let (id, data) = WriteInput::decode_request(&payload)
.map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)?;
let writer = {
let guard = state.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(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)
}
})
.await
.map_err(|e| format!("register WriteInput: {e:?}"))?;
let (input_tx, mut input_rx) = mpsc::channel::<Vec<u8>>(INPUT_CHANNEL_CAPACITY);
endpoint
.register_stream_handler(STREAM_INPUT_METHOD_ID, move |event, _responder, _ctx| {
if let RpcStreamEvent::PayloadChunk { bytes, .. } = event
&& let Err(e) = input_tx.try_send(bytes)
{
tracing::warn!(error = %e, "server input buffer full; dropping input chunk");
}
})
.await
.map_err(|e| format!("register stream handler STREAM_INPUT: {e:?}"))?;
let input_st = Arc::clone(&state);
tokio::spawn(async move {
while let Some(data) = input_rx.recv().await {
let writer = {
let guard = input_st.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);
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);
tokio::spawn(async move {
let mut guard = st.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: ctx.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);
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 guard = st.lock().await;
let handle_clone = handle.clone();
guard.clients.insert(
handle.0.conn_id,
ClientEntry {
caller: Some(handle_clone),
cols: u16::MAX,
rows: u16::MAX,
},
);
}
RpcIpcServerEvent::ClientDisconnected(conn_id) => {
tracing::info!("Client {conn_id} disconnected");
let mut guard = st.lock().await;
guard.clients.remove(&conn_id);
guard.subscribers.retain(|s| s.conn_id != conn_id);
guard.recalculate_pty_size();
let (ncols, nrows) = guard
.session
.as_ref()
.map(|s| (s.cols, s.rows))
.unwrap_or((FALLBACK_COLS, FALLBACK_ROWS));
let targets: Vec<ClientEntry> = guard.clients.values().cloned().collect();
drop(guard);
ServerState::notify_clients(&targets, ncols, nrows);
}
}
}
});
let (exit_tx, mut exit_rx) = oneshot::channel::<i32>();
let st = Arc::clone(&state);
tokio::spawn(async move {
loop {
tokio::select! {
_ = notify.notified() => {}
_ = tokio::time::sleep(SESSION_EXIT_POLL_INTERVAL) => {}
}
let mut guard = st.lock().await;
if guard.subscribers.is_empty() {
let mut exited = false;
if let Some(session) = guard.session.as_mut() {
session.sync_screen();
exited = session.check_exited();
}
if exited {
tracing::info!("Session exited, tearing down");
guard.session = None;
}
continue;
}
let (raw, exited, code) = {
let Some(session) = guard.session.as_mut() else {
let _ = exit_tx.send(0);
break;
};
let raw = session.read_output();
let exited = session.check_exited();
let code = session.exit_code;
(raw, exited, code)
};
if raw.is_empty() && !guard.subscribers.is_empty() {
tracing::debug!(
"PTY output empty with {} subscribers",
guard.subscribers.len()
);
}
if !raw.is_empty() {
for sub in &guard.subscribers {
sub.respond.respond(raw.clone(), false);
}
}
if exited {
for sub in &guard.subscribers {
sub.respond.respond(Vec::new(), true);
}
guard.subscribers.clear();
let _ = exit_tx.send(code.unwrap_or(0));
tracing::info!("Session exited with code {:?}", code);
break;
}
}
});
tracing::info!("Session server listening on channel {}", config.channel);
let exit_code = tokio::select! {
result = async {
server
.serve(&socket_name)
.await
.map_err(|e| format!("serve: {e:?}"))
} => {
result?;
0
}
code = &mut exit_rx => {
tokio::time::sleep(SESSION_EXIT_FLUSH_GRACE).await;
code.unwrap_or(0)
}
};
Ok(exit_code)
}