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#[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
97pub 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 #[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 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_locked_task(&self.init_timeout_task).await;
263
264 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 self.state
288 .try_transition(ConnectionState::Initiating, ConnectionState::Activating)?;
289
290 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 self.state.ensure(ConnectionState::Activating, |state| {
308 format!("Cannot activate connection in invalid state: {state:?}")
309 })?;
310 }
311
312 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 self.state
331 .try_transition(ConnectionState::Activating, ConnectionState::Ready)?;
332
333 self.namespace.insert_connection(self);
335
336 #[cfg(feature = "tracing")]
337 tracing::debug!(connection_id = self.id, "server connection is ready");
338
339 self.send_packet(&WsIoPacket::new_ready()).await?;
341
342 if let Some(on_ready_handler) = self.namespace.config.on_ready_handler.clone() {
344 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 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 self.state.store(ConnectionState::Closing);
363
364 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 self.namespace.remove_connection(self.id);
373
374 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_locked_task(&self.init_timeout_task).await;
384
385 self.cancel_token.cancel();
387
388 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 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 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 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 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 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 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 self.state
490 .try_transition(ConnectionState::Created, ConnectionState::AwaitingInit)?;
491
492 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 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 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
695static 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 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 let packet_codec = connection.namespace.config.packet_codec;
750 let encoded = packet_codec.encode(&WsIoPacket::new_init(None)).unwrap();
751
752 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 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 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 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}