Skip to main content

shell_tunnel/api/
websocket.rs

1//! WebSocket handler for real-time command streaming.
2
3use std::time::Duration;
4
5use axum::{
6    extract::{
7        ws::{Message, WebSocket, WebSocketUpgrade},
8        Path, State,
9    },
10    response::IntoResponse,
11};
12use futures_util::{SinkExt, StreamExt};
13
14use super::handlers::AppState;
15use super::types::WsMessage;
16use crate::execution::Command;
17use crate::session::SessionId;
18
19/// WebSocket upgrade handler.
20pub async fn ws_handler(
21    ws: WebSocketUpgrade,
22    State(state): State<AppState>,
23    Path(session_id): Path<u64>,
24) -> impl IntoResponse {
25    ws.on_upgrade(move |socket| handle_socket(socket, state, session_id))
26}
27
28/// Handle WebSocket connection.
29async fn handle_socket(socket: WebSocket, state: AppState, session_id: u64) {
30    let id = SessionId::from_raw(session_id);
31
32    // Verify session exists
33    if state.store.get(&id).ok().flatten().is_none() {
34        let (mut sink, _) = socket.split();
35        let err = WsMessage::Error {
36            code: "SESSION_NOT_FOUND".to_string(),
37            message: format!("Session {} not found", session_id),
38        };
39        if let Ok(json) = serde_json::to_string(&err) {
40            let _ = sink.send(Message::Text(json.into())).await;
41        }
42        return;
43    }
44
45    let (mut sink, mut stream) = socket.split();
46
47    // Process incoming messages
48    while let Some(msg) = stream.next().await {
49        let msg = match msg {
50            Ok(Message::Text(text)) => text.to_string(),
51            Ok(Message::Close(_)) => break,
52            Ok(Message::Ping(data)) => {
53                let _ = sink.send(Message::Pong(data)).await;
54                continue;
55            }
56            Ok(_) => continue,
57            Err(_) => break,
58        };
59
60        // Parse WebSocket message
61        let ws_msg: WsMessage = match serde_json::from_str(&msg) {
62            Ok(m) => m,
63            Err(e) => {
64                let err = WsMessage::Error {
65                    code: "PARSE_ERROR".to_string(),
66                    message: e.to_string(),
67                };
68                if let Ok(json) = serde_json::to_string(&err) {
69                    let _ = sink.send(Message::Text(json.into())).await;
70                }
71                continue;
72            }
73        };
74
75        match ws_msg {
76            WsMessage::Execute {
77                command,
78                timeout_secs,
79            } => {
80                // Build command
81                let mut cmd = Command::new(&command);
82                if let Some(secs) = timeout_secs {
83                    cmd = cmd.timeout(Duration::from_secs(secs));
84                }
85
86                // Execute with streaming
87                match state.executor.execute_async(&cmd).await {
88                    Ok((mut rx, handle)) => {
89                        // Stream output chunks
90                        while let Some(chunk) = rx.recv().await {
91                            let output = WsMessage::Output {
92                                data: String::from_utf8_lossy(&chunk.raw).to_string(),
93                                is_final: false,
94                            };
95                            if let Ok(json) = serde_json::to_string(&output) {
96                                if sink.send(Message::Text(json.into())).await.is_err() {
97                                    break;
98                                }
99                            }
100                        }
101
102                        // Wait for completion and send result
103                        match handle.await {
104                            Ok(Ok(result)) => {
105                                // Update session context
106                                state
107                                    .store
108                                    .update(&id, |s| {
109                                        s.context.record_execution(&command, result.exit_code);
110                                    })
111                                    .ok();
112
113                                let result_msg = WsMessage::Result {
114                                    success: result.exit_code.map(|c| c == 0).unwrap_or(false)
115                                        && !result.timed_out,
116                                    exit_code: result.exit_code,
117                                    duration_ms: result.duration.as_millis() as u64,
118                                    timed_out: result.timed_out,
119                                };
120                                if let Ok(json) = serde_json::to_string(&result_msg) {
121                                    let _ = sink.send(Message::Text(json.into())).await;
122                                }
123                            }
124                            Ok(Err(e)) => {
125                                let err = WsMessage::Error {
126                                    code: "EXECUTION_ERROR".to_string(),
127                                    message: e.to_string(),
128                                };
129                                if let Ok(json) = serde_json::to_string(&err) {
130                                    let _ = sink.send(Message::Text(json.into())).await;
131                                }
132                            }
133                            Err(e) => {
134                                let err = WsMessage::Error {
135                                    code: "TASK_ERROR".to_string(),
136                                    message: e.to_string(),
137                                };
138                                if let Ok(json) = serde_json::to_string(&err) {
139                                    let _ = sink.send(Message::Text(json.into())).await;
140                                }
141                            }
142                        }
143                    }
144                    Err(e) => {
145                        let err = WsMessage::Error {
146                            code: "EXECUTION_ERROR".to_string(),
147                            message: e.to_string(),
148                        };
149                        if let Ok(json) = serde_json::to_string(&err) {
150                            let _ = sink.send(Message::Text(json.into())).await;
151                        }
152                    }
153                }
154            }
155            WsMessage::Ping => {
156                let pong = WsMessage::Pong;
157                if let Ok(json) = serde_json::to_string(&pong) {
158                    let _ = sink.send(Message::Text(json.into())).await;
159                }
160            }
161            _ => {
162                // Ignore other message types from client
163            }
164        }
165    }
166}
167
168/// One-shot WebSocket execution (no session required).
169pub async fn ws_oneshot_handler(
170    ws: WebSocketUpgrade,
171    State(state): State<AppState>,
172) -> impl IntoResponse {
173    ws.on_upgrade(move |socket| handle_oneshot_socket(socket, state))
174}
175
176/// Handle one-shot WebSocket connection.
177async fn handle_oneshot_socket(socket: WebSocket, state: AppState) {
178    let (mut sink, mut stream) = socket.split();
179
180    while let Some(msg) = stream.next().await {
181        let msg = match msg {
182            Ok(Message::Text(text)) => text.to_string(),
183            Ok(Message::Close(_)) => break,
184            Ok(Message::Ping(data)) => {
185                let _ = sink.send(Message::Pong(data)).await;
186                continue;
187            }
188            Ok(_) => continue,
189            Err(_) => break,
190        };
191
192        let ws_msg: WsMessage = match serde_json::from_str(&msg) {
193            Ok(m) => m,
194            Err(e) => {
195                let err = WsMessage::Error {
196                    code: "PARSE_ERROR".to_string(),
197                    message: e.to_string(),
198                };
199                if let Ok(json) = serde_json::to_string(&err) {
200                    let _ = sink.send(Message::Text(json.into())).await;
201                }
202                continue;
203            }
204        };
205
206        match ws_msg {
207            WsMessage::Execute {
208                command,
209                timeout_secs,
210            } => {
211                let mut cmd = Command::new(&command);
212                if let Some(secs) = timeout_secs {
213                    cmd = cmd.timeout(Duration::from_secs(secs));
214                }
215
216                match state.executor.execute_async(&cmd).await {
217                    Ok((mut rx, handle)) => {
218                        while let Some(chunk) = rx.recv().await {
219                            let output = WsMessage::Output {
220                                data: String::from_utf8_lossy(&chunk.raw).to_string(),
221                                is_final: false,
222                            };
223                            if let Ok(json) = serde_json::to_string(&output) {
224                                if sink.send(Message::Text(json.into())).await.is_err() {
225                                    break;
226                                }
227                            }
228                        }
229
230                        match handle.await {
231                            Ok(Ok(result)) => {
232                                let result_msg = WsMessage::Result {
233                                    success: result.exit_code.map(|c| c == 0).unwrap_or(false)
234                                        && !result.timed_out,
235                                    exit_code: result.exit_code,
236                                    duration_ms: result.duration.as_millis() as u64,
237                                    timed_out: result.timed_out,
238                                };
239                                if let Ok(json) = serde_json::to_string(&result_msg) {
240                                    let _ = sink.send(Message::Text(json.into())).await;
241                                }
242                            }
243                            Ok(Err(e)) => {
244                                let err = WsMessage::Error {
245                                    code: "EXECUTION_ERROR".to_string(),
246                                    message: e.to_string(),
247                                };
248                                if let Ok(json) = serde_json::to_string(&err) {
249                                    let _ = sink.send(Message::Text(json.into())).await;
250                                }
251                            }
252                            Err(e) => {
253                                let err = WsMessage::Error {
254                                    code: "TASK_ERROR".to_string(),
255                                    message: e.to_string(),
256                                };
257                                if let Ok(json) = serde_json::to_string(&err) {
258                                    let _ = sink.send(Message::Text(json.into())).await;
259                                }
260                            }
261                        }
262                    }
263                    Err(e) => {
264                        let err = WsMessage::Error {
265                            code: "EXECUTION_ERROR".to_string(),
266                            message: e.to_string(),
267                        };
268                        if let Ok(json) = serde_json::to_string(&err) {
269                            let _ = sink.send(Message::Text(json.into())).await;
270                        }
271                    }
272                }
273            }
274            WsMessage::Ping => {
275                let pong = WsMessage::Pong;
276                if let Ok(json) = serde_json::to_string(&pong) {
277                    let _ = sink.send(Message::Text(json.into())).await;
278                }
279            }
280            _ => {}
281        }
282    }
283}
284
285#[cfg(test)]
286mod tests {
287    use super::*;
288
289    #[test]
290    fn test_ws_message_execute_parse() {
291        let json = r#"{"type": "execute", "command": "echo hello"}"#;
292        let msg: WsMessage = serde_json::from_str(json).unwrap();
293        match msg {
294            WsMessage::Execute { command, .. } => assert_eq!(command, "echo hello"),
295            _ => panic!("Expected Execute message"),
296        }
297    }
298
299    #[test]
300    fn test_ws_message_ping_parse() {
301        let json = r#"{"type": "ping"}"#;
302        let msg: WsMessage = serde_json::from_str(json).unwrap();
303        assert!(matches!(msg, WsMessage::Ping));
304    }
305}