Skip to main content

term_session_server/
session_server.rs

1use std::sync::Arc;
2use std::time::Duration;
3
4use muxio_core::rpc::rpc_internals::RpcStreamEvent;
5use muxio_rpc_service::prebuffered::RpcMethodPrebuffered;
6use muxio_rpc_service_endpoint::{RpcServiceEndpointInterface, StreamResponder};
7use muxio_tokio_rpc_ipc_server::{RpcIpcServer, RpcIpcServerEvent};
8use portable_pty::PtySize;
9use tokio::sync::{Mutex, mpsc, oneshot};
10
11use term_session_muxio_service_definitions::{
12    CloseSession, ListSessions, ResizePty, STREAM_INPUT_METHOD_ID, SUBSCRIBE_OUTPUT_METHOD_ID,
13    Spawn, WriteInput,
14};
15
16use crate::session::Session;
17
18pub struct SessionServerConfig {
19    pub socket_path: String,
20    pub cmd: Vec<String>,
21    pub cols: u16,
22    pub rows: u16,
23}
24
25struct ClientEntry {
26    conn_id: usize,
27}
28
29struct SubscriberEntry {
30    conn_id: usize,
31    respond: StreamResponder,
32}
33
34struct ServerState {
35    session: Option<Session>,
36    clients: Vec<ClientEntry>,
37    subscribers: Vec<SubscriberEntry>,
38}
39
40impl ServerState {
41    fn new() -> Self {
42        Self {
43            session: None,
44            clients: Vec::new(),
45            subscribers: Vec::new(),
46        }
47    }
48}
49
50type SharedState = Arc<Mutex<ServerState>>;
51
52/// Run the session server. Returns the PTY child's exit code on success.
53pub async fn run_server(
54    config: SessionServerConfig,
55) -> Result<i32, Box<dyn std::error::Error + Send + Sync>> {
56    let state: SharedState = Arc::new(Mutex::new(ServerState::new()));
57
58    // Spawn initial session
59    {
60        let mut st = state.lock().await;
61        let cmd = if config.cmd.is_empty() {
62            None
63        } else {
64            Some(config.cmd.clone())
65        };
66        let session = Session::spawn(1, cmd, config.cols, config.rows)?;
67        st.session = Some(session);
68    }
69
70    let (event_tx, mut event_rx) = mpsc::unbounded_channel();
71    let server = RpcIpcServer::new(Some(event_tx));
72    let endpoint = server.endpoint();
73
74    // Register Spawn
75    let st = Arc::clone(&state);
76    endpoint
77        .register_prebuffered(Spawn::METHOD_ID, move |payload, _ctx| {
78            let state = Arc::clone(&st);
79            async move {
80                let mut guard = state.lock().await;
81
82                let (cmd, cols, rows) = Spawn::decode_request(&payload)
83                    .map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)?;
84
85                // If a session already exists and hasn't exited, resize and return it.
86                if let Some(ref mut session) = guard.session
87                    && !session.exited
88                {
89                    let size = PtySize {
90                        rows,
91                        cols,
92                        pixel_width: 0,
93                        pixel_height: 0,
94                    };
95                    let _ = session.pty.resize(size);
96                    session.cols = cols;
97                    session.rows = rows;
98
99                    return Spawn::encode_response(session.id)
100                        .map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>);
101                }
102
103                let id = 1;
104                let session = Session::spawn(id, cmd, cols, rows)?;
105                guard.session = Some(session);
106                Spawn::encode_response(id)
107                    .map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)
108            }
109        })
110        .await
111        .map_err(|e| format!("register Spawn: {e:?}"))?;
112
113    // Register ResizePty
114    let st = Arc::clone(&state);
115    endpoint
116        .register_prebuffered(ResizePty::METHOD_ID, move |payload, _ctx| {
117            let state = Arc::clone(&st);
118            async move {
119                let (_id, cols, rows) = ResizePty::decode_request(&payload)
120                    .map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)?;
121                let mut guard = state.lock().await;
122                if let Some(session) = guard.session.as_mut() {
123                    let size = portable_pty::PtySize {
124                        rows,
125                        cols,
126                        pixel_width: 0,
127                        pixel_height: 0,
128                    };
129                    let _ = session.pty.resize(size);
130                }
131                ResizePty::encode_response(())
132                    .map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)
133            }
134        })
135        .await
136        .map_err(|e| format!("register ResizePty: {e:?}"))?;
137
138    // Register CloseSession
139    let st = Arc::clone(&state);
140    endpoint
141        .register_prebuffered(CloseSession::METHOD_ID, move |payload, _ctx| {
142            let state = Arc::clone(&st);
143            async move {
144                let _id = CloseSession::decode_request(&payload)
145                    .map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)?;
146                let mut guard = state.lock().await;
147                if let Some(session) = guard.session.as_mut() {
148                    let _ = session.pty.kill_child();
149                }
150                guard.session = None;
151                CloseSession::encode_response(())
152                    .map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)
153            }
154        })
155        .await
156        .map_err(|e| format!("register CloseSession: {e:?}"))?;
157
158    // Register ListSessions
159    let st = Arc::clone(&state);
160    endpoint
161        .register_prebuffered(ListSessions::METHOD_ID, move |_payload, _ctx| {
162            let state = Arc::clone(&st);
163            async move {
164                let guard = state.lock().await;
165                let sessions = match &guard.session {
166                    Some(s) => vec![(s.id, String::new(), s.exited)],
167                    None => vec![],
168                };
169                ListSessions::encode_response(sessions)
170                    .map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)
171            }
172        })
173        .await
174        .map_err(|e| format!("register ListSessions: {e:?}"))?;
175
176    // Register WriteInput
177    let st = Arc::clone(&state);
178    endpoint
179        .register_prebuffered(WriteInput::METHOD_ID, move |payload, _ctx| {
180            let state = Arc::clone(&st);
181            async move {
182                let (id, data) = WriteInput::decode_request(&payload)
183                    .map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)?;
184                let mut guard = state.lock().await;
185                if let Some(session) = guard.session.as_mut()
186                    && session.id == id
187                {
188                    let _ = session.pty.write_bytes(&data);
189                }
190                WriteInput::encode_response(())
191                    .map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)
192            }
193        })
194        .await
195        .map_err(|e| format!("register WriteInput: {e:?}"))?;
196
197    // Register StreamInput (streaming handler for PTY input)
198    // The channel persists across client disconnects so reconnecting
199    // clients can still send input — we drop it only when the server
200    // shuts down.
201    let (input_tx, mut input_rx) = mpsc::unbounded_channel::<Vec<u8>>();
202    endpoint
203        .register_stream_handler(STREAM_INPUT_METHOD_ID, move |event, _responder, _ctx| {
204            if let RpcStreamEvent::PayloadChunk { bytes, .. } = event {
205                let _ = input_tx.send(bytes);
206            }
207            // Intentionally ignore End/Error — the channel stays alive.
208        })
209        .await
210        .map_err(|e| format!("register stream handler STREAM_INPUT: {e:?}"))?;
211
212    // Background task: write received input bytes to the PTY session
213    let input_st = Arc::clone(&state);
214    tokio::spawn(async move {
215        while let Some(data) = input_rx.recv().await {
216            let mut guard = input_st.lock().await;
217            if let Some(session) = guard.session.as_mut() {
218                let _ = session.pty.write_bytes(&data);
219            }
220        }
221    });
222
223    // Register SubscribeOutput (streaming handler for PTY output pushes)
224    let st = Arc::clone(&state);
225    endpoint
226        .register_stream_handler(SUBSCRIBE_OUTPUT_METHOD_ID, move |event, respond, ctx| {
227            let is_new = matches!(&event, RpcStreamEvent::Header { .. });
228            if is_new {
229                let st = Arc::clone(&st);
230                tokio::spawn(async move {
231                    let mut guard = st.lock().await;
232
233                    // Drain accumulated PTY output and capture the raw bytes
234                    // so they can be sent to the new subscriber (not just the snapshot).
235                    let early = guard.session.as_mut().and_then(|s| {
236                        let data = s.read_output();
237                        if data.is_empty() { None } else { Some(data) }
238                    });
239                    let snapshot = guard.session.as_mut().map(|s| s.generate_snapshot());
240
241                    guard.subscribers.push(SubscriberEntry {
242                        conn_id: ctx.conn_id,
243                        respond: respond.clone(),
244                    });
245
246                    drop(guard);
247
248                    if let Some(data) = snapshot
249                        && !data.is_empty()
250                    {
251                        respond.respond(data, false);
252                    }
253                    if let Some(data) = early {
254                        respond.respond(data, false);
255                    }
256                });
257            }
258        })
259        .await
260        .map_err(|e| format!("register SubscribeOutput: {e:?}"))?;
261
262    // Connection event handler
263    let st = Arc::clone(&state);
264    tokio::spawn(async move {
265        while let Some(event) = event_rx.recv().await {
266            match event {
267                RpcIpcServerEvent::ClientConnected(handle) => {
268                    tracing::info!("Client {} connected", handle.0.conn_id);
269
270                    let mut guard = st.lock().await;
271                    guard.clients.push(ClientEntry {
272                        conn_id: handle.0.conn_id,
273                    });
274                }
275                RpcIpcServerEvent::ClientDisconnected(conn_id) => {
276                    tracing::info!("Client {conn_id} disconnected");
277                    let mut guard = st.lock().await;
278                    guard.clients.retain(|c| c.conn_id != conn_id);
279                    guard.subscribers.retain(|s| s.conn_id != conn_id);
280                }
281            }
282        }
283    });
284
285    // Output polling and push via stored StreamResponders.
286    // When the session exits, the exit code is sent back through this
287    // channel so run_server can return it.
288    let (exit_tx, mut exit_rx) = oneshot::channel::<i32>();
289    let st = Arc::clone(&state);
290    tokio::spawn(async move {
291        let mut interval = tokio::time::interval(Duration::from_millis(8));
292        loop {
293            interval.tick().await;
294
295            let mut guard = st.lock().await;
296
297            if guard.subscribers.is_empty() {
298                if let Some(session) = guard.session.as_mut() {
299                    session.sync_screen();
300                }
301                continue;
302            }
303
304            let Some(session) = guard.session.as_mut() else {
305                break;
306            };
307
308            let raw = session.read_output();
309            let exited = session.check_exited();
310            let code = session.exit_code;
311
312            if raw.is_empty() && !guard.subscribers.is_empty() {
313                tracing::debug!(
314                    "PTY output empty with {} subscribers",
315                    guard.subscribers.len()
316                );
317            }
318
319            // Push raw PTY output to all subscribers via StreamResponder
320            if !raw.is_empty() {
321                for sub in &guard.subscribers {
322                    sub.respond.respond(raw.clone(), false);
323                }
324            }
325
326            // On exit: finalize all streams and clean up
327            if exited {
328                for sub in &guard.subscribers {
329                    sub.respond.respond(Vec::new(), true);
330                }
331                guard.subscribers.clear();
332                let _ = exit_tx.send(code.unwrap_or(0));
333                tracing::info!("Session exited with code {:?}", code);
334                break;
335            }
336        }
337    });
338
339    tracing::info!("Session server listening on {}", config.socket_path);
340
341    // Wait for either the server to finish or the session to exit.
342    let exit_code = tokio::select! {
343        result = server.serve(&config.socket_path) => {
344            result.map_err(|e| format!("serve: {e:?}"))?;
345            0
346        }
347        code = &mut exit_rx => {
348            code.unwrap_or(0)
349        }
350    };
351
352    Ok(exit_code)
353}