1use crate::{
7 Error as PjsError, Result as PjsResult, StreamFrame, domain::Priority, security::RateLimitGuard,
8};
9use futures::{Sink, SinkExt};
10use serde::{Deserialize, Serialize};
11use serde_json::Value;
12use sha2::{Digest, Sha256};
13use std::{
14 collections::HashMap,
15 future::Future,
16 sync::Arc,
17 time::{Duration, Instant},
18};
19use tokio::sync::{RwLock, broadcast};
20use tracing::{debug, error, info, warn};
21use uuid::Uuid;
22
23#[cfg(feature = "websocket-client")]
24pub mod client;
25pub mod security;
26#[cfg(feature = "http-server")]
27pub mod server;
28
29#[cfg(feature = "websocket-client")]
30pub use client::{PjsWebSocketClient, StreamStats};
31pub use security::SecureWebSocketHandler;
32#[cfg(feature = "http-server")]
33pub use server::{AxumWebSocketTransport, create_websocket_router};
34
35pub(crate) const WRITE_TIMEOUT: Duration = Duration::from_secs(10);
66
67pub(crate) async fn send_with_write_timeout<S, M>(
80 sink: &mut S,
81 message: M,
82 timeout: Duration,
83) -> Result<(), String>
84where
85 S: Sink<M> + Unpin,
86 S::Error: std::fmt::Display,
87{
88 match tokio::time::timeout(timeout, sink.send(message)).await {
89 Ok(Ok(())) => Ok(()),
90 Ok(Err(e)) => Err(format!("send failed: {e}")),
91 Err(_elapsed) => Err(format!("write stalled for {timeout:?}")),
92 }
93}
94
95#[cfg(test)]
96mod write_timeout_tests {
97 use super::*;
98
99 #[tokio::test(start_paused = true)]
100 async fn test_send_with_write_timeout_times_out_on_stalled_sink() {
101 let handle = tokio::spawn(async {
102 let mut sink = futures::sink::unfold((), |_, _item: &str| {
103 futures::future::pending::<Result<(), std::io::Error>>()
104 });
105 send_with_write_timeout(&mut sink, "hello", Duration::from_millis(200)).await
106 });
107
108 tokio::time::advance(Duration::from_millis(201)).await;
109
110 let result = handle.await.expect("task panicked");
111 assert!(
112 result.is_err(),
113 "a write that never completes must time out, not hang forever"
114 );
115 }
116
117 #[tokio::test]
118 async fn test_send_with_write_timeout_succeeds_on_ready_sink() {
119 let mut sink = futures::sink::drain();
120 send_with_write_timeout(&mut sink, "hello", WRITE_TIMEOUT)
121 .await
122 .expect("a sink that accepts immediately must not be treated as stalled");
123 }
124}
125
126#[derive(Debug, Clone, Serialize, Deserialize)]
128#[serde(tag = "type", content = "data")]
129pub enum WsMessage {
130 StreamInit {
132 session_id: String,
134 data: Value,
136 options: StreamOptions,
138 },
139 StreamFrame {
141 session_id: String,
143 frame_id: u32,
145 priority: u8,
147 payload: Value,
149 is_complete: bool,
151 },
152 FrameAck {
154 session_id: String,
156 frame_id: u32,
158 processing_time_ms: u64,
160 },
161 StreamComplete {
163 session_id: String,
165 checksum: String,
167 },
168 Error {
170 session_id: Option<String>,
172 error: String,
174 code: u16,
176 },
177 Ping {
179 timestamp: u64,
181 },
182 Pong {
184 timestamp: u64,
186 },
187}
188
189#[derive(Debug, Clone, Serialize, Deserialize)]
191pub struct StreamOptions {
192 pub max_frame_size: usize,
194 pub client_fps: Option<u32>,
196 pub compression: bool,
198 pub priority_mapping: Option<HashMap<String, u8>>,
200}
201
202impl Default for StreamOptions {
203 fn default() -> Self {
204 Self {
205 max_frame_size: 64 * 1024, client_fps: None, compression: true,
208 priority_mapping: None,
209 }
210 }
211}
212
213#[derive(Debug)]
228pub struct WebSocketStreamSession {
229 pub id: String,
231 pub created_at: Instant,
233 pub options: StreamOptions,
235 pub plan: Vec<StreamFrame>,
237 pub current_frame: u32,
239 pub acknowledged_frames: Vec<u32>,
241 pub client_metrics: ClientMetrics,
243 pub rate_limit_guard: Option<RateLimitGuard>,
245 stream_task: Option<tokio::task::AbortHandle>,
251}
252
253#[derive(Debug, Default)]
255pub struct ClientMetrics {
256 pub average_processing_time_ms: f64,
258 pub frames_acknowledged: u32,
260 pub last_ack_time: Option<Instant>,
262 pub estimated_bandwidth_kbps: Option<f64>,
264 pub connection_rtt_ms: Option<u64>,
266}
267
268impl ClientMetrics {
269 pub fn update_processing_time(&mut self, processing_time_ms: u64) {
271 let new_time = processing_time_ms as f64;
272 if self.frames_acknowledged == 0 {
273 self.average_processing_time_ms = new_time;
274 } else {
275 let alpha = 0.3;
277 self.average_processing_time_ms =
278 alpha * new_time + (1.0 - alpha) * self.average_processing_time_ms;
279 }
280 self.frames_acknowledged += 1;
281 self.last_ack_time = Some(Instant::now());
282 }
283
284 pub fn is_client_slow(&self) -> bool {
286 self.average_processing_time_ms > 100.0 }
288
289 pub fn recommended_frame_delay(&self) -> Duration {
298 if self.is_client_slow() {
299 Duration::from_millis((self.average_processing_time_ms * 0.5) as u64)
300 .min(MAX_ADAPTIVE_FRAME_DELAY)
301 } else {
302 Duration::from_millis(10) }
304 }
305}
306
307const MAX_ADAPTIVE_FRAME_DELAY: Duration = Duration::from_secs(1);
309
310pub trait WebSocketTransport: Send + Sync {
312 type Connection: Send + Sync;
314
315 type StartStreamFuture<'a>: Future<Output = PjsResult<String>> + Send + 'a
317 where
318 Self: 'a;
319
320 type SendFrameFuture<'a>: Future<Output = PjsResult<()>> + Send + 'a
322 where
323 Self: 'a;
324
325 type HandleMessageFuture<'a>: Future<Output = PjsResult<()>> + Send + 'a
327 where
328 Self: 'a;
329
330 type CloseStreamFuture<'a>: Future<Output = PjsResult<()>> + Send + 'a
332 where
333 Self: 'a;
334
335 fn start_stream(
337 &self,
338 connection: Arc<Self::Connection>,
339 data: Value,
340 options: StreamOptions,
341 ) -> Self::StartStreamFuture<'_>;
342
343 fn send_frame(
353 &self,
354 connection: Arc<Self::Connection>,
355 message: WsMessage,
356 ) -> Self::SendFrameFuture<'_>;
357
358 fn handle_message(
360 &self,
361 connection: Arc<Self::Connection>,
362 message: WsMessage,
363 ) -> Self::HandleMessageFuture<'_>;
364
365 fn close_stream(&self, session_id: &str) -> Self::CloseStreamFuture<'_>;
367}
368
369pub struct AdaptiveStreamController {
371 sessions: Arc<RwLock<HashMap<String, WebSocketStreamSession>>>,
372 frame_tx: broadcast::Sender<(String, WsMessage)>,
373}
374
375impl AdaptiveStreamController {
376 pub fn new() -> Self {
378 let (frame_tx, _) = broadcast::channel(1000);
379
380 Self {
381 sessions: Arc::new(RwLock::new(HashMap::new())),
382 frame_tx,
383 }
384 }
385
386 pub async fn create_session(&self, data: Value, options: StreamOptions) -> PjsResult<String> {
388 let session_id = Uuid::new_v4().to_string();
389 let plan = vec![StreamFrame {
390 data: data.clone(),
391 priority: Priority::HIGH,
392 metadata: std::collections::HashMap::new(),
393 }]; let session = WebSocketStreamSession {
396 id: session_id.clone(),
397 created_at: Instant::now(),
398 options,
399 plan,
400 current_frame: 0,
401 acknowledged_frames: Vec::new(),
402 client_metrics: ClientMetrics::default(),
403 rate_limit_guard: None, stream_task: None, };
406
407 self.sessions
408 .write()
409 .await
410 .insert(session_id.clone(), session);
411
412 info!("Created streaming session: {}", session_id);
413 Ok(session_id)
414 }
415
416 pub async fn start_streaming(&self, session_id: &str) -> PjsResult<()> {
418 let mut sessions = self.sessions.write().await;
419 let session = sessions
420 .get_mut(session_id)
421 .ok_or_else(|| PjsError::InvalidSession(session_id.to_string()))?;
422
423 let session_id = session_id.to_string();
425 let frame_tx = self.frame_tx.clone();
426 let plan = session.plan.clone();
427
428 let task_session_id = session_id.clone();
429 let sessions_for_task = self.sessions.clone();
430 let handle = tokio::spawn(async move {
431 if let Err(e) =
432 Self::stream_frames(task_session_id, plan, frame_tx, sessions_for_task).await
433 {
434 error!("Error streaming frames: {}", e);
435 }
436 });
437
438 if let Some(previous) = session.stream_task.replace(handle.abort_handle()) {
444 previous.abort();
445 }
446 tokio::spawn(async move {
447 match handle.await {
448 Ok(()) => {}
449 Err(join_err) if join_err.is_panic() => {
450 error!(
451 "Streaming task panicked for session {}: {}",
452 session_id, join_err
453 );
454 }
455 Err(_) => {} }
457 });
458
459 Ok(())
460 }
461
462 async fn stream_frames(
463 session_id: String,
464 plan: Vec<StreamFrame>, frame_tx: broadcast::Sender<(String, WsMessage)>,
466 sessions: Arc<RwLock<HashMap<String, WebSocketStreamSession>>>,
467 ) -> Result<(), PjsError> {
468 let mut frames_data = Vec::new();
469
470 for (frame_id, frame) in plan.iter().enumerate() {
471 let payload_bytes =
473 serde_json::to_vec(&frame.data).map_err(|e| PjsError::Other(e.to_string()))?;
474 frames_data.push(payload_bytes);
475
476 let ws_message = WsMessage::StreamFrame {
477 session_id: session_id.clone(),
478 frame_id: frame_id as u32,
479 priority: frame.priority.value(),
480 payload: frame.data.clone(),
481 is_complete: frame_id == (plan.len() - 1),
482 };
483
484 if let Err(e) = frame_tx.send((session_id.clone(), ws_message)) {
485 error!("Failed to send frame {}: {}", frame_id, e);
486 break;
487 }
488
489 let delay = sessions
490 .read()
491 .await
492 .get(&session_id)
493 .map(|session| session.client_metrics.recommended_frame_delay())
494 .unwrap_or(Duration::from_millis(10));
495 tokio::time::sleep(delay).await;
496 }
497
498 let complete_message = WsMessage::StreamComplete {
500 session_id: session_id.clone(),
501 checksum: calculate_stream_checksum(&frames_data),
502 };
503
504 let _ = frame_tx.send((session_id, complete_message));
505 Ok(())
506 }
507
508 pub async fn handle_frame_ack(
510 &self,
511 session_id: &str,
512 frame_id: u32,
513 processing_time_ms: u64,
514 ) -> PjsResult<()> {
515 let mut sessions = self.sessions.write().await;
516 let session = sessions
517 .get_mut(session_id)
518 .ok_or_else(|| PjsError::InvalidSession(session_id.to_string()))?;
519
520 session.acknowledged_frames.push(frame_id);
521 session
522 .client_metrics
523 .update_processing_time(processing_time_ms);
524
525 debug!(
526 "Frame {} acknowledged for session {} (processing: {}ms, avg: {:.1}ms)",
527 frame_id,
528 session_id,
529 processing_time_ms,
530 session.client_metrics.average_processing_time_ms
531 );
532
533 if session.client_metrics.is_client_slow() {
534 warn!(
535 "Client {} is processing slowly (avg: {:.1}ms)",
536 session_id, session.client_metrics.average_processing_time_ms
537 );
538 }
539
540 Ok(())
541 }
542
543 pub fn subscribe_frames(&self) -> broadcast::Receiver<(String, WsMessage)> {
545 self.frame_tx.subscribe()
546 }
547
548 pub async fn set_rate_limit_guard(
550 &self,
551 session_id: &str,
552 guard: RateLimitGuard,
553 ) -> PjsResult<()> {
554 let mut sessions = self.sessions.write().await;
555 let session = sessions
556 .get_mut(session_id)
557 .ok_or_else(|| PjsError::InvalidSession(session_id.to_string()))?;
558
559 session.rate_limit_guard = Some(guard);
560 Ok(())
561 }
562
563 pub async fn validate_message(&self, session_id: &str, frame_size: usize) -> PjsResult<()> {
565 let sessions = self.sessions.read().await;
566 let session = sessions
567 .get(session_id)
568 .ok_or_else(|| PjsError::InvalidSession(session_id.to_string()))?;
569
570 if let Some(guard) = &session.rate_limit_guard {
571 guard
572 .check_message(frame_size)
573 .map_err(|e| PjsError::SecurityError(format!("Rate limit violation: {}", e)))?;
574 }
575
576 Ok(())
577 }
578
579 pub async fn remove_session(&self, session_id: &str) -> bool {
599 let mut sessions = self.sessions.write().await;
600 let removed = sessions.remove(session_id);
601 match &removed {
602 Some(session) => {
603 if let Some(abort_handle) = &session.stream_task {
604 abort_handle.abort();
605 }
606 info!("Removed streaming session: {}", session_id);
607 }
608 None => debug!("remove_session called on unknown id: {}", session_id),
609 }
610 removed.is_some()
611 }
612
613 pub async fn cleanup_expired_sessions(&self, max_age: Duration) {
615 let mut sessions = self.sessions.write().await;
616 let now = Instant::now();
617
618 sessions.retain(|id, session| {
619 let expired = now.duration_since(session.created_at) > max_age;
620 if expired {
621 if let Some(abort_handle) = &session.stream_task {
622 abort_handle.abort();
623 }
624 info!("Cleaning up expired session: {}", id);
625 }
626 !expired
627 });
628 }
629}
630
631impl Default for AdaptiveStreamController {
632 fn default() -> Self {
633 Self::new()
634 }
635}
636
637fn calculate_stream_checksum(frames_data: &[Vec<u8>]) -> String {
639 let mut hasher = Sha256::new();
640
641 for frame_data in frames_data {
643 hasher.update(frame_data);
644 }
645
646 hasher.update((frames_data.len() as u64).to_le_bytes());
648
649 let result = hasher.finalize();
650 let hex: String = result.iter().map(|byte| format!("{byte:02x}")).collect();
651 format!("sha256:{hex}")
652}
653
654#[cfg(test)]
655mod tests {
656 use super::*;
657 use serde_json::json;
658
659 #[tokio::test]
660 async fn test_create_session() {
661 let controller = AdaptiveStreamController::new();
662 let data = json!({
663 "critical": {"id": 1, "status": "active"},
664 "details": {"name": "test", "description": "test data"}
665 });
666
667 let session_id = controller
668 .create_session(data, StreamOptions::default())
669 .await
670 .unwrap();
671
672 assert!(!session_id.is_empty());
673
674 let sessions = controller.sessions.read().await;
675 assert!(sessions.contains_key(&session_id));
676 }
677
678 #[tokio::test]
679 async fn test_frame_acknowledgment() {
680 let controller = AdaptiveStreamController::new();
681 let data = json!({"test": "data"});
682
683 let session_id = controller
684 .create_session(data, StreamOptions::default())
685 .await
686 .unwrap();
687
688 controller
689 .handle_frame_ack(&session_id, 0, 50)
690 .await
691 .unwrap();
692
693 let sessions = controller.sessions.read().await;
694 let session = sessions.get(&session_id).unwrap();
695 assert_eq!(session.acknowledged_frames, vec![0]);
696 assert_eq!(session.client_metrics.average_processing_time_ms, 50.0);
697 }
698
699 #[tokio::test]
704 async fn test_remove_session_aborts_streaming_task_before_completion() {
705 let controller = AdaptiveStreamController::new();
706 let session_id = controller
707 .create_session(json!({"test": "data"}), StreamOptions::default())
708 .await
709 .unwrap();
710
711 {
714 let mut sessions = controller.sessions.write().await;
715 let session = sessions.get_mut(&session_id).unwrap();
716 session.plan = (0..200)
717 .map(|_| StreamFrame {
718 data: json!({}),
719 priority: Priority::HIGH,
720 metadata: std::collections::HashMap::new(),
721 })
722 .collect();
723 }
724
725 let mut frames_rx = controller.subscribe_frames();
726
727 controller.start_streaming(&session_id).await.unwrap();
728 assert!(controller.remove_session(&session_id).await);
729
730 let mut saw_complete = false;
734 let mut frame_count = 0;
735 let drain_deadline = tokio::time::Instant::now() + Duration::from_millis(300);
736 while tokio::time::Instant::now() < drain_deadline {
737 match tokio::time::timeout(Duration::from_millis(20), frames_rx.recv()).await {
738 Ok(Ok((_, WsMessage::StreamComplete { .. }))) => {
739 saw_complete = true;
740 break;
741 }
742 Ok(Ok(_)) => frame_count += 1,
743 Ok(Err(_)) => break, Err(_) => {} }
746 }
747
748 assert!(
749 !saw_complete,
750 "streaming task must not run to completion after remove_session aborts it"
751 );
752 assert!(
753 frame_count < 200,
754 "streaming task must stop well short of the full plan once aborted, sent {frame_count} frames"
755 );
756 }
757
758 #[test]
759 fn test_client_metrics() {
760 let mut metrics = ClientMetrics::default();
761
762 metrics.update_processing_time(100);
763 assert_eq!(metrics.average_processing_time_ms, 100.0);
764
765 metrics.update_processing_time(200);
766 assert!((metrics.average_processing_time_ms - 130.0).abs() < 0.1);
768
769 assert!(metrics.is_client_slow());
770 }
771
772 #[test]
777 fn test_recommended_frame_delay_clamps_extreme_processing_time() {
778 let mut metrics = ClientMetrics::default();
779 metrics.update_processing_time(100_000_000_000);
780
781 assert_eq!(
782 metrics.recommended_frame_delay(),
783 MAX_ADAPTIVE_FRAME_DELAY,
784 "delay must be clamped to MAX_ADAPTIVE_FRAME_DELAY, not scale unbounded with client-supplied input"
785 );
786 }
787
788 #[tokio::test]
793 async fn test_stream_frames_completes_promptly_under_malicious_client_metrics() {
794 let controller = AdaptiveStreamController::new();
795 let session_id = controller
796 .create_session(json!({"test": "data"}), StreamOptions::default())
797 .await
798 .unwrap();
799
800 {
801 let mut sessions = controller.sessions.write().await;
802 sessions
803 .get_mut(&session_id)
804 .unwrap()
805 .client_metrics
806 .update_processing_time(100_000_000_000);
807 }
808
809 let mut frames_rx = controller.subscribe_frames();
810 controller.start_streaming(&session_id).await.unwrap();
811
812 let result = tokio::time::timeout(MAX_ADAPTIVE_FRAME_DELAY * 2, async {
813 loop {
814 match frames_rx
815 .recv()
816 .await
817 .expect("channel must not close early")
818 {
819 (sid, WsMessage::StreamComplete { .. }) if sid == session_id => break,
820 _ => continue,
821 }
822 }
823 })
824 .await;
825
826 assert!(
827 result.is_ok(),
828 "stream must complete within 2x MAX_ADAPTIVE_FRAME_DELAY, not stall on malicious client metrics"
829 );
830 }
831
832 #[test]
833 fn test_checksum_calculation() {
834 let empty_frames: Vec<Vec<u8>> = vec![];
836 let checksum = calculate_stream_checksum(&empty_frames);
837 assert!(checksum.starts_with("sha256:"));
838
839 let single_frame = vec![vec![1, 2, 3, 4]];
841 let checksum1 = calculate_stream_checksum(&single_frame);
842 assert!(checksum1.starts_with("sha256:"));
843
844 let multi_frames = vec![vec![1, 2], vec![3, 4], vec![5, 6]];
846 let checksum2 = calculate_stream_checksum(&multi_frames);
847 assert!(checksum2.starts_with("sha256:"));
848
849 let same_frames = vec![vec![1, 2], vec![3, 4], vec![5, 6]];
851 let checksum3 = calculate_stream_checksum(&same_frames);
852 assert_eq!(checksum2, checksum3);
853
854 let diff_frames = vec![vec![1, 2], vec![3, 4], vec![5, 7]]; let checksum4 = calculate_stream_checksum(&diff_frames);
857 assert_ne!(checksum2, checksum4);
858
859 let reordered_frames = vec![vec![3, 4], vec![1, 2], vec![5, 6]];
861 let checksum5 = calculate_stream_checksum(&reordered_frames);
862 assert_ne!(checksum2, checksum5);
863 }
864}