Skip to main content

libfw_server/
ws.rs

1//! WebSocket handler: the unified block-transfer transport.
2//!
3//! Every transfer — **upload or download** — uses the identical block
4//! protocol from [`libfw_core::ws`]:
5//!
6//! - The sender pipelines fixed-size blocks without waiting for a per-block
7//!   acknowledgment (no-ack), so blocks may travel out of order.
8//! - The receiver verifies **every** block in real time (CRC32 + bounds),
9//!   marks bad blocks with `FRAME_NAK` and asks the sender to re-queue them.
10//! - A wave boundary (`FRAME_WAVE_DONE`) triggers reconciliation: the
11//!   receiver replies with `FRAME_REQ` (still-missing blocks) or
12//!   `FRAME_COMPLETE` (everything verified). The sender re-adds requested
13//!   blocks to its transfer queue and re-sends until the receiver is happy.
14//!
15//! All control commands (hello, directory listing, file metadata) travel over
16//! the same WebSocket; there are no separate HTTP calls on the transfer path.
17//! One connection may carry any number of transfers **sequentially** (the
18//! browser client currently opens one connection per file, which lets
19//! multiple files transfer concurrently across separate sockets).
20
21use std::collections::VecDeque;
22use std::io::Read;
23use std::sync::Arc;
24
25use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
26use axum::extract::State;
27use axum::response::Response;
28use bytes::Bytes;
29use futures::SinkExt;
30use libfw_core::auth::{Action, AuthError};
31use libfw_core::claims::TokenClaims;
32use libfw_core::compress::{
33    CompressionFormat, MAX_FRAME_OUTPUT, compressor, decompressor_with_limit,
34};
35use libfw_core::storage::WriteMode;
36use libfw_core::ws::*;
37use libfw_core::{RangeSpec, protocol_compatible};
38
39use crate::{ServerState, validate_rel_path};
40
41/// Default block size for a WebSocket download stream (256 KiB), kept small
42/// so the browser's out-of-order reorder buffer stays modest.
43const DEFAULT_DOWNLOAD_BLOCK: u64 = 256 * 1024;
44/// Default in-flight blocks per wave when the client doesn't specify one
45/// (bounds the receiver's buffering; each wave costs one RTT, so raise it on
46/// high-latency links).
47const SENDER_WINDOW: usize = 16;
48
49/// axum route handler for `GET /ws`.
50///
51/// The upgrade is not gated by the `x-libfw-protocol` HTTP header (a browser
52/// WebSocket handshake cannot carry it); instead the protocol version is
53/// checked inside the `FRAME_HELLO` handshake.
54pub async fn ws_handler(
55    ws: WebSocketUpgrade,
56    State(state): State<Arc<ServerState>>,
57) -> Response {
58    ws.on_upgrade(move |socket| run_socket(socket, state))
59}
60
61// ---------------------------------------------------------------------------
62// Connection lifecycle
63// ---------------------------------------------------------------------------
64
65async fn run_socket(mut socket: WebSocket, state: Arc<ServerState>) {
66    // 1. Handshake: HELLO (protocol + token) → HELLO_OK.
67    let claims = match handshake(&mut socket, &state).await {
68        Ok(claims) => claims,
69        Err(err) => {
70            let _ = send_frame(&mut socket, error_frame("handshake", &err)).await;
71            let _ = socket.close().await;
72            return;
73        }
74    };
75
76    // 2. Control + any number of sequential transfers. All frames are binary
77    //    (first byte = frame type); Text is accepted for robustness.
78    loop {
79        let msg = match socket.recv().await {
80            Some(Ok(msg)) => msg,
81            _ => break,
82        };
83        let frame: Vec<u8> = match msg {
84            Message::Binary(data) => data.to_vec(),
85            Message::Text(text) => text.as_bytes().to_vec(),
86            Message::Close(_) => break,
87            _ => continue,
88        };
89        match frame_type(&frame) {
90            Some(FRAME_LIST_REQ) => {
91                let reply = list_reply(&state, &claims, &frame).await;
92                if send_frame(&mut socket, reply).await.is_err() {
93                    break;
94                }
95            }
96            Some(FRAME_META_REQ) => {
97                let reply = meta_reply(&state, &claims, &frame).await;
98                if send_frame(&mut socket, reply).await.is_err() {
99                    break;
100                }
101            }
102            Some(FRAME_START) => {
103                let start: StartRequest = match parse_control(&frame, FRAME_START) {
104                    Some(s) => s,
105                    None => {
106                        let _ = send_frame(
107                            &mut socket,
108                            error_frame("protocol", "malformed START"),
109                        )
110                        .await;
111                        break;
112                    }
113                };
114                match start.kind {
115                    TransferKind::Download => {
116                        run_download(&mut socket, &state, &claims, start).await;
117                    }
118                    TransferKind::Upload => {
119                        run_upload(&mut socket, &state, &claims, start).await;
120                    }
121                }
122                // Keep the connection open for further control/transfers.
123            }
124            Some(FRAME_COMPLETE) => {
125                // A client abort before any transfer started.
126                break;
127            }
128            // BLOCK/NAK/REQ outside an active transfer are out of protocol.
129            _ => {}
130        }
131    }
132}
133
134/// Perform the `FRAME_HELLO` handshake and return the verified claims.
135async fn handshake(socket: &mut WebSocket, state: &ServerState) -> Result<TokenClaims, String> {
136    let msg = match socket.recv().await {
137        Some(Ok(Message::Binary(data))) => data.to_vec(),
138        Some(Ok(Message::Text(text))) => text.as_bytes().to_vec(),
139        _ => return Err("expected FRAME_HELLO".into()),
140    };
141    let hello: Hello = parse_control(&msg, FRAME_HELLO).ok_or("expected FRAME_HELLO")?;
142    if !protocol_compatible(&hello.protocol) {
143        return Err(format!("unsupported protocol `{}`", hello.protocol));
144    }
145    let claims = state
146        .verifier
147        .verify(&hello.token)
148        .map_err(|e| format!("authentication failed: {e}"))?;
149    let ok = control_frame(FRAME_HELLO_OK, &serde_json::json!({ "ok": true }));
150    send_frame(socket, ok).await.map_err(|_| "send failed".to_string())?;
151    Ok(claims)
152}
153
154/// Send one raw frame over the socket.
155async fn send_frame(socket: &mut WebSocket, frame: Vec<u8>) -> Result<(), ()> {
156    socket
157        .send(Message::Binary(Bytes::from(frame)))
158        .await
159        .map_err(|_| ())
160}
161
162/// Send a `FRAME_COMPLETE` message.
163async fn send_complete(socket: &mut WebSocket, ok: bool, size: u64, err: Option<&str>) {
164    let msg = CompleteMessage {
165        ok,
166        size,
167        error: err.map(str::to_string),
168    };
169    let _ = send_frame(socket, control_frame(FRAME_COMPLETE, &msg)).await;
170}
171
172/// Build a `FRAME_ERROR` frame.
173fn error_frame(code: &str, message: &str) -> Vec<u8> {
174    control_frame(
175        FRAME_ERROR,
176        &ErrorMessage {
177            code: code.to_string(),
178            message: message.to_string(),
179        },
180    )
181}
182
183fn authorize(
184    state: &ServerState,
185    claims: &TokenClaims,
186    path: &str,
187    action: Action,
188) -> Result<(), String> {
189    state.authorize(claims, path, action).map_err(|err| match err {
190        AuthError::Forbidden { path, action } => {
191            format!("permission denied: {action} on `{path}`")
192        }
193        other => format!("unauthorized: {other}"),
194    })
195}
196
197// ---------------------------------------------------------------------------
198// Control: directory listing & metadata
199// ---------------------------------------------------------------------------
200
201async fn list_reply(state: &ServerState, claims: &TokenClaims, frame: &[u8]) -> Vec<u8> {
202    #[derive(serde::Deserialize)]
203    struct ListReq {
204        #[serde(default)]
205        path: String,
206    }
207    let req: ListReq = match serde_json::from_slice(frame_payload(frame)) {
208        Ok(r) => r,
209        Err(_) => return error_frame("protocol", "malformed LIST_REQ"),
210    };
211    let path = match validate_rel_path(&req.path) {
212        Ok(p) => p,
213        Err(e) => return error_frame("path", e),
214    };
215    if let Err(e) = authorize(state, claims, &path, Action::Read) {
216        return error_frame("auth", &e);
217    }
218    match state.storage.list_dir(&path).await {
219        Ok(entries) => control_frame(
220            FRAME_LIST_REPLY,
221            &serde_json::json!({ "path": req.path, "entries": entries }),
222        ),
223        Err(e) => error_frame("storage", &e.to_string()),
224    }
225}
226
227async fn meta_reply(state: &ServerState, claims: &TokenClaims, frame: &[u8]) -> Vec<u8> {
228    #[derive(serde::Deserialize)]
229    struct MetaReq {
230        #[serde(default)]
231        path: String,
232    }
233    let req: MetaReq = match serde_json::from_slice(frame_payload(frame)) {
234        Ok(r) => r,
235        Err(_) => return error_frame("protocol", "malformed META_REQ"),
236    };
237    let path = match validate_rel_path(&req.path) {
238        Ok(p) => p,
239        Err(e) => return error_frame("path", e),
240    };
241    if let Err(e) = authorize(state, claims, &path, Action::Read) {
242        return error_frame("auth", &e);
243    }
244    match state.storage.file_meta(&path).await {
245        Ok(Some(meta)) => control_frame(
246            FRAME_META_REPLY,
247            &serde_json::json!({
248                "path": meta.path,
249                "size": meta.size,
250                "mtime": meta.mtime,
251                "etag": meta.etag,
252            }),
253        ),
254        Ok(None) => error_frame("not_found", &format!("file not found: {path}")),
255        Err(e) => error_frame("storage", &e.to_string()),
256    }
257}
258
259// ---------------------------------------------------------------------------
260// Download (server is the SENDER)
261// ---------------------------------------------------------------------------
262
263/// Serve one download stream: slice `[offset, size)` into blocks and push
264/// them to the client (receiver), handling real-time NAK/REQ retransmission.
265async fn run_download(
266    socket: &mut WebSocket,
267    state: &ServerState,
268    claims: &TokenClaims,
269    start: StartRequest,
270) {
271    let path = match validate_rel_path(&start.path) {
272        Ok(p) => p,
273        Err(e) => {
274            let _ = send_frame(socket, error_frame("path", e)).await;
275            return;
276        }
277    };
278    if let Err(e) = authorize(state, claims, &path, Action::Read) {
279        let _ = send_frame(socket, error_frame("auth", &e)).await;
280        return;
281    }
282    let meta = match state.storage.file_meta(&path).await {
283        Ok(Some(m)) => m,
284        Ok(None) => {
285            let _ = send_frame(socket, error_frame("not_found", &path)).await;
286            return;
287        }
288        Err(e) => {
289            let _ = send_frame(socket, error_frame("storage", &e.to_string())).await;
290            return;
291        }
292    };
293
294    let block_size = if start.block_size > 0 {
295        start.block_size
296    } else {
297        DEFAULT_DOWNLOAD_BLOCK
298    };
299    let window = if start.window > 0 {
300        start.window as usize
301    } else {
302        SENDER_WINDOW
303    };
304    // Whether the server actually compresses (only when the client asked AND
305    // the server's configured download compression is zrip).
306    let compress = start.compress && state.compression == CompressionFormat::Zrip;
307    let start_off = start.offset.min(meta.size);
308    let total = meta.size - start_off;
309    let total_blocks = block_count(total, block_size);
310
311    let ready = ReadyReply {
312        kind: TransferKind::Download,
313        path: path.clone(),
314        size: meta.size,
315        mtime: meta.mtime,
316        etag: meta.etag.clone(),
317        compress,
318        block_size,
319        total_blocks,
320        offset: start_off,
321        received: Vec::new(),
322    };
323    if send_frame(socket, control_frame(FRAME_READY, &ready)).await.is_err() {
324        return;
325    }
326
327    // Sender transfer queue: indices to send. NAK/REQ re-add to this queue.
328    let mut queue: VecDeque<u32> = (0..total_blocks).collect();
329
330    loop {
331        // 1. Read one wave of blocks into memory, then send them.
332        //
333        // Each block is read with a FRESH reader opened at its absolute
334        // offset and drained in a tight loop. This avoids a Windows/tokio
335        // quirk where a long-lived `read_stream` reader returns early EOF
336        // once its reads are interleaved with socket awaits — each block
337        // read is independent, and memory stays bounded (one wave).
338        let mut wave: Vec<(u32, Vec<u8>, u32, u32)> = Vec::new();
339        let mut sent = 0usize;
340        while sent < window {
341            let Some(idx) = queue.pop_front() else {
342                break;
343            };
344            let (start, end) = block_bounds(idx, block_size, total);
345            let abs = block_offset(idx, block_size, start_off);
346            let mut data = vec![0u8; (end - start) as usize];
347            let mut reader = match state
348                .storage
349                .read_stream(&path, RangeSpec { start: abs, end: abs + (end - start) })
350                .await
351            {
352                Ok(r) => r,
353                Err(e) => {
354                    let _ = send_frame(socket, error_frame("storage", &e.to_string())).await;
355                    return;
356                }
357            };
358            if read_exact(&mut reader, &mut data).is_err() {
359                let _ = send_frame(socket, error_frame("io", "read failed")).await;
360                return;
361            }
362            let raw_len = data.len() as u32;
363            let payload: Vec<u8> = if compress {
364                match compress_frame(&data) {
365                    Some(p) => p,
366                    None => {
367                        let _ =
368                            send_frame(socket, error_frame("compress", "compress failed")).await;
369                        return;
370                    }
371                }
372            } else {
373                data
374            };
375            let crc = crc32(&payload);
376            wave.push((idx, payload, crc, raw_len));
377            sent += 1;
378        }
379
380        // 2. Send the wave's blocks (no per-block ack).
381        for (idx, payload, crc, raw_len) in wave {
382            let frame = block_frame(idx, crc, raw_len, &payload);
383            if send_frame(socket, frame).await.is_err() {
384                return;
385            }
386        }
387
388        // 3. Wave boundary: the receiver reconciles.
389        if send_frame(socket, wave_done_frame()).await.is_err() {
390            return;
391        }
392
393        // 4. Read events until the receiver asks for more (REQ) or is done
394        //    (COMPLETE). NAKs re-queue immediately ("实时核验 → 重传队列").
395        loop {
396            let msg = match socket.recv().await {
397                Some(Ok(msg)) => msg,
398                _ => return, // client gone
399            };
400            let frame: Vec<u8> = match msg {
401                Message::Binary(data) => data.to_vec(),
402                Message::Text(text) => text.as_bytes().to_vec(),
403                Message::Close(_) => return,
404                _ => continue,
405            };
406            match frame_type(&frame) {
407                Some(FRAME_NAK) => {
408                    if let Some(index) = parse_nak(&frame) {
409                        queue.push_back(index);
410                    }
411                }
412                Some(FRAME_REQ) => {
413                    if let Some(indices) = parse_req(&frame) {
414                        queue.extend(indices);
415                    }
416                    break; // next wave
417                }
418                Some(FRAME_COMPLETE) => return, // receiver finished → done
419                _ => {}
420            }
421        }
422    }
423}
424
425// ---------------------------------------------------------------------------
426// Upload (server is the RECEIVER)
427// ---------------------------------------------------------------------------
428
429/// Receive one upload stream: verify every block (CRC32 + bounds), write it
430/// at its absolute offset into a shared session temp, NAK bad blocks and
431/// commit once all blocks are verified.
432async fn run_upload(
433    socket: &mut WebSocket,
434    state: &ServerState,
435    claims: &TokenClaims,
436    start: StartRequest,
437) {
438    let path = match validate_rel_path(&start.path) {
439        Ok(p) => p,
440        Err(e) => {
441            let _ = send_frame(socket, error_frame("path", e)).await;
442            return;
443        }
444    };
445    if let Err(e) = authorize(state, claims, &path, Action::Write) {
446        let _ = send_frame(socket, error_frame("auth", &e)).await;
447        return;
448    }
449    if start.size > state.max_upload_size {
450        let _ = send_frame(
451            socket,
452            error_frame(
453                "too_large",
454                &format!("upload exceeds limit of {} bytes", state.max_upload_size),
455            ),
456        )
457        .await;
458        return;
459    }
460
461    let block_size = if start.block_size > 0 {
462        start.block_size
463    } else {
464        libfw_core::CHUNK_SIZE
465    };
466    let total_blocks = block_count(start.size, block_size);
467    let mode = if start.mode.eq_ignore_ascii_case("create") {
468        WriteMode::Create
469    } else {
470        WriteMode::Overwrite
471    };
472    // Deterministic session id from the ETag → an interrupted upload of the
473    // same file version resumes the same shared temp on the server.
474    let session = start.etag.trim_matches('"');
475
476    let mut sink = match state.storage.write_stream_session(&path, session, mode).await {
477        Ok(s) => s,
478        Err(e) => {
479            let _ = send_frame(socket, error_frame("storage", &e.to_string())).await;
480            return;
481        }
482    };
483
484    // Seed the verified set + resume ranges from what the server already
485    // holds, so the client only retransmits the missing blocks.
486    let received = sink.received_ranges().await.unwrap_or_default();
487    let received_pairs: Vec<[u64; 2]> = received.iter().map(|r| [r.start, r.end]).collect();
488    let received_slices: Vec<(u64, u64)> = received.iter().map(|r| (r.start, r.end)).collect();
489    let mut verified = BlockSet::new(total_blocks);
490    verified.seed_from_ranges(block_size, &received_slices);
491
492    let ready = ReadyReply {
493        kind: TransferKind::Upload,
494        path: path.clone(),
495        size: start.size,
496        mtime: start.mtime,
497        etag: start.etag.clone(),
498        compress: start.compress,
499        block_size,
500        total_blocks,
501        offset: 0,
502        received: received_pairs,
503    };
504    if send_frame(socket, control_frame(FRAME_READY, &ready)).await.is_err() {
505        return; // keep the session temp for a later resume
506    }
507
508    let compress = start.compress;
509    loop {
510        let msg = match socket.recv().await {
511            Some(Ok(msg)) => msg,
512            _ => {
513                // Client gone: drop the sink WITHOUT abort so the session
514                // temp survives for a resumable retry (tus expiration).
515                return;
516            }
517        };
518        let frame: Vec<u8> = match msg {
519            Message::Binary(data) => data.to_vec(),
520            Message::Text(text) => text.as_bytes().to_vec(),
521            Message::Close(_) => {
522                // Keep the session temp for a possible resume.
523                return;
524            }
525            _ => continue,
526        };
527        match frame_type(&frame) {
528            Some(FRAME_BLOCK) => {
529                let Some(block) = parse_block(&frame) else {
530                    continue;
531                };
532                // Already-verified (duplicate / re-sent) or out of range →
533                // idempotent no-op.
534                if block.index >= total_blocks || verified.contains(block.index) {
535                    continue;
536                }
537                // Real-time verification: CRC + length + bounds.
538                let crc_ok = crc32(&block.data) == block.crc;
539                let raw: Vec<u8> = if compress {
540                    match decompress_frame(&block.data) {
541                        Ok(d) => d,
542                        Err(_) => {
543                            let _ = send_frame(socket, nak_frame(block.index)).await;
544                            continue;
545                        }
546                    }
547                } else {
548                    block.data
549                };
550                let len_ok = !compress || raw.len() as u32 == block.raw_len;
551                let abs = block_offset(block.index, block_size, 0);
552                let end = abs.saturating_add(raw.len() as u64);
553                let in_bounds = end <= start.size && end <= state.max_upload_size;
554                if !(crc_ok && len_ok && in_bounds) {
555                    // Mark bad → ask the sender to re-queue it.
556                    let _ = send_frame(socket, nak_frame(block.index)).await;
557                    continue;
558                }
559                if let Err(e) = sink.write_at(abs, &raw).await {
560                    let _ = send_frame(
561                        socket,
562                        control_frame(
563                            FRAME_COMPLETE,
564                            &CompleteMessage::err(format!("write failed: {e}")),
565                        ),
566                    )
567                    .await;
568                    let _ = sink.abort().await;
569                    return;
570                }
571                verified.insert(block.index);
572            }
573            Some(FRAME_WAVE_DONE) => {
574                // Reconciliation: everything verified → commit, else ask the
575                // sender to re-send the missing blocks.
576                let missing = verified.missing();
577                if missing.is_empty() {
578                    let len = match sink.len().await {
579                        Ok(l) => l,
580                        Err(e) => {
581                            let _ = send_complete(
582                                socket,
583                                false,
584                                0,
585                                Some(&format!("len failed: {e}")),
586                            )
587                            .await;
588                            let _ = sink.abort().await;
589                            return;
590                        }
591                    };
592                    if len != start.size {
593                        let _ = send_complete(socket, false, 0, Some("commit size mismatch")).await;
594                        let _ = sink.abort().await;
595                        return;
596                    }
597                    // `commit` consumes the sink; on failure there is nothing
598                    // left to abort (temp is best-effort).
599                    match sink.commit().await {
600                        Ok(_) => {
601                            let _ = send_complete(socket, true, start.size, None).await;
602                            return;
603                        }
604                        Err(e) => {
605                            let _ = send_complete(
606                                socket,
607                                false,
608                                0,
609                                Some(&format!("commit failed: {e}")),
610                            )
611                            .await;
612                            return;
613                        }
614                    }
615                } else {
616                    let _ = send_frame(socket, req_frame(&missing)).await;
617                }
618            }
619            Some(FRAME_COMPLETE) => {
620                // A client abort mid-upload keeps the session temp.
621                return;
622            }
623            _ => {}
624        }
625    }
626}
627
628// ---------------------------------------------------------------------------
629// Small helpers
630// ---------------------------------------------------------------------------
631
632/// Read exactly `buf.len()` bytes (or fail on early EOF) from a reader.
633fn read_exact(reader: &mut Box<dyn Read + Send>, buf: &mut [u8]) -> Result<(), ()> {
634    let mut filled = 0usize;
635    while filled < buf.len() {
636        match reader.read(&mut buf[filled..]) {
637            Ok(0) => return Err(()),
638            Ok(n) => filled += n,
639            Err(_) => return Err(()),
640        }
641    }
642    Ok(())
643}
644
645/// Compress `data` into one independent zrip frame.
646fn compress_frame(data: &[u8]) -> Option<Vec<u8>> {
647    let mut enc = compressor(CompressionFormat::Zrip).ok()?;
648    let mut out = Vec::with_capacity(data.len());
649    enc.compress(data, &mut out).ok()?;
650    enc.finish(&mut out).ok()?;
651    Some(out)
652}
653
654/// Decompress one independent zrip frame.
655fn decompress_frame(data: &[u8]) -> Result<Vec<u8>, libfw_core::StorageError> {
656    let mut dec = decompressor_with_limit(CompressionFormat::Zrip, MAX_FRAME_OUTPUT);
657    let mut out: Vec<u8> = Vec::new();
658    dec.decompress(data, &mut out)
659        .map_err(|e| libfw_core::StorageError::Other(std::io::Error::other(format!("{e}"))))?;
660    dec.finish(&mut out)
661        .map_err(|e| libfw_core::StorageError::Other(std::io::Error::other(format!("{e}"))))?;
662    Ok(out)
663}