Skip to main content

sz_rust_core/
gateway.rs

1//! Gateway 模块 — GatewayWorker Gateway API 抽象(对齐 PHP `GatewayWorker\Gateway`)
2//!
3//! 提供 WebSocket 客户端管理、群组广播、消息推送等能力。
4//!
5//! ## PHP 对齐
6//!
7//! ### 核心 API 映射
8//!
9//! | PHP 方法 | Rust 方法 | 说明 |
10//! |---------|-----------|------|
11//! | `Gateway::sendToClient($client_id, $message)` | [`Gateway::send_to_client`] | 向指定客户端发送消息 |
12//! | `Gateway::sendToAll($message)` | [`Gateway::send_to_all`] | 向所有在线客户端广播 |
13//! | `Gateway::sendToGroup($group, $message)` | [`Gateway::send_to_group`] | 向指定群组发送消息 |
14//! | `Gateway::joinGroup($client_id, $group)` | [`Gateway::join_group`] | 将客户端加入群组 |
15//! | `Gateway::leaveGroup($client_id, $group)` | [`Gateway::leave_group`] | 将客户端离开群组 |
16//! | `Gateway::ungroup($group)` | [`Gateway::ungroup`] | 解散指定群组 |
17//! | `Gateway::isOnline($client_id)` | [`Gateway::is_online`] | 判断客户端是否在线 |
18//! | `Gateway::getClientCount()` | [`Gateway::get_client_count`] | 获取在线客户端数量 |
19//! | `Gateway::getClientCountByGroup($group)` | [`Gateway::get_client_count_by_group`] | 获取群组客户端数量 |
20//! | `Gateway::getAllClientIds()` | [`Gateway::get_all_client_ids`] | 获取所有在线 client_id |
21//! | `Gateway::getClientIdListByGroup($group)` | [`Gateway::get_client_id_list_by_group`] | 获取群组中的 client_id |
22//! | `Gateway::closeClient($client_id)` | [`Gateway::close_client`] | 关闭客户端连接 |
23//!
24//! ### PHP 行为对齐
25//!
26//! - **静态方法 → 实例方法**:PHP `Gateway` 全部为静态方法,Rust 通过 [`Gateway`] 实例方法
27//!   表达,配置与传输层通过构造函数注入,便于测试和多实例隔离。
28//! - **Register 转发 → Transport 抽象**:PHP GatewayWorker 通过 Register 服务器转发命令到
29//!   BusinessWorker/Worker。Rust 端通过 [`GatewayTransport`] trait 抽象底层通信,业务方可
30//!   实现具体通信(如 Redis pub/sub、HTTP、TCP)。
31//! - **client_id 格式**:对齐 PHP GatewayWorker 的 20 字符 hex 字符串。
32//!
33//! ## 架构说明
34//!
35//! - [`GatewayTransport`] trait:抽象 Gateway API 的底层通信
36//! - [`MemoryGatewayTransport`]:内存实现,用于测试和开发环境
37//! - [`Gateway`]:面向业务的 Gateway API 客户端,委托 [`GatewayTransport`] 执行具体操作
38
39use std::collections::HashMap;
40use std::sync::Arc;
41
42use parking_lot::Mutex;
43use thiserror::Error;
44
45// ============================================================================
46// 错误类型
47// ============================================================================
48
49/// Gateway 错误
50///
51/// 对齐 PHP `GatewayWorker\Gateway` 在各类异常场景下抛出的错误。
52#[derive(Debug, Error)]
53pub enum GatewayError {
54    /// 客户端未找到(client_id 不在线或不存在)
55    #[error("客户端未找到: {0}")]
56    ClientNotFound(String),
57    /// 群组未找到
58    #[error("群组未找到: {0}")]
59    GroupNotFound(String),
60    /// 发送失败
61    #[error("发送失败: {0}")]
62    SendFailed(String),
63    /// 无效的 client_id(格式非法或为空)
64    #[error("无效的 client_id: {0}")]
65    InvalidClientId(String),
66    /// 传输错误
67    #[error("传输错误: {0}")]
68    Transport(String),
69    /// 序列化失败
70    #[error("序列化失败: {0}")]
71    Serialize(String),
72}
73
74// ============================================================================
75// ClientId 类型
76// ============================================================================
77
78/// 客户端 ID — 对齐 PHP GatewayWorker 的 client_id(20 字符 hex 字符串)
79pub type ClientId = String;
80
81// ============================================================================
82// GatewayConfig
83// ============================================================================
84
85/// Gateway 配置 — 对齐 PHP `GatewayWorker\Gateway` 的 Register 地址配置
86///
87/// # PHP 对齐
88///
89/// ```php
90/// // PHP GatewayWorker Register 地址配置
91/// Gateway::$registerAddress = '127.0.0.1:1238';
92/// Gateway::$defaultGroup = 'default';
93/// ```
94///
95/// # Rust 用法
96///
97/// ```rust,ignore
98/// use sz_rust_core::gateway::GatewayConfig;
99///
100/// let config = GatewayConfig::new("127.0.0.1:1238")
101///     .with_heartbeat_interval(30)
102///     .with_default_group("default");
103/// ```
104#[derive(Debug, Clone)]
105pub struct GatewayConfig {
106    /// Register 服务器地址(如 `127.0.0.1:1238`)
107    pub register_address: String,
108    /// 心跳间隔(秒,默认 55)
109    pub heartbeat_interval: u64,
110    /// 默认群组名(对齐 PHP `Gateway::$defaultGroup`)
111    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    /// 创建新配置
126    ///
127    /// # 参数
128    ///
129    /// - `register_address`: Register 服务器地址
130    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    /// 设置心跳间隔(秒)
139    pub fn with_heartbeat_interval(mut self, interval: u64) -> Self {
140        self.heartbeat_interval = interval;
141        self
142    }
143
144    /// 设置默认群组名
145    pub fn with_default_group(mut self, group: impl Into<String>) -> Self {
146        self.default_group = Some(group.into());
147        self
148    }
149}
150
151// ============================================================================
152// GatewayTransport trait
153// ============================================================================
154
155/// Gateway 传输层 trait — 抽象 Gateway API 的底层通信
156///
157/// PHP GatewayWorker 通过 Register 服务器转发命令到 BusinessWorker/Worker。
158/// Rust 端通过此 trait 抽象,业务方可实现具体通信(如 Redis pub/sub、HTTP、TCP)。
159///
160/// # 线程安全
161///
162/// 实现者必须保证 `Send + Sync`,因为 [`Gateway`] 通常作为单例在多线程下使用。
163pub trait GatewayTransport: Send + Sync {
164    /// 向指定 client_id 发送消息
165    ///
166    /// # 参数
167    ///
168    /// - `client_id`: 目标客户端 ID
169    /// - `message`: 消息内容
170    ///
171    /// # 返回
172    ///
173    /// 成功返回 `Ok(())`,客户端不在线返回 [`GatewayError::ClientNotFound`]。
174    fn send_to_client(&self, client_id: &str, message: &str) -> Result<(), GatewayError>;
175
176    /// 向多个 client_id 发送消息
177    ///
178    /// # 参数
179    ///
180    /// - `client_ids`: 目标客户端 ID 列表
181    /// - `message`: 消息内容
182    ///
183    /// # 返回
184    ///
185    /// 任一客户端不在线即返回 [`GatewayError::ClientNotFound`]。
186    fn send_to_clients(&self, client_ids: &[String], message: &str) -> Result<(), GatewayError>;
187
188    /// 向所有在线 client_id 发送广播
189    ///
190    /// # 参数
191    ///
192    /// - `message`: 消息内容
193    fn send_to_all(&self, message: &str) -> Result<(), GatewayError>;
194
195    /// 向指定群组发送消息
196    ///
197    /// # 参数
198    ///
199    /// - `group`: 群组名
200    /// - `message`: 消息内容
201    fn send_to_group(&self, group: &str, message: &str) -> Result<(), GatewayError>;
202
203    /// 获取所有在线 client_id
204    fn get_all_client_ids(&self) -> Result<Vec<String>, GatewayError>;
205
206    /// 获取指定群组中的 client_id
207    ///
208    /// # 参数
209    ///
210    /// - `group`: 群组名
211    fn get_client_id_list_by_group(&self, group: &str) -> Result<Vec<String>, GatewayError>;
212
213    /// 获取 client_id 所在的群组
214    ///
215    /// # 参数
216    ///
217    /// - `client_id`: 客户端 ID
218    ///
219    /// # 返回
220    ///
221    /// 客户端不在线返回 [`GatewayError::ClientNotFound`]。
222    fn get_groups_by_client_id(&self, client_id: &str) -> Result<Vec<String>, GatewayError>;
223
224    /// 将 client_id 加入群组
225    ///
226    /// # 参数
227    ///
228    /// - `client_id`: 客户端 ID
229    /// - `group`: 群组名
230    ///
231    /// # 返回
232    ///
233    /// 客户端不在线返回 [`GatewayError::ClientNotFound`]。
234    fn join_group(&self, client_id: &str, group: &str) -> Result<(), GatewayError>;
235
236    /// 将 client_id 离开群组
237    ///
238    /// # 参数
239    ///
240    /// - `client_id`: 客户端 ID
241    /// - `group`: 群组名
242    ///
243    /// # 返回
244    ///
245    /// 客户端不在线返回 [`GatewayError::ClientNotFound`],群组不存在返回 [`GatewayError::GroupNotFound`]。
246    fn leave_group(&self, client_id: &str, group: &str) -> Result<(), GatewayError>;
247
248    /// 解散指定群组(对齐 PHP `Gateway::ungroup()`)
249    ///
250    /// 移除群组,并将群组内所有 client_id 的群组成员关系清除。
251    ///
252    /// # 参数
253    ///
254    /// - `group`: 群组名
255    ///
256    /// # 返回
257    ///
258    /// 群组不存在返回 [`GatewayError::GroupNotFound`]。
259    fn ungroup(&self, group: &str) -> Result<(), GatewayError>;
260
261    /// 判断 client_id 是否在线
262    fn is_online(&self, client_id: &str) -> Result<bool, GatewayError>;
263
264    /// 获取在线 client_id 数量
265    fn get_client_count(&self) -> Result<usize, GatewayError>;
266
267    /// 获取指定群组的在线 client_id 数量
268    ///
269    /// # 参数
270    ///
271    /// - `group`: 群组名
272    fn get_client_count_by_group(&self, group: &str) -> Result<usize, GatewayError>;
273
274    /// 关闭指定 client_id 的连接
275    ///
276    /// # 参数
277    ///
278    /// - `client_id`: 客户端 ID
279    ///
280    /// # 返回
281    ///
282    /// 客户端不在线返回 [`GatewayError::ClientNotFound`]。
283    fn close_client(&self, client_id: &str) -> Result<(), GatewayError>;
284}
285
286// ============================================================================
287// MemoryGatewayTransport
288// ============================================================================
289
290/// 内存 Gateway 传输实现 — 用于测试和开发环境
291///
292/// 不进行实际网络通信,所有状态保存在内存中,供测试断言使用。
293///
294/// # 数据结构
295///
296/// - `client_messages`: client_id → messages 队列(键存在即表示 client 在线)
297/// - `client_groups`: client_id → groups 映射
298/// - `group_clients`: group → client_ids 映射
299///
300/// # 线程安全
301///
302/// 通过 `Arc<Mutex<GatewayState>>` 保护,支持并发访问。所有操作在单个锁内完成,
303/// 保证多字段状态的一致性。
304///
305/// # 用法
306///
307/// ```rust,ignore
308/// use sz_rust_core::gateway::{MemoryGatewayTransport, GatewayTransport};
309///
310/// let transport = MemoryGatewayTransport::new();
311/// transport.register_client("7f00000108fc00000001");
312/// transport.send_to_client("7f00000108fc00000001", "hello").unwrap();
313/// assert_eq!(transport.client_messages("7f00000108fc00000001"), vec!["hello".to_string()]);
314/// ```
315#[derive(Debug, Default)]
316pub struct MemoryGatewayTransport {
317    /// Gateway 内存状态
318    state: Mutex<GatewayState>,
319}
320
321/// Gateway 内存状态
322#[derive(Debug, Default)]
323struct GatewayState {
324    /// client_id → messages 队列(键存在即表示 client 在线)
325    client_messages: HashMap<ClientId, Vec<String>>,
326    /// client_id → groups 映射
327    client_groups: HashMap<ClientId, Vec<String>>,
328    /// group → client_ids 映射
329    group_clients: HashMap<String, Vec<ClientId>>,
330}
331
332impl MemoryGatewayTransport {
333    /// 创建新的内存传输
334    pub fn new() -> Self {
335        Self::default()
336    }
337
338    /// 注册客户端(标记为在线)
339    ///
340    /// 非 trait 方法,供测试/开发注册在线客户端。对齐 PHP GatewayWorker 中客户端连接
341    /// 建立后自动上线的行为。重复注册为幂等操作。
342    ///
343    /// # 参数
344    ///
345    /// - `client_id`: 客户端 ID
346    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    /// 获取客户端收到的所有消息(快照)
355    ///
356    /// 非 trait 方法,供测试断言。客户端不在线时返回空 Vec。
357    ///
358    /// # 参数
359    ///
360    /// - `client_id`: 客户端 ID
361    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        // 先校验所有客户端在线,再统一发送,避免部分发送
389        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        // 群组不存在时为空操作(对齐 PHP 广播到空群组的行为)
416        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        // 维护 client_id → groups 映射(去重)
461        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        // 维护 group → client_ids 映射(去重)
469        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        // 从 client_id → groups 移除
489        if let Some(groups) = state.client_groups.get_mut(client_id) {
490            groups.retain(|g| g != group);
491        }
492        // 从 group → client_ids 移除
493        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        // 移除群组
505        state.group_clients.remove(group);
506        // 从所有 client 的群组列表中移除该群组
507        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        // 移除 client 的群组列表
537        state.client_groups.remove(client_id);
538        // 从所有群组中移除该 client
539        for client_ids in state.group_clients.values_mut() {
540            client_ids.retain(|c| c != client_id);
541        }
542        Ok(())
543    }
544}
545
546// ============================================================================
547// Gateway
548// ============================================================================
549
550/// Gateway API 客户端 — 对齐 PHP `GatewayWorker\Gateway`
551///
552/// 提供 WebSocket 客户端管理、群组广播、消息推送等能力。
553/// 通过 [`GatewayTransport`] trait 抽象底层通信。
554///
555/// # PHP 对齐
556///
557/// PHP `Gateway` 全部为静态方法,Rust 通过实例方法表达,配置与传输层通过构造函数注入,
558/// 便于测试和多实例隔离。
559///
560/// # 用法
561///
562/// ```rust,ignore
563/// use sz_rust_core::gateway::{
564///     Gateway, GatewayConfig, MemoryGatewayTransport, GatewayTransport,
565/// };
566/// use std::sync::Arc;
567///
568/// let transport = Arc::new(MemoryGatewayTransport::new());
569/// let gateway = Gateway::new(GatewayConfig::new("127.0.0.1:1238"), transport.clone());
570///
571/// transport.register_client("7f00000108fc00000001");
572/// gateway.send_to_client("7f00000108fc00000001", "hello").unwrap();
573/// ```
574pub struct Gateway {
575    /// Gateway 配置
576    config: GatewayConfig,
577    /// Gateway 传输层实现
578    transport: Arc<dyn GatewayTransport>,
579}
580
581impl Gateway {
582    /// 创建 Gateway 客户端
583    ///
584    /// # 参数
585    ///
586    /// - `config`: Gateway 配置
587    /// - `transport`: 传输层实现(业务方注入 Redis pub/sub / HTTP / TCP 等具体实现)
588    pub fn new(config: GatewayConfig, transport: Arc<dyn GatewayTransport>) -> Self {
589        Self { config, transport }
590    }
591
592    /// 向指定 client_id 发送消息 — 对齐 `Gateway::sendToClient()`
593    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    /// 向所有在线 client_id 发送广播 — 对齐 `Gateway::sendToAll()`
598    pub fn send_to_all(&self, message: &str) -> Result<(), GatewayError> {
599        self.transport.send_to_all(message)
600    }
601
602    /// 向指定群组发送消息 — 对齐 `Gateway::sendToGroup()`
603    pub fn send_to_group(&self, group: &str, message: &str) -> Result<(), GatewayError> {
604        self.transport.send_to_group(group, message)
605    }
606
607    /// 将 client_id 加入群组 — 对齐 `Gateway::joinGroup()`
608    pub fn join_group(&self, client_id: &str, group: &str) -> Result<(), GatewayError> {
609        self.transport.join_group(client_id, group)
610    }
611
612    /// 将 client_id 离开群组 — 对齐 `Gateway::leaveGroup()`
613    pub fn leave_group(&self, client_id: &str, group: &str) -> Result<(), GatewayError> {
614        self.transport.leave_group(client_id, group)
615    }
616
617    /// 解散指定群组 — 对齐 `Gateway::ungroup()`
618    pub fn ungroup(&self, group: &str) -> Result<(), GatewayError> {
619        self.transport.ungroup(group)
620    }
621
622    /// 判断 client_id 是否在线 — 对齐 `Gateway::isOnline()`
623    pub fn is_online(&self, client_id: &str) -> Result<bool, GatewayError> {
624        self.transport.is_online(client_id)
625    }
626
627    /// 获取在线 client_id 数量 — 对齐 `Gateway::getClientCount()`
628    pub fn get_client_count(&self) -> Result<usize, GatewayError> {
629        self.transport.get_client_count()
630    }
631
632    /// 获取指定群组的在线 client_id 数量 — 对齐 `Gateway::getClientCountByGroup()`
633    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    /// 获取所有在线 client_id — 对齐 `Gateway::getAllClientIds()`
638    pub fn get_all_client_ids(&self) -> Result<Vec<String>, GatewayError> {
639        self.transport.get_all_client_ids()
640    }
641
642    /// 获取指定群组中的 client_id — 对齐 `Gateway::getClientIdListByGroup()`
643    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    /// 关闭指定 client_id 的连接 — 对齐 `Gateway::closeClient()`
648    pub fn close_client(&self, client_id: &str) -> Result<(), GatewayError> {
649        self.transport.close_client(client_id)
650    }
651
652    /// 获取 Gateway 配置
653    pub fn config(&self) -> &GatewayConfig {
654        &self.config
655    }
656}
657
658// ============================================================================
659// 单元测试
660// ============================================================================
661
662#[cfg(test)]
663mod tests {
664    use super::*;
665
666    // ------------------------------------------------------------------------
667    // 辅助常量
668    // ------------------------------------------------------------------------
669
670    /// 测试用 client_id(对齐 PHP GatewayWorker 20 字符 hex 格式)
671    const CLIENT_A: &str = "7f00000108fc00000001";
672    /// 测试用 client_id B
673    const CLIENT_B: &str = "7f00000108fc00000002";
674    /// 测试用 client_id C
675    const CLIENT_C: &str = "7f00000108fc00000003";
676
677    // ------------------------------------------------------------------------
678    // GatewayConfig 测试
679    // ------------------------------------------------------------------------
680
681    /// 测试 GatewayConfig builder 模式
682    #[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    // ------------------------------------------------------------------------
694    // MemoryGatewayTransport — send_to_client
695    // ------------------------------------------------------------------------
696
697    /// 测试 MemoryGatewayTransport 向单个客户端发送消息
698    #[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    // ------------------------------------------------------------------------
711    // MemoryGatewayTransport — send_to_clients
712    // ------------------------------------------------------------------------
713
714    /// 测试 MemoryGatewayTransport 向多个客户端发送消息
715    #[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    // ------------------------------------------------------------------------
736    // MemoryGatewayTransport — send_to_all
737    // ------------------------------------------------------------------------
738
739    /// 测试 MemoryGatewayTransport 向所有在线客户端广播
740    #[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    // ------------------------------------------------------------------------
764    // MemoryGatewayTransport — send_to_group
765    // ------------------------------------------------------------------------
766
767    /// 测试 MemoryGatewayTransport 向群组发送消息
768    #[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        // A、B 加入 room1,C 不加入
776        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        // C 不在群组,不应收到消息
790        assert!(transport.client_messages(CLIENT_C).is_empty());
791
792        // 向不存在的群组发送消息应为空操作(不报错)
793        transport.send_to_group("nonexistent", "msg").unwrap();
794    }
795
796    // ------------------------------------------------------------------------
797    // MemoryGatewayTransport — join_group / leave_group
798    // ------------------------------------------------------------------------
799
800    /// 测试 MemoryGatewayTransport 加入和离开群组
801    #[test]
802    fn test_memory_gateway_transport_join_leave_group() {
803        let transport = MemoryGatewayTransport::new();
804        transport.register_client(CLIENT_A);
805
806        // 加入群组
807        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        // 重复加入同一群组应为幂等操作
817        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        // 离开群组
822        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        // 离开不存在的群组应返回 GroupNotFound
830        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    // ------------------------------------------------------------------------
838    // MemoryGatewayTransport — ungroup
839    // ------------------------------------------------------------------------
840
841    /// 测试 MemoryGatewayTransport 解散群组
842    #[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        // 解散群组
852        transport.ungroup("room1").unwrap();
853
854        // 群组已不存在
855        let clients = transport.get_client_id_list_by_group("room1").unwrap();
856        assert!(clients.is_empty());
857
858        // client 的群组列表中应移除 room1
859        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        // 解散不存在的群组应返回 GroupNotFound
865        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    // ------------------------------------------------------------------------
873    // MemoryGatewayTransport — is_online
874    // ------------------------------------------------------------------------
875
876    /// 测试 MemoryGatewayTransport 判断客户端在线状态
877    #[test]
878    fn test_memory_gateway_transport_is_online() {
879        let transport = MemoryGatewayTransport::new();
880
881        // 未注册客户端不在线
882        assert!(!transport.is_online(CLIENT_A).unwrap());
883
884        transport.register_client(CLIENT_A);
885        assert!(transport.is_online(CLIENT_A).unwrap());
886
887        // 关闭后不在线
888        transport.close_client(CLIENT_A).unwrap();
889        assert!(!transport.is_online(CLIENT_A).unwrap());
890    }
891
892    // ------------------------------------------------------------------------
893    // MemoryGatewayTransport — get_client_count
894    // ------------------------------------------------------------------------
895
896    /// 测试 MemoryGatewayTransport 获取在线客户端数量
897    #[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        // 重复注册不应增加计数
910        transport.register_client(CLIENT_A);
911        assert_eq!(transport.get_client_count().unwrap(), 3);
912
913        // 关闭后计数减少
914        transport.close_client(CLIENT_B).unwrap();
915        assert_eq!(transport.get_client_count().unwrap(), 2);
916    }
917
918    // ------------------------------------------------------------------------
919    // MemoryGatewayTransport — get_client_count_by_group
920    // ------------------------------------------------------------------------
921
922    /// 测试 MemoryGatewayTransport 获取群组客户端数量
923    #[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        // 群组不存在时返回 0
931        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        // A 离开 room1
941        transport.leave_group(CLIENT_A, "room1").unwrap();
942        assert_eq!(transport.get_client_count_by_group("room1").unwrap(), 1);
943    }
944
945    // ------------------------------------------------------------------------
946    // MemoryGatewayTransport — get_all_client_ids
947    // ------------------------------------------------------------------------
948
949    /// 测试 MemoryGatewayTransport 获取所有在线 client_id
950    #[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    // ------------------------------------------------------------------------
964    // MemoryGatewayTransport — get_client_id_list_by_group
965    // ------------------------------------------------------------------------
966
967    /// 测试 MemoryGatewayTransport 获取群组中的 client_id
968    #[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        // 群组不存在时返回空 Vec
975        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    // ------------------------------------------------------------------------
989    // MemoryGatewayTransport — get_groups_by_client_id
990    // ------------------------------------------------------------------------
991
992    /// 测试 MemoryGatewayTransport 获取客户端所在的群组
993    #[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        // 无群组时返回空 Vec
999        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        // 不在线的客户端返回 ClientNotFound
1011        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    // ------------------------------------------------------------------------
1019    // MemoryGatewayTransport — close_client
1020    // ------------------------------------------------------------------------
1021
1022    /// 测试 MemoryGatewayTransport 关闭客户端连接
1023    #[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        // 关闭 A
1033        transport.close_client(CLIENT_A).unwrap();
1034
1035        // A 不再在线
1036        assert!(!transport.is_online(CLIENT_A).unwrap());
1037        assert_eq!(transport.get_client_count().unwrap(), 1);
1038
1039        // A 应从群组中移除
1040        let clients = transport.get_client_id_list_by_group("room1").unwrap();
1041        assert_eq!(clients, vec![CLIENT_B.to_string()]);
1042
1043        // 关闭不在线的客户端返回 ClientNotFound
1044        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    // ------------------------------------------------------------------------
1052    // MemoryGatewayTransport — ClientNotFound
1053    // ------------------------------------------------------------------------
1054
1055    /// 测试 MemoryGatewayTransport 客户端未找到错误
1056    #[test]
1057    fn test_memory_gateway_transport_client_not_found() {
1058        let transport = MemoryGatewayTransport::new();
1059
1060        // 向未注册客户端发送消息
1061        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        // 向未注册客户端加入群组
1068        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        // send_to_clients 中任一客户端不在线
1075        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        // send_to_clients 失败时不应部分发送
1085        assert!(transport.client_messages(CLIENT_A).is_empty());
1086    }
1087
1088    // ------------------------------------------------------------------------
1089    // Gateway 测试
1090    // ------------------------------------------------------------------------
1091
1092    /// 测试 Gateway 向单个客户端发送消息
1093    #[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    /// 测试 Gateway 向所有在线客户端广播
1108    #[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    /// 测试 Gateway 向群组发送消息
1129    #[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    /// 测试 Gateway 加入群组
1152    #[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    /// 测试 Gateway 判断客户端在线状态
1169    #[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    /// 测试 Gateway 获取在线客户端数量
1181    #[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        // 验证 config 访问器
1193        assert_eq!(gateway.config().register_address, "127.0.0.1:1238");
1194        assert_eq!(gateway.config().heartbeat_interval, 55);
1195    }
1196
1197    /// 测试 Gateway 关闭客户端连接
1198    #[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}