1use 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 identity: Option<axum::Extension<crate::audit::Identity>>,
25) -> impl IntoResponse {
26 let identity = identity.map(|axum::Extension(id)| id);
29 ws.on_upgrade(move |socket| handle_socket(socket, state, session_id, identity))
30}
31
32async 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 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 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 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 let mut cmd = Command::new(&command);
91 if let Some(secs) = timeout_secs {
92 cmd = cmd.timeout(Duration::from_secs(secs));
93 }
94
95 match state.executor.execute_async(&cmd).await {
97 Ok((mut rx, handle)) => {
98 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 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 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 }
186 }
187 }
188}
189
190pub 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
200async 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}