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