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#[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
98pub 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 #[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 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_locked_task(&self.init_timeout_task).await;
265
266 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 self.state
290 .try_transition(ConnectionState::Initiating, ConnectionState::Activating)?;
291
292 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 self.state.ensure(ConnectionState::Activating, |state| {
310 format!("Cannot activate connection in invalid state: {state:?}")
311 })?;
312 }
313
314 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 self.state
333 .try_transition(ConnectionState::Activating, ConnectionState::Ready)?;
334
335 self.namespace.insert_connection(self);
337
338 #[cfg(feature = "tracing")]
339 tracing::debug!(connection_id = self.id, "server connection is ready");
340
341 self.send_packet(&WsIoPacket::new_ready()).await?;
343
344 if let Some(on_ready_handler) = self.namespace.config.on_ready_handler.clone() {
346 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 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 self.state.store(ConnectionState::Closing);
365
366 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 self.namespace.remove_connection(self.id);
375
376 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_locked_task(&self.init_timeout_task).await;
386
387 self.cancel_token.cancel();
389
390 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 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 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 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 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 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 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 self.state
499 .try_transition(ConnectionState::Created, ConnectionState::AwaitingInit)?;
500
501 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 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 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
707static 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 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 let encoded = connection
762 .namespace
763 .config
764 .packet_codec
765 .encode(&WsIoPacket::new_init(None))
766 .unwrap();
767
768 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 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 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 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}