Skip to main content

wsio_server/connection/
mod.rs

1use std::{
2    fmt::{
3        Debug as FmtDebug,
4        Formatter,
5        Result as FmtResult,
6    },
7    panic::AssertUnwindSafe,
8    sync::{
9        Arc,
10        LazyLock,
11        atomic::{
12            AtomicU64,
13            Ordering,
14        },
15    },
16};
17
18use anyhow::{
19    Result,
20    anyhow,
21    bail,
22};
23use futures_util::FutureExt;
24use http::{
25    HeaderMap,
26    Uri,
27};
28use kikiutils::{
29    atomic::enum_cell::AtomicEnumCell,
30    types::fx_collections::FxDashSet,
31};
32use num_enum::{
33    IntoPrimitive,
34    TryFromPrimitive,
35};
36use serde::{
37    Serialize,
38    de::DeserializeOwned,
39};
40use tokio::{
41    select,
42    spawn,
43    sync::{
44        Mutex,
45        mpsc::{
46            Receiver,
47            Sender,
48            channel,
49        },
50    },
51    task::JoinHandle,
52    time::{
53        sleep,
54        timeout,
55    },
56};
57use tokio_tungstenite::tungstenite::Message;
58use tokio_util::sync::CancellationToken;
59
60#[cfg(feature = "connection-extensions")]
61mod extensions;
62
63#[cfg(feature = "connection-extensions")]
64use self::extensions::ConnectionExtensions;
65use crate::{
66    WsIoServer,
67    core::{
68        channel_capacity_from_websocket_config,
69        event::registry::WsIoEventRegistry,
70        packet::{
71            WsIoPacket,
72            WsIoPacketType,
73        },
74        traits::task::spawner::TaskSpawner,
75        types::BoxAsyncUnaryResultHandler,
76        utils::task::abort_locked_task,
77    },
78    namespace::{
79        WsIoServerNamespace,
80        operators::broadcast::WsIoServerNamespaceBroadcastOperator,
81    },
82};
83
84// Enums
85#[repr(u8)]
86#[derive(Debug, Eq, IntoPrimitive, PartialEq, TryFromPrimitive)]
87enum ConnectionState {
88    Activating,
89    AwaitingInit,
90    Closed,
91    Closing,
92    Created,
93    Initiating,
94    Ready,
95}
96
97// Structs
98pub struct WsIoServerConnection {
99    cancel_token: CancellationToken,
100    event_dispatcher_task: Mutex<Option<JoinHandle<()>>>,
101    event_queue_tx: Sender<WsIoPacket>,
102    event_registry: WsIoEventRegistry<WsIoServerConnection>,
103    #[cfg(feature = "connection-extensions")]
104    extensions: ConnectionExtensions,
105    headers: HeaderMap,
106    id: u64,
107    init_timeout_task: Mutex<Option<JoinHandle<()>>>,
108    joined_rooms: FxDashSet<String>,
109    message_tx: Sender<Arc<Message>>,
110    namespace: Arc<WsIoServerNamespace>,
111    on_close_handler: Mutex<Option<BoxAsyncUnaryResultHandler<Self>>>,
112    request_uri: Uri,
113    state: AtomicEnumCell<ConnectionState>,
114}
115
116impl FmtDebug for WsIoServerConnection {
117    fn fmt(&self, f: &mut Formatter<'_>) -> FmtResult {
118        let event_dispatcher_task = match self.event_dispatcher_task.try_lock() {
119            Ok(task) => {
120                if task.is_some() {
121                    "<task>"
122                } else {
123                    "<none>"
124                }
125            },
126            Err(_) => "<locked>",
127        };
128
129        let init_timeout_task = match self.init_timeout_task.try_lock() {
130            Ok(task) => {
131                if task.is_some() {
132                    "<task>"
133                } else {
134                    "<none>"
135                }
136            },
137            Err(_) => "<locked>",
138        };
139
140        let on_close_handler = match self.on_close_handler.try_lock() {
141            Ok(handler) => {
142                if handler.is_some() {
143                    "<handler>"
144                } else {
145                    "<none>"
146                }
147            },
148            Err(_) => "<locked>",
149        };
150
151        let mut debug = f.debug_struct("WsIoServerConnection");
152        debug
153            .field("id", &self.id)
154            .field("state", &self.state)
155            .field("request_uri", &self.request_uri)
156            .field("headers", &self.headers)
157            .field("joined_rooms_len", &self.joined_rooms.len())
158            .field("message_tx", &self.message_tx)
159            .field("event_queue_tx", &self.event_queue_tx)
160            .field("event_dispatcher_task", &event_dispatcher_task)
161            .field("cancel_token", &"<cancel_token>")
162            .field("namespace", &"<namespace>")
163            .field("event_registry", &self.event_registry)
164            .field("init_timeout_task", &init_timeout_task)
165            .field("on_close_handler", &on_close_handler);
166
167        #[cfg(feature = "connection-extensions")]
168        debug.field("extensions", &self.extensions);
169
170        debug.finish()
171    }
172}
173
174impl TaskSpawner for WsIoServerConnection {
175    #[inline]
176    fn cancel_token(&self) -> CancellationToken {
177        self.cancel_token.clone()
178    }
179}
180
181impl WsIoServerConnection {
182    #[inline]
183    pub(crate) fn new(
184        headers: HeaderMap,
185        namespace: Arc<WsIoServerNamespace>,
186        request_uri: Uri,
187    ) -> (Arc<Self>, Receiver<Arc<Message>>, Receiver<WsIoPacket>) {
188        let channel_capacity = channel_capacity_from_websocket_config(&namespace.config.websocket_config);
189        let (event_queue_tx, event_queue_rx) = channel(channel_capacity);
190        let (message_tx, message_rx) = channel(channel_capacity);
191        let id = NEXT_CONNECTION_ID.fetch_add(1, Ordering::Relaxed);
192
193        #[cfg(feature = "tracing")]
194        tracing::debug!(
195            connection_id = id,
196            namespace = %namespace.path(),
197            request_path = request_uri.path(),
198            "creating server connection"
199        );
200
201        (
202            Arc::new(Self {
203                cancel_token: CancellationToken::new(),
204                event_dispatcher_task: Mutex::new(None),
205                event_queue_tx,
206                event_registry: WsIoEventRegistry::new(),
207                #[cfg(feature = "connection-extensions")]
208                extensions: ConnectionExtensions::new(),
209                headers,
210                id,
211                init_timeout_task: Mutex::new(None),
212                joined_rooms: FxDashSet::default(),
213                message_tx,
214                namespace,
215                on_close_handler: Mutex::new(None),
216                request_uri,
217                state: AtomicEnumCell::new(ConnectionState::Created),
218            }),
219            message_rx,
220            event_queue_rx,
221        )
222    }
223
224    // Private methods
225    #[inline]
226    async fn handle_event_packet(self: &Arc<Self>, packet: WsIoPacket) -> Result<()> {
227        #[cfg(feature = "tracing")]
228        tracing::trace!(
229            connection_id = self.id,
230            event = packet.key.as_deref().unwrap_or_default(),
231            has_data = packet.data.is_some(),
232            "received client event packet"
233        );
234
235        let cancel_token = self.cancel_token();
236        select! {
237            () = cancel_token.cancelled() => Ok(()),
238            result = self.event_queue_tx.send(packet) => result.map_err(|_| anyhow!("event dispatcher is closed")),
239        }
240    }
241
242    async fn handle_init_packet(self: &Arc<Self>, packet_data: Option<&[u8]>) -> Result<()> {
243        // Verify current state; only valid from AwaitingInit → Initiating
244        let state = self.state.get();
245        if state == ConnectionState::AwaitingInit {
246            self.state.try_transition(state, ConnectionState::Initiating)?;
247        } else {
248            #[cfg(feature = "tracing")]
249            tracing::debug!(
250                connection_id = self.id,
251                ?state,
252                "received init packet in invalid server connection state"
253            );
254
255            bail!("Received init packet in invalid state: {state:?}");
256        }
257
258        #[cfg(feature = "tracing")]
259        tracing::debug!(connection_id = self.id, "received client init packet");
260
261        // Abort init-timeout task
262        abort_locked_task(&self.init_timeout_task).await;
263
264        // Invoke init_response_handler with timeout protection if configured
265        if let Some(init_response_handler) = &self.namespace.config.init_response_handler {
266            match timeout(
267                self.namespace.config.init_response_handler_timeout,
268                init_response_handler(self.clone(), packet_data, &self.namespace.config.packet_codec),
269            )
270            .await
271            {
272                Ok(result) => result?,
273                Err(err) => {
274                    #[cfg(feature = "tracing")]
275                    tracing::warn!(
276                        connection_id = self.id,
277                        error = %err,
278                        "server init response handler timed out"
279                    );
280
281                    return Err(err.into());
282                },
283            }
284        }
285
286        // Activate connection
287        self.state
288            .try_transition(ConnectionState::Initiating, ConnectionState::Activating)?;
289
290        // Invoke middleware with timeout protection if configured
291        if let Some(middleware) = &self.namespace.config.middleware {
292            match timeout(
293                self.namespace.config.middleware_execution_timeout,
294                middleware(self.clone()),
295            )
296            .await
297            {
298                Ok(result) => result?,
299                Err(err) => {
300                    #[cfg(feature = "tracing")]
301                    tracing::warn!(connection_id = self.id, error = %err, "server middleware timed out");
302                    return Err(err.into());
303                },
304            }
305
306            // Ensure connection is still in Activating state
307            self.state.ensure(ConnectionState::Activating, |state| {
308                format!("Cannot activate connection in invalid state: {state:?}")
309            })?;
310        }
311
312        // Invoke on_connect_handler with timeout protection if configured
313        if let Some(on_connect_handler) = &self.namespace.config.on_connect_handler {
314            match timeout(
315                self.namespace.config.on_connect_handler_timeout,
316                on_connect_handler(self.clone()),
317            )
318            .await
319            {
320                Ok(result) => result?,
321                Err(err) => {
322                    #[cfg(feature = "tracing")]
323                    tracing::warn!(connection_id = self.id, error = %err, "server on-connect handler timed out");
324                    return Err(err.into());
325                },
326            }
327        }
328
329        // Transition state to Ready
330        self.state
331            .try_transition(ConnectionState::Activating, ConnectionState::Ready)?;
332
333        // Insert connection into namespace
334        self.namespace.insert_connection(self);
335
336        #[cfg(feature = "tracing")]
337        tracing::debug!(connection_id = self.id, "server connection is ready");
338
339        // Send ready packet
340        self.send_packet(&WsIoPacket::new_ready()).await?;
341
342        // Invoke on_ready_handler if configured
343        if let Some(on_ready_handler) = self.namespace.config.on_ready_handler.clone() {
344            // Run handler asynchronously in a detached task
345            self.spawn_task(on_ready_handler(self.clone()));
346        }
347
348        Ok(())
349    }
350
351    async fn send_packet(&self, packet: &WsIoPacket) -> Result<()> {
352        self.send_message(self.namespace.encode_packet_to_message(packet)?)
353            .await
354    }
355
356    // Protected methods
357    pub(super) async fn cleanup(self: &Arc<Self>) {
358        #[cfg(feature = "tracing")]
359        tracing::debug!(connection_id = self.id, "cleaning up server connection");
360
361        // Set connection state to Closing
362        self.state.store(ConnectionState::Closing);
363
364        // Stop event dispatch before mutating namespace and room membership.
365        let event_dispatcher_task = self.event_dispatcher_task.lock().await.take();
366        if let Some(event_dispatcher_task) = event_dispatcher_task {
367            event_dispatcher_task.abort();
368            let _ = event_dispatcher_task.await;
369        }
370
371        // Remove connection from namespace
372        self.namespace.remove_connection(self.id);
373
374        // Leave all joined rooms
375        let joined_rooms = self.joined_rooms.iter().map(|entry| entry.clone()).collect::<Vec<_>>();
376        for room_name in &joined_rooms {
377            self.namespace.remove_connection_id_from_room(room_name, self.id);
378        }
379
380        self.joined_rooms.clear();
381
382        // Abort init-timeout task
383        abort_locked_task(&self.init_timeout_task).await;
384
385        // Cancel all ongoing operations via cancel token
386        self.cancel_token.cancel();
387
388        // Invoke on_close_handler with timeout protection if configured
389        if let Some(on_close_handler) = self.on_close_handler.lock().await.take()
390            && let Err(_err) = timeout(
391                self.namespace.config.on_close_handler_timeout,
392                on_close_handler(self.clone()),
393            )
394            .await
395        {
396            #[cfg(feature = "tracing")]
397            tracing::warn!(connection_id = self.id, error = %_err, "server close handler timed out");
398        }
399
400        // Set connection state to Closed
401        self.state.store(ConnectionState::Closed);
402
403        #[cfg(feature = "tracing")]
404        tracing::debug!(connection_id = self.id, "server connection closed");
405    }
406
407    #[inline]
408    pub(super) fn close(&self) {
409        // Skip if connection is already Closing or Closed, otherwise set connection state to Closing
410        match self.state.get() {
411            ConnectionState::Closed | ConnectionState::Closing => return,
412            _state => {
413                #[cfg(feature = "tracing")]
414                tracing::debug!(connection_id = self.id, state = ?_state, "closing server connection");
415                self.state.store(ConnectionState::Closing);
416            },
417        }
418
419        // Send websocket close frame to initiate graceful shutdown
420        let _ = self.message_tx.try_send(Arc::new(Message::Close(None)));
421    }
422
423    pub(super) async fn emit_event_message(&self, message: Arc<Message>) -> Result<()> {
424        self.state.ensure(ConnectionState::Ready, |state| {
425            format!("Cannot emit in invalid state: {state:?}")
426        })?;
427
428        self.send_message(message).await
429    }
430
431    pub(super) async fn handle_incoming_packet(self: &Arc<Self>, encoded_packet: &[u8]) -> Result<()> {
432        // TODO: lazy load
433        let packet = match self.namespace.config.packet_codec.decode(encoded_packet) {
434            Ok(packet) => packet,
435            Err(err) => {
436                #[cfg(feature = "tracing")]
437                tracing::debug!(connection_id = self.id, error = %err, "failed to decode client packet");
438                return Err(err);
439            },
440        };
441
442        match &packet.r#type {
443            WsIoPacketType::Event => {
444                if self.is_ready() {
445                    return self.handle_event_packet(packet).await;
446                }
447
448                Ok(())
449            },
450            WsIoPacketType::Init => self.handle_init_packet(packet.data.as_deref()).await,
451            _ => Ok(()),
452        }
453    }
454
455    pub(super) async fn init(self: &Arc<Self>) -> Result<()> {
456        // Verify current state; only valid Created
457        self.state.ensure(ConnectionState::Created, |state| {
458            format!("Cannot init connection in invalid state: {state:?}")
459        })?;
460
461        #[cfg(feature = "tracing")]
462        tracing::debug!(connection_id = self.id, "initializing server connection");
463
464        // Generate init request data if init request handler is configured
465        let init_request_data = if let Some(init_request_handler) = &self.namespace.config.init_request_handler {
466            match timeout(
467                self.namespace.config.init_request_handler_timeout,
468                init_request_handler(self.clone(), &self.namespace.config.packet_codec),
469            )
470            .await
471            {
472                Ok(result) => result?,
473                Err(err) => {
474                    #[cfg(feature = "tracing")]
475                    tracing::warn!(
476                        connection_id = self.id,
477                        error = %err,
478                        "server init request handler timed out"
479                    );
480
481                    return Err(err.into());
482                },
483            }
484        } else {
485            None
486        };
487
488        // Transition state to AwaitingInit
489        self.state
490            .try_transition(ConnectionState::Created, ConnectionState::AwaitingInit)?;
491
492        // Spawn init-response-timeout watchdog to close connection if init not received in time
493        let connection = self.clone();
494        *self.init_timeout_task.lock().await = Some(spawn(async move {
495            sleep(connection.namespace.config.init_response_timeout).await;
496            if connection.state.is(ConnectionState::AwaitingInit) {
497                #[cfg(feature = "tracing")]
498                tracing::warn!(
499                    connection_id = connection.id,
500                    "timed out waiting for client init response packet"
501                );
502
503                connection.close();
504            }
505        }));
506
507        // Send init packet
508        self.send_packet(&WsIoPacket::new_init(init_request_data)).await
509    }
510
511    pub(super) async fn send_message(&self, message: Arc<Message>) -> Result<()> {
512        Ok(self.message_tx.send(message).await?)
513    }
514
515    pub(super) async fn start_event_dispatcher(self: &Arc<Self>, mut event_queue_rx: Receiver<WsIoPacket>) {
516        let cancel_token = self.cancel_token();
517        let connection = self.clone();
518        *self.event_dispatcher_task.lock().await = Some(spawn(async move {
519            let dispatcher = async {
520                loop {
521                    let event_packet = select! {
522                        () = cancel_token.cancelled() => break,
523                        event_packet = event_queue_rx.recv() => event_packet,
524                    };
525
526                    let Some(event_packet) = event_packet else {
527                        break;
528                    };
529
530                    let Some(event) = event_packet.key else {
531                        continue;
532                    };
533
534                    if let Err(_err) = connection
535                        .event_registry
536                        .dispatch_event_packet(
537                            connection.clone(),
538                            event,
539                            &connection.namespace.config.packet_codec,
540                            event_packet.data,
541                            &cancel_token,
542                        )
543                        .await
544                    {
545                        #[cfg(feature = "tracing")]
546                        tracing::warn!(
547                            connection_id = connection.id,
548                            error = %_err,
549                            "server event dispatcher failed; closing connection"
550                        );
551
552                        connection.close();
553                        break;
554                    }
555                }
556            };
557
558            if AssertUnwindSafe(dispatcher).catch_unwind().await.is_err() {
559                #[cfg(feature = "tracing")]
560                tracing::error!(
561                    connection_id = connection.id,
562                    "server event dispatcher panicked; closing connection"
563                );
564
565                connection.close();
566            }
567        }));
568    }
569
570    // Public methods
571    pub async fn disconnect(&self) {
572        #[cfg(feature = "tracing")]
573        tracing::debug!(connection_id = self.id, "disconnecting server connection");
574        let _ = self.send_packet(&WsIoPacket::new_disconnect()).await;
575        self.close();
576    }
577
578    pub async fn emit<D: Serialize>(&self, event: impl AsRef<str>, data: Option<&D>) -> Result<()> {
579        self.emit_event_message(
580            self.namespace.encode_packet_to_message(&WsIoPacket::new_event(
581                event.as_ref(),
582                data.map(|data| self.namespace.config.packet_codec.encode_data(data))
583                    .transpose()?,
584            ))?,
585        )
586        .await
587    }
588
589    #[inline]
590    pub fn except(
591        self: &Arc<Self>,
592        room_names: impl IntoIterator<Item = impl Into<String>>,
593    ) -> WsIoServerNamespaceBroadcastOperator {
594        self.namespace.except(room_names).except_connection_ids([self.id])
595    }
596
597    #[cfg(feature = "connection-extensions")]
598    #[inline]
599    pub fn extensions(&self) -> &ConnectionExtensions {
600        &self.extensions
601    }
602
603    #[inline]
604    pub fn headers(&self) -> &HeaderMap {
605        &self.headers
606    }
607
608    #[inline]
609    pub fn id(&self) -> u64 {
610        self.id
611    }
612
613    #[inline]
614    pub fn is_ready(&self) -> bool {
615        self.state.is(ConnectionState::Ready)
616    }
617
618    #[inline]
619    pub fn join(self: &Arc<Self>, room_names: impl IntoIterator<Item = impl Into<String>>) {
620        for room_name in room_names {
621            let room_name = room_name.into();
622            self.namespace.add_connection_id_to_room(&room_name, self.id);
623
624            #[cfg(feature = "tracing")]
625            tracing::trace!(connection_id = self.id, room = %room_name, "connection joined room");
626            self.joined_rooms.insert(room_name);
627        }
628    }
629
630    #[inline]
631    pub fn leave(self: &Arc<Self>, room_names: impl IntoIterator<Item = impl Into<String>>) {
632        for room_name in room_names {
633            let room_name = room_name.into();
634            self.namespace.remove_connection_id_from_room(&room_name, self.id);
635
636            self.joined_rooms.remove(&room_name);
637
638            #[cfg(feature = "tracing")]
639            tracing::trace!(connection_id = self.id, room = %room_name, "connection left room");
640        }
641    }
642
643    #[inline]
644    pub fn namespace(&self) -> Arc<WsIoServerNamespace> {
645        self.namespace.clone()
646    }
647
648    #[inline]
649    pub fn off(&self, event: impl AsRef<str>) {
650        self.event_registry.off(event.as_ref());
651    }
652
653    #[inline]
654    pub fn off_by_handler_id(&self, event: impl AsRef<str>, handler_id: u32) {
655        self.event_registry.off_by_handler_id(event.as_ref(), handler_id);
656    }
657
658    #[inline]
659    pub fn on<H, Fut, D>(&self, event: impl AsRef<str>, handler: H) -> u32
660    where
661        H: Fn(Arc<WsIoServerConnection>, Arc<D>) -> Fut + Send + Sync + 'static,
662        Fut: Future<Output = Result<()>> + Send + 'static,
663        D: DeserializeOwned + Send + Sync + 'static,
664    {
665        self.event_registry.on(event.as_ref(), handler)
666    }
667
668    pub async fn on_close<H, Fut>(&self, handler: H)
669    where
670        H: Fn(Arc<WsIoServerConnection>) -> Fut + Send + Sync + 'static,
671        Fut: Future<Output = Result<()>> + Send + 'static,
672    {
673        *self.on_close_handler.lock().await = Some(Box::new(move |connection| Box::pin(handler(connection))));
674    }
675
676    #[inline]
677    pub fn request_uri(&self) -> &Uri {
678        &self.request_uri
679    }
680
681    #[inline]
682    pub fn server(&self) -> WsIoServer {
683        self.namespace.server()
684    }
685
686    #[inline]
687    pub fn to(
688        self: &Arc<Self>,
689        room_names: impl IntoIterator<Item = impl Into<String>>,
690    ) -> WsIoServerNamespaceBroadcastOperator {
691        self.namespace.to(room_names).except_connection_ids([self.id])
692    }
693}
694
695// Constants/Statics
696static NEXT_CONNECTION_ID: LazyLock<AtomicU64> = LazyLock::new(|| AtomicU64::new(0));
697
698#[cfg(test)]
699mod tests {
700    use std::time::Duration;
701
702    use http::{
703        HeaderMap,
704        Uri,
705    };
706    use tokio::{
707        sync::mpsc::unbounded_channel,
708        time::{
709            sleep,
710            timeout,
711        },
712    };
713
714    use super::*;
715
716    fn create_test_connection() -> Arc<WsIoServerConnection> {
717        let server = Arc::new(WsIoServer::builder().build());
718        let namespace = server.new_namespace_builder("/socket").register().unwrap();
719        let (connection, _rx, _event_rx) =
720            WsIoServerConnection::new(HeaderMap::new(), namespace, Uri::from_static("http://localhost"));
721
722        connection
723    }
724
725    fn create_test_connection_with_event_queue_rx() -> (Arc<WsIoServerConnection>, Receiver<WsIoPacket>) {
726        let server = Arc::new(WsIoServer::builder().build());
727        let namespace = server.new_namespace_builder("/socket").register().unwrap();
728        let (connection, _rx, event_queue_rx) =
729            WsIoServerConnection::new(HeaderMap::new(), namespace, Uri::from_static("http://localhost"));
730
731        (connection, event_queue_rx)
732    }
733
734    #[tokio::test]
735    async fn test_handle_incoming_packet_decode_error() {
736        let connection = create_test_connection();
737        let garbage_data = b"obviously not valid messagepack";
738        // Should seamlessly return a Result::Err, not panic
739        let result = connection.handle_incoming_packet(garbage_data).await;
740        assert!(result.is_err(), "Decoding garbage payload should trigger an error");
741    }
742
743    #[tokio::test]
744    async fn test_handle_init_packet_in_invalid_state() {
745        let connection = create_test_connection();
746        assert_eq!(connection.state.get(), ConnectionState::Created);
747
748        // Sending an init packet when the connection is merely `Created` (not yet `AwaitingInit`) should throw an error
749        let packet_codec = connection.namespace.config.packet_codec;
750        let encoded = packet_codec.encode(&WsIoPacket::new_init(None)).unwrap();
751
752        // This simulates a manual client Init push before server starts the handshake buffer
753        let result = connection.handle_incoming_packet(&encoded).await;
754        assert!(
755            result.is_err(),
756            "Should error because state is Created, not AwaitingInit"
757        );
758
759        assert!(result.unwrap_err().to_string().contains("invalid state"));
760    }
761
762    #[tokio::test]
763    async fn test_handle_event_packet_rejects_missing_or_empty_key() {
764        let connection = create_test_connection();
765
766        // Force the connection into the Ready state so it accepts Event packets
767        connection.state.store(ConnectionState::Ready);
768
769        let packet_codec = connection.namespace.config.packet_codec;
770
771        for key in [None, Some("")] {
772            let encoded = packet_codec
773                .encode(&WsIoPacket::new(WsIoPacketType::Event, key, None))
774                .unwrap();
775
776            let result = connection.handle_incoming_packet(&encoded).await;
777            assert!(result.is_err(), "Should reject an invalid event key");
778            assert_eq!(result.unwrap_err().to_string(), "Event packet missing key");
779        }
780    }
781
782    #[tokio::test]
783    async fn test_event_dispatcher_preserves_packet_order() {
784        let (connection, event_queue_rx) = create_test_connection_with_event_queue_rx();
785        connection.state.store(ConnectionState::Ready);
786
787        let (handled_tx, mut handled_rx) = unbounded_channel();
788        connection.on("ordered", move |_connection, payload: Arc<String>| {
789            let handled_tx = handled_tx.clone();
790            async move {
791                handled_tx.send(format!("start:{payload}")).unwrap();
792                if payload.as_str() == "first" {
793                    sleep(Duration::from_millis(25)).await;
794                }
795
796                handled_tx.send(format!("end:{payload}")).unwrap();
797                Ok(())
798            }
799        });
800
801        connection.start_event_dispatcher(event_queue_rx).await;
802        let packet_codec = connection.namespace.config.packet_codec;
803        for payload in ["first", "second"] {
804            let packet_data = packet_codec.encode_data(&payload).unwrap();
805            let encoded_packet = packet_codec
806                .encode(&WsIoPacket::new_event("ordered", Some(packet_data)))
807                .unwrap();
808
809            connection.handle_incoming_packet(&encoded_packet).await.unwrap();
810        }
811
812        let mut handled = Vec::with_capacity(4);
813        for _ in 0..4 {
814            handled.push(
815                timeout(Duration::from_secs(1), handled_rx.recv())
816                    .await
817                    .unwrap()
818                    .unwrap(),
819            );
820        }
821
822        assert_eq!(handled, ["start:first", "end:first", "start:second", "end:second"]);
823
824        connection.cleanup().await;
825    }
826
827    #[tokio::test]
828    async fn test_connection_close_state_transitions() {
829        let connection = create_test_connection();
830        assert_eq!(connection.state.get(), ConnectionState::Created);
831
832        connection.close();
833        assert_eq!(connection.state.get(), ConnectionState::Closing);
834
835        // Calling close again when Closing shouldn't alter anything
836        connection.close();
837        assert_eq!(connection.state.get(), ConnectionState::Closing);
838    }
839
840    #[tokio::test]
841    async fn test_connection_cleanup() {
842        let connection = create_test_connection();
843        let namespace = connection.namespace();
844
845        // Insert connection manually for test
846        namespace.insert_connection(&connection);
847        assert_eq!(namespace.connection_count(), 1);
848
849        connection.join(["room_a", "room_b"]);
850        assert!(connection.joined_rooms.contains("room_a"));
851
852        connection.cleanup().await;
853
854        assert_eq!(connection.state.get(), ConnectionState::Closed);
855        assert!(connection.joined_rooms.is_empty());
856        assert_eq!(namespace.connection_count(), 0);
857    }
858}