shell_tunnel/api/
websocket.rs1use 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
19pub 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
28async fn handle_socket(socket: WebSocket, state: AppState, session_id: u64) {
30 let id = SessionId::from_raw(session_id);
31
32 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 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 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 let mut cmd = Command::new(&command);
82 if let Some(secs) = timeout_secs {
83 cmd = cmd.timeout(Duration::from_secs(secs));
84 }
85
86 match state.executor.execute_async(&cmd).await {
88 Ok((mut rx, handle)) => {
89 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 match handle.await {
104 Ok(Ok(result)) => {
105 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 }
164 }
165 }
166}
167
168pub 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
176async 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}