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