1use 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
41const DEFAULT_DOWNLOAD_BLOCK: u64 = 256 * 1024;
44const SENDER_WINDOW: usize = 16;
48
49pub 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
61async fn run_socket(mut socket: WebSocket, state: Arc<ServerState>) {
66 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 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 }
124 Some(FRAME_COMPLETE) => {
125 break;
127 }
128 _ => {}
130 }
131 }
132}
133
134async 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
154async 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
162async 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
172fn 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
197async 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
259async 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 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 let mut queue: VecDeque<u32> = (0..total_blocks).collect();
329
330 loop {
331 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 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 if send_frame(socket, wave_done_frame()).await.is_err() {
390 return;
391 }
392
393 loop {
396 let msg = match socket.recv().await {
397 Some(Ok(msg)) => msg,
398 _ => return, };
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; }
418 Some(FRAME_COMPLETE) => return, _ => {}
420 }
421 }
422 }
423}
424
425async 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 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 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; }
507
508 let compress = start.compress;
509 loop {
510 let msg = match socket.recv().await {
511 Some(Ok(msg)) => msg,
512 _ => {
513 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 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 if block.index >= total_blocks || verified.contains(block.index) {
535 continue;
536 }
537 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 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 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 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 return;
622 }
623 _ => {}
624 }
625 }
626}
627
628fn 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
645fn 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
654fn 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}