1use std::collections::HashMap;
40use std::sync::Arc;
41
42use parking_lot::Mutex;
43use thiserror::Error;
44
45#[derive(Debug, Error)]
53pub enum GatewayError {
54 #[error("客户端未找到: {0}")]
56 ClientNotFound(String),
57 #[error("群组未找到: {0}")]
59 GroupNotFound(String),
60 #[error("发送失败: {0}")]
62 SendFailed(String),
63 #[error("无效的 client_id: {0}")]
65 InvalidClientId(String),
66 #[error("传输错误: {0}")]
68 Transport(String),
69 #[error("序列化失败: {0}")]
71 Serialize(String),
72}
73
74pub type ClientId = String;
80
81#[derive(Debug, Clone)]
105pub struct GatewayConfig {
106 pub register_address: String,
108 pub heartbeat_interval: u64,
110 pub default_group: Option<String>,
112}
113
114impl Default for GatewayConfig {
115 fn default() -> Self {
116 Self {
117 register_address: "127.0.0.1:1238".to_string(),
118 heartbeat_interval: 55,
119 default_group: None,
120 }
121 }
122}
123
124impl GatewayConfig {
125 pub fn new(register_address: impl Into<String>) -> Self {
131 Self {
132 register_address: register_address.into(),
133 heartbeat_interval: 55,
134 default_group: None,
135 }
136 }
137
138 pub fn with_heartbeat_interval(mut self, interval: u64) -> Self {
140 self.heartbeat_interval = interval;
141 self
142 }
143
144 pub fn with_default_group(mut self, group: impl Into<String>) -> Self {
146 self.default_group = Some(group.into());
147 self
148 }
149}
150
151pub trait GatewayTransport: Send + Sync {
164 fn send_to_client(&self, client_id: &str, message: &str) -> Result<(), GatewayError>;
175
176 fn send_to_clients(&self, client_ids: &[String], message: &str) -> Result<(), GatewayError>;
187
188 fn send_to_all(&self, message: &str) -> Result<(), GatewayError>;
194
195 fn send_to_group(&self, group: &str, message: &str) -> Result<(), GatewayError>;
202
203 fn get_all_client_ids(&self) -> Result<Vec<String>, GatewayError>;
205
206 fn get_client_id_list_by_group(&self, group: &str) -> Result<Vec<String>, GatewayError>;
212
213 fn get_groups_by_client_id(&self, client_id: &str) -> Result<Vec<String>, GatewayError>;
223
224 fn join_group(&self, client_id: &str, group: &str) -> Result<(), GatewayError>;
235
236 fn leave_group(&self, client_id: &str, group: &str) -> Result<(), GatewayError>;
247
248 fn ungroup(&self, group: &str) -> Result<(), GatewayError>;
260
261 fn is_online(&self, client_id: &str) -> Result<bool, GatewayError>;
263
264 fn get_client_count(&self) -> Result<usize, GatewayError>;
266
267 fn get_client_count_by_group(&self, group: &str) -> Result<usize, GatewayError>;
273
274 fn close_client(&self, client_id: &str) -> Result<(), GatewayError>;
284}
285
286#[derive(Debug, Default)]
316pub struct MemoryGatewayTransport {
317 state: Mutex<GatewayState>,
319}
320
321#[derive(Debug, Default)]
323struct GatewayState {
324 client_messages: HashMap<ClientId, Vec<String>>,
326 client_groups: HashMap<ClientId, Vec<String>>,
328 group_clients: HashMap<String, Vec<ClientId>>,
330}
331
332impl MemoryGatewayTransport {
333 pub fn new() -> Self {
335 Self::default()
336 }
337
338 pub fn register_client(&self, client_id: &str) {
347 self.state
348 .lock()
349 .client_messages
350 .entry(client_id.to_string())
351 .or_default();
352 }
353
354 pub fn client_messages(&self, client_id: &str) -> Vec<String> {
362 self.state
363 .lock()
364 .client_messages
365 .get(client_id)
366 .cloned()
367 .unwrap_or_default()
368 }
369}
370
371impl GatewayTransport for MemoryGatewayTransport {
372 fn send_to_client(&self, client_id: &str, message: &str) -> Result<(), GatewayError> {
373 if client_id.is_empty() {
374 return Err(GatewayError::InvalidClientId(client_id.to_string()));
375 }
376 let mut state = self.state.lock();
377 match state.client_messages.get_mut(client_id) {
378 Some(messages) => {
379 messages.push(message.to_string());
380 Ok(())
381 }
382 None => Err(GatewayError::ClientNotFound(client_id.to_string())),
383 }
384 }
385
386 fn send_to_clients(&self, client_ids: &[String], message: &str) -> Result<(), GatewayError> {
387 let mut state = self.state.lock();
388 for client_id in client_ids {
390 if client_id.is_empty() {
391 return Err(GatewayError::InvalidClientId(client_id.to_string()));
392 }
393 if !state.client_messages.contains_key(client_id) {
394 return Err(GatewayError::ClientNotFound(client_id.to_string()));
395 }
396 }
397 for client_id in client_ids {
398 if let Some(messages) = state.client_messages.get_mut(client_id) {
399 messages.push(message.to_string());
400 }
401 }
402 Ok(())
403 }
404
405 fn send_to_all(&self, message: &str) -> Result<(), GatewayError> {
406 let mut state = self.state.lock();
407 for messages in state.client_messages.values_mut() {
408 messages.push(message.to_string());
409 }
410 Ok(())
411 }
412
413 fn send_to_group(&self, group: &str, message: &str) -> Result<(), GatewayError> {
414 let mut state = self.state.lock();
415 if let Some(client_ids) = state.group_clients.get(group) {
417 let client_ids = client_ids.clone();
418 for client_id in client_ids {
419 if let Some(messages) = state.client_messages.get_mut(&client_id) {
420 messages.push(message.to_string());
421 }
422 }
423 }
424 Ok(())
425 }
426
427 fn get_all_client_ids(&self) -> Result<Vec<String>, GatewayError> {
428 let state = self.state.lock();
429 Ok(state.client_messages.keys().cloned().collect())
430 }
431
432 fn get_client_id_list_by_group(&self, group: &str) -> Result<Vec<String>, GatewayError> {
433 let state = self.state.lock();
434 Ok(state.group_clients.get(group).cloned().unwrap_or_default())
435 }
436
437 fn get_groups_by_client_id(&self, client_id: &str) -> Result<Vec<String>, GatewayError> {
438 if client_id.is_empty() {
439 return Err(GatewayError::InvalidClientId(client_id.to_string()));
440 }
441 let state = self.state.lock();
442 if !state.client_messages.contains_key(client_id) {
443 return Err(GatewayError::ClientNotFound(client_id.to_string()));
444 }
445 Ok(state
446 .client_groups
447 .get(client_id)
448 .cloned()
449 .unwrap_or_default())
450 }
451
452 fn join_group(&self, client_id: &str, group: &str) -> Result<(), GatewayError> {
453 if client_id.is_empty() {
454 return Err(GatewayError::InvalidClientId(client_id.to_string()));
455 }
456 let mut state = self.state.lock();
457 if !state.client_messages.contains_key(client_id) {
458 return Err(GatewayError::ClientNotFound(client_id.to_string()));
459 }
460 let client_groups = state
462 .client_groups
463 .entry(client_id.to_string())
464 .or_default();
465 if !client_groups.iter().any(|g| g == group) {
466 client_groups.push(group.to_string());
467 }
468 let group_clients = state.group_clients.entry(group.to_string()).or_default();
470 if !group_clients.iter().any(|c| c == client_id) {
471 group_clients.push(client_id.to_string());
472 }
473 Ok(())
474 }
475
476 fn leave_group(&self, client_id: &str, group: &str) -> Result<(), GatewayError> {
477 if client_id.is_empty() {
478 return Err(GatewayError::InvalidClientId(client_id.to_string()));
479 }
480 let mut state = self.state.lock();
481 if !state.client_messages.contains_key(client_id) {
482 return Err(GatewayError::ClientNotFound(client_id.to_string()));
483 }
484 let group_exists = state.group_clients.contains_key(group);
485 if !group_exists {
486 return Err(GatewayError::GroupNotFound(group.to_string()));
487 }
488 if let Some(groups) = state.client_groups.get_mut(client_id) {
490 groups.retain(|g| g != group);
491 }
492 if let Some(client_ids) = state.group_clients.get_mut(group) {
494 client_ids.retain(|c| c != client_id);
495 }
496 Ok(())
497 }
498
499 fn ungroup(&self, group: &str) -> Result<(), GatewayError> {
500 let mut state = self.state.lock();
501 if !state.group_clients.contains_key(group) {
502 return Err(GatewayError::GroupNotFound(group.to_string()));
503 }
504 state.group_clients.remove(group);
506 for groups in state.client_groups.values_mut() {
508 groups.retain(|g| g != group);
509 }
510 Ok(())
511 }
512
513 fn is_online(&self, client_id: &str) -> Result<bool, GatewayError> {
514 let state = self.state.lock();
515 Ok(state.client_messages.contains_key(client_id))
516 }
517
518 fn get_client_count(&self) -> Result<usize, GatewayError> {
519 let state = self.state.lock();
520 Ok(state.client_messages.len())
521 }
522
523 fn get_client_count_by_group(&self, group: &str) -> Result<usize, GatewayError> {
524 let state = self.state.lock();
525 Ok(state.group_clients.get(group).map(|v| v.len()).unwrap_or(0))
526 }
527
528 fn close_client(&self, client_id: &str) -> Result<(), GatewayError> {
529 if client_id.is_empty() {
530 return Err(GatewayError::InvalidClientId(client_id.to_string()));
531 }
532 let mut state = self.state.lock();
533 if state.client_messages.remove(client_id).is_none() {
534 return Err(GatewayError::ClientNotFound(client_id.to_string()));
535 }
536 state.client_groups.remove(client_id);
538 for client_ids in state.group_clients.values_mut() {
540 client_ids.retain(|c| c != client_id);
541 }
542 Ok(())
543 }
544}
545
546pub struct Gateway {
575 config: GatewayConfig,
577 transport: Arc<dyn GatewayTransport>,
579}
580
581impl Gateway {
582 pub fn new(config: GatewayConfig, transport: Arc<dyn GatewayTransport>) -> Self {
589 Self { config, transport }
590 }
591
592 pub fn send_to_client(&self, client_id: &str, message: &str) -> Result<(), GatewayError> {
594 self.transport.send_to_client(client_id, message)
595 }
596
597 pub fn send_to_all(&self, message: &str) -> Result<(), GatewayError> {
599 self.transport.send_to_all(message)
600 }
601
602 pub fn send_to_group(&self, group: &str, message: &str) -> Result<(), GatewayError> {
604 self.transport.send_to_group(group, message)
605 }
606
607 pub fn join_group(&self, client_id: &str, group: &str) -> Result<(), GatewayError> {
609 self.transport.join_group(client_id, group)
610 }
611
612 pub fn leave_group(&self, client_id: &str, group: &str) -> Result<(), GatewayError> {
614 self.transport.leave_group(client_id, group)
615 }
616
617 pub fn ungroup(&self, group: &str) -> Result<(), GatewayError> {
619 self.transport.ungroup(group)
620 }
621
622 pub fn is_online(&self, client_id: &str) -> Result<bool, GatewayError> {
624 self.transport.is_online(client_id)
625 }
626
627 pub fn get_client_count(&self) -> Result<usize, GatewayError> {
629 self.transport.get_client_count()
630 }
631
632 pub fn get_client_count_by_group(&self, group: &str) -> Result<usize, GatewayError> {
634 self.transport.get_client_count_by_group(group)
635 }
636
637 pub fn get_all_client_ids(&self) -> Result<Vec<String>, GatewayError> {
639 self.transport.get_all_client_ids()
640 }
641
642 pub fn get_client_id_list_by_group(&self, group: &str) -> Result<Vec<String>, GatewayError> {
644 self.transport.get_client_id_list_by_group(group)
645 }
646
647 pub fn close_client(&self, client_id: &str) -> Result<(), GatewayError> {
649 self.transport.close_client(client_id)
650 }
651
652 pub fn config(&self) -> &GatewayConfig {
654 &self.config
655 }
656}
657
658#[cfg(test)]
663mod tests {
664 use super::*;
665
666 const CLIENT_A: &str = "7f00000108fc00000001";
672 const CLIENT_B: &str = "7f00000108fc00000002";
674 const CLIENT_C: &str = "7f00000108fc00000003";
676
677 #[test]
683 fn test_gateway_config_builder() {
684 let config = GatewayConfig::new("127.0.0.1:1238")
685 .with_heartbeat_interval(30)
686 .with_default_group("default");
687
688 assert_eq!(config.register_address, "127.0.0.1:1238");
689 assert_eq!(config.heartbeat_interval, 30);
690 assert_eq!(config.default_group.as_deref(), Some("default"));
691 }
692
693 #[test]
699 fn test_memory_gateway_transport_send_to_client() {
700 let transport = MemoryGatewayTransport::new();
701 transport.register_client(CLIENT_A);
702
703 transport.send_to_client(CLIENT_A, "hello").unwrap();
704 transport.send_to_client(CLIENT_A, "world").unwrap();
705
706 let messages = transport.client_messages(CLIENT_A);
707 assert_eq!(messages, vec!["hello".to_string(), "world".to_string()]);
708 }
709
710 #[test]
716 fn test_memory_gateway_transport_send_to_clients() {
717 let transport = MemoryGatewayTransport::new();
718 transport.register_client(CLIENT_A);
719 transport.register_client(CLIENT_B);
720
721 transport
722 .send_to_clients(&[CLIENT_A.to_string(), CLIENT_B.to_string()], "broadcast")
723 .unwrap();
724
725 assert_eq!(
726 transport.client_messages(CLIENT_A),
727 vec!["broadcast".to_string()]
728 );
729 assert_eq!(
730 transport.client_messages(CLIENT_B),
731 vec!["broadcast".to_string()]
732 );
733 }
734
735 #[test]
741 fn test_memory_gateway_transport_send_to_all() {
742 let transport = MemoryGatewayTransport::new();
743 transport.register_client(CLIENT_A);
744 transport.register_client(CLIENT_B);
745 transport.register_client(CLIENT_C);
746
747 transport.send_to_all("announcement").unwrap();
748
749 assert_eq!(
750 transport.client_messages(CLIENT_A),
751 vec!["announcement".to_string()]
752 );
753 assert_eq!(
754 transport.client_messages(CLIENT_B),
755 vec!["announcement".to_string()]
756 );
757 assert_eq!(
758 transport.client_messages(CLIENT_C),
759 vec!["announcement".to_string()]
760 );
761 }
762
763 #[test]
769 fn test_memory_gateway_transport_send_to_group() {
770 let transport = MemoryGatewayTransport::new();
771 transport.register_client(CLIENT_A);
772 transport.register_client(CLIENT_B);
773 transport.register_client(CLIENT_C);
774
775 transport.join_group(CLIENT_A, "room1").unwrap();
777 transport.join_group(CLIENT_B, "room1").unwrap();
778
779 transport.send_to_group("room1", "group-msg").unwrap();
780
781 assert_eq!(
782 transport.client_messages(CLIENT_A),
783 vec!["group-msg".to_string()]
784 );
785 assert_eq!(
786 transport.client_messages(CLIENT_B),
787 vec!["group-msg".to_string()]
788 );
789 assert!(transport.client_messages(CLIENT_C).is_empty());
791
792 transport.send_to_group("nonexistent", "msg").unwrap();
794 }
795
796 #[test]
802 fn test_memory_gateway_transport_join_leave_group() {
803 let transport = MemoryGatewayTransport::new();
804 transport.register_client(CLIENT_A);
805
806 transport.join_group(CLIENT_A, "room1").unwrap();
808 transport.join_group(CLIENT_A, "room2").unwrap();
809
810 let groups = transport.get_groups_by_client_id(CLIENT_A).unwrap();
811 assert_eq!(groups, vec!["room1".to_string(), "room2".to_string()]);
812
813 let clients = transport.get_client_id_list_by_group("room1").unwrap();
814 assert_eq!(clients, vec![CLIENT_A.to_string()]);
815
816 transport.join_group(CLIENT_A, "room1").unwrap();
818 let groups = transport.get_groups_by_client_id(CLIENT_A).unwrap();
819 assert_eq!(groups.len(), 2);
820
821 transport.leave_group(CLIENT_A, "room1").unwrap();
823 let groups = transport.get_groups_by_client_id(CLIENT_A).unwrap();
824 assert_eq!(groups, vec!["room2".to_string()]);
825
826 let clients = transport.get_client_id_list_by_group("room1").unwrap();
827 assert!(clients.is_empty());
828
829 let err = transport.leave_group(CLIENT_A, "nonexistent").unwrap_err();
831 match err {
832 GatewayError::GroupNotFound(group) => assert_eq!(group, "nonexistent"),
833 other => panic!("期望 GroupNotFound, 实际 {other:?}"),
834 }
835 }
836
837 #[test]
843 fn test_memory_gateway_transport_ungroup() {
844 let transport = MemoryGatewayTransport::new();
845 transport.register_client(CLIENT_A);
846 transport.register_client(CLIENT_B);
847
848 transport.join_group(CLIENT_A, "room1").unwrap();
849 transport.join_group(CLIENT_B, "room1").unwrap();
850
851 transport.ungroup("room1").unwrap();
853
854 let clients = transport.get_client_id_list_by_group("room1").unwrap();
856 assert!(clients.is_empty());
857
858 let groups_a = transport.get_groups_by_client_id(CLIENT_A).unwrap();
860 assert!(!groups_a.iter().any(|g| g == "room1"));
861 let groups_b = transport.get_groups_by_client_id(CLIENT_B).unwrap();
862 assert!(!groups_b.iter().any(|g| g == "room1"));
863
864 let err = transport.ungroup("nonexistent").unwrap_err();
866 match err {
867 GatewayError::GroupNotFound(group) => assert_eq!(group, "nonexistent"),
868 other => panic!("期望 GroupNotFound, 实际 {other:?}"),
869 }
870 }
871
872 #[test]
878 fn test_memory_gateway_transport_is_online() {
879 let transport = MemoryGatewayTransport::new();
880
881 assert!(!transport.is_online(CLIENT_A).unwrap());
883
884 transport.register_client(CLIENT_A);
885 assert!(transport.is_online(CLIENT_A).unwrap());
886
887 transport.close_client(CLIENT_A).unwrap();
889 assert!(!transport.is_online(CLIENT_A).unwrap());
890 }
891
892 #[test]
898 fn test_memory_gateway_transport_get_client_count() {
899 let transport = MemoryGatewayTransport::new();
900 assert_eq!(transport.get_client_count().unwrap(), 0);
901
902 transport.register_client(CLIENT_A);
903 assert_eq!(transport.get_client_count().unwrap(), 1);
904
905 transport.register_client(CLIENT_B);
906 transport.register_client(CLIENT_C);
907 assert_eq!(transport.get_client_count().unwrap(), 3);
908
909 transport.register_client(CLIENT_A);
911 assert_eq!(transport.get_client_count().unwrap(), 3);
912
913 transport.close_client(CLIENT_B).unwrap();
915 assert_eq!(transport.get_client_count().unwrap(), 2);
916 }
917
918 #[test]
924 fn test_memory_gateway_transport_get_client_count_by_group() {
925 let transport = MemoryGatewayTransport::new();
926 transport.register_client(CLIENT_A);
927 transport.register_client(CLIENT_B);
928 transport.register_client(CLIENT_C);
929
930 assert_eq!(transport.get_client_count_by_group("room1").unwrap(), 0);
932
933 transport.join_group(CLIENT_A, "room1").unwrap();
934 transport.join_group(CLIENT_B, "room1").unwrap();
935 transport.join_group(CLIENT_C, "room2").unwrap();
936
937 assert_eq!(transport.get_client_count_by_group("room1").unwrap(), 2);
938 assert_eq!(transport.get_client_count_by_group("room2").unwrap(), 1);
939
940 transport.leave_group(CLIENT_A, "room1").unwrap();
942 assert_eq!(transport.get_client_count_by_group("room1").unwrap(), 1);
943 }
944
945 #[test]
951 fn test_memory_gateway_transport_get_all_client_ids() {
952 let transport = MemoryGatewayTransport::new();
953 assert!(transport.get_all_client_ids().unwrap().is_empty());
954
955 transport.register_client(CLIENT_A);
956 transport.register_client(CLIENT_B);
957
958 let mut ids = transport.get_all_client_ids().unwrap();
959 ids.sort();
960 assert_eq!(ids, vec![CLIENT_A.to_string(), CLIENT_B.to_string()]);
961 }
962
963 #[test]
969 fn test_memory_gateway_transport_get_client_id_list_by_group() {
970 let transport = MemoryGatewayTransport::new();
971 transport.register_client(CLIENT_A);
972 transport.register_client(CLIENT_B);
973
974 assert!(transport
976 .get_client_id_list_by_group("room1")
977 .unwrap()
978 .is_empty());
979
980 transport.join_group(CLIENT_A, "room1").unwrap();
981 transport.join_group(CLIENT_B, "room1").unwrap();
982
983 let mut clients = transport.get_client_id_list_by_group("room1").unwrap();
984 clients.sort();
985 assert_eq!(clients, vec![CLIENT_A.to_string(), CLIENT_B.to_string()]);
986 }
987
988 #[test]
994 fn test_memory_gateway_transport_get_groups_by_client_id() {
995 let transport = MemoryGatewayTransport::new();
996 transport.register_client(CLIENT_A);
997
998 assert!(transport
1000 .get_groups_by_client_id(CLIENT_A)
1001 .unwrap()
1002 .is_empty());
1003
1004 transport.join_group(CLIENT_A, "room1").unwrap();
1005 transport.join_group(CLIENT_A, "room2").unwrap();
1006
1007 let groups = transport.get_groups_by_client_id(CLIENT_A).unwrap();
1008 assert_eq!(groups, vec!["room1".to_string(), "room2".to_string()]);
1009
1010 let err = transport.get_groups_by_client_id(CLIENT_B).unwrap_err();
1012 match err {
1013 GatewayError::ClientNotFound(client_id) => assert_eq!(client_id, CLIENT_B),
1014 other => panic!("期望 ClientNotFound, 实际 {other:?}"),
1015 }
1016 }
1017
1018 #[test]
1024 fn test_memory_gateway_transport_close_client() {
1025 let transport = MemoryGatewayTransport::new();
1026 transport.register_client(CLIENT_A);
1027 transport.register_client(CLIENT_B);
1028
1029 transport.join_group(CLIENT_A, "room1").unwrap();
1030 transport.join_group(CLIENT_B, "room1").unwrap();
1031
1032 transport.close_client(CLIENT_A).unwrap();
1034
1035 assert!(!transport.is_online(CLIENT_A).unwrap());
1037 assert_eq!(transport.get_client_count().unwrap(), 1);
1038
1039 let clients = transport.get_client_id_list_by_group("room1").unwrap();
1041 assert_eq!(clients, vec![CLIENT_B.to_string()]);
1042
1043 let err = transport.close_client(CLIENT_A).unwrap_err();
1045 match err {
1046 GatewayError::ClientNotFound(client_id) => assert_eq!(client_id, CLIENT_A),
1047 other => panic!("期望 ClientNotFound, 实际 {other:?}"),
1048 }
1049 }
1050
1051 #[test]
1057 fn test_memory_gateway_transport_client_not_found() {
1058 let transport = MemoryGatewayTransport::new();
1059
1060 let err = transport.send_to_client(CLIENT_A, "msg").unwrap_err();
1062 match err {
1063 GatewayError::ClientNotFound(client_id) => assert_eq!(client_id, CLIENT_A),
1064 other => panic!("期望 ClientNotFound, 实际 {other:?}"),
1065 }
1066
1067 let err = transport.join_group(CLIENT_A, "room1").unwrap_err();
1069 match err {
1070 GatewayError::ClientNotFound(client_id) => assert_eq!(client_id, CLIENT_A),
1071 other => panic!("期望 ClientNotFound, 实际 {other:?}"),
1072 }
1073
1074 transport.register_client(CLIENT_A);
1076 let err = transport
1077 .send_to_clients(&[CLIENT_A.to_string(), CLIENT_B.to_string()], "msg")
1078 .unwrap_err();
1079 match err {
1080 GatewayError::ClientNotFound(client_id) => assert_eq!(client_id, CLIENT_B),
1081 other => panic!("期望 ClientNotFound, 实际 {other:?}"),
1082 }
1083
1084 assert!(transport.client_messages(CLIENT_A).is_empty());
1086 }
1087
1088 #[test]
1094 fn test_gateway_send_to_client() {
1095 let transport = Arc::new(MemoryGatewayTransport::new());
1096 let gateway = Gateway::new(GatewayConfig::new("127.0.0.1:1238"), transport.clone());
1097
1098 transport.register_client(CLIENT_A);
1099 gateway.send_to_client(CLIENT_A, "hello").unwrap();
1100
1101 assert_eq!(
1102 transport.client_messages(CLIENT_A),
1103 vec!["hello".to_string()]
1104 );
1105 }
1106
1107 #[test]
1109 fn test_gateway_send_to_all() {
1110 let transport = Arc::new(MemoryGatewayTransport::new());
1111 let gateway = Gateway::new(GatewayConfig::new("127.0.0.1:1238"), transport.clone());
1112
1113 transport.register_client(CLIENT_A);
1114 transport.register_client(CLIENT_B);
1115
1116 gateway.send_to_all("broadcast").unwrap();
1117
1118 assert_eq!(
1119 transport.client_messages(CLIENT_A),
1120 vec!["broadcast".to_string()]
1121 );
1122 assert_eq!(
1123 transport.client_messages(CLIENT_B),
1124 vec!["broadcast".to_string()]
1125 );
1126 }
1127
1128 #[test]
1130 fn test_gateway_send_to_group() {
1131 let transport = Arc::new(MemoryGatewayTransport::new());
1132 let gateway = Gateway::new(GatewayConfig::new("127.0.0.1:1238"), transport.clone());
1133
1134 transport.register_client(CLIENT_A);
1135 transport.register_client(CLIENT_B);
1136 transport.join_group(CLIENT_A, "room1").unwrap();
1137 transport.join_group(CLIENT_B, "room1").unwrap();
1138
1139 gateway.send_to_group("room1", "group-msg").unwrap();
1140
1141 assert_eq!(
1142 transport.client_messages(CLIENT_A),
1143 vec!["group-msg".to_string()]
1144 );
1145 assert_eq!(
1146 transport.client_messages(CLIENT_B),
1147 vec!["group-msg".to_string()]
1148 );
1149 }
1150
1151 #[test]
1153 fn test_gateway_join_group() {
1154 let transport = Arc::new(MemoryGatewayTransport::new());
1155 let gateway = Gateway::new(GatewayConfig::new("127.0.0.1:1238"), transport.clone());
1156
1157 transport.register_client(CLIENT_A);
1158
1159 gateway.join_group(CLIENT_A, "room1").unwrap();
1160
1161 let groups = transport.get_groups_by_client_id(CLIENT_A).unwrap();
1162 assert_eq!(groups, vec!["room1".to_string()]);
1163
1164 let clients = gateway.get_client_id_list_by_group("room1").unwrap();
1165 assert_eq!(clients, vec![CLIENT_A.to_string()]);
1166 }
1167
1168 #[test]
1170 fn test_gateway_is_online() {
1171 let transport = Arc::new(MemoryGatewayTransport::new());
1172 let gateway = Gateway::new(GatewayConfig::new("127.0.0.1:1238"), transport.clone());
1173
1174 assert!(!gateway.is_online(CLIENT_A).unwrap());
1175
1176 transport.register_client(CLIENT_A);
1177 assert!(gateway.is_online(CLIENT_A).unwrap());
1178 }
1179
1180 #[test]
1182 fn test_gateway_get_client_count() {
1183 let transport = Arc::new(MemoryGatewayTransport::new());
1184 let gateway = Gateway::new(GatewayConfig::new("127.0.0.1:1238"), transport.clone());
1185
1186 assert_eq!(gateway.get_client_count().unwrap(), 0);
1187
1188 transport.register_client(CLIENT_A);
1189 transport.register_client(CLIENT_B);
1190 assert_eq!(gateway.get_client_count().unwrap(), 2);
1191
1192 assert_eq!(gateway.config().register_address, "127.0.0.1:1238");
1194 assert_eq!(gateway.config().heartbeat_interval, 55);
1195 }
1196
1197 #[test]
1199 fn test_gateway_close_client() {
1200 let transport = Arc::new(MemoryGatewayTransport::new());
1201 let gateway = Gateway::new(GatewayConfig::new("127.0.0.1:1238"), transport.clone());
1202
1203 transport.register_client(CLIENT_A);
1204 transport.join_group(CLIENT_A, "room1").unwrap();
1205
1206 gateway.close_client(CLIENT_A).unwrap();
1207
1208 assert!(!gateway.is_online(CLIENT_A).unwrap());
1209 assert_eq!(gateway.get_client_count().unwrap(), 0);
1210 assert!(gateway
1211 .get_client_id_list_by_group("room1")
1212 .unwrap()
1213 .is_empty());
1214 }
1215}