package ws import ( "context" "encoding/json" "sync" "testing" "time" "github.com/alicebob/miniredis/v2" "github.com/redis/go-redis/v9" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) // --- Mock MessageHandler --- // mockHandler 记录所有通过 SendToAccount 和 SendToRoom 发送的消息, // 用于验证 BroadcastRelay 是否正确转发 Redis Pub/Sub 消息。 type mockHandler struct { mu sync.Mutex accounts map[uint][]byte // accountID → 发送的数据 rooms map[string][]byte // room → 发送的数据 } func newMockHandler() *mockHandler { return &mockHandler{ accounts: make(map[uint][]byte), rooms: make(map[string][]byte), } } func (m *mockHandler) SendToAccount(accountID uint, data []byte) { m.mu.Lock() defer m.mu.Unlock() m.accounts[accountID] = data } func (m *mockHandler) SendToRoom(room string, data []byte) { m.mu.Lock() defer m.mu.Unlock() m.rooms[room] = data } func (m *mockHandler) getAccountData(accountID uint) []byte { m.mu.Lock() defer m.mu.Unlock() return m.accounts[accountID] } func (m *mockHandler) getRoomData(room string) []byte { m.mu.Lock() defer m.mu.Unlock() return m.rooms[room] } // --- 测试辅助函数 --- // setupBroadcastTest 创建一个 miniredis 实例、redis.Client 和 BroadcastRelay, // 返回这些对象以便在测试中使用。测试结束后需要调用 cleanup。 func setupBroadcastTest(t *testing.T) (*miniredis.Miniredis, *redis.Client, *mockHandler, func()) { t.Helper() mr := miniredis.RunT(t) rdb := redis.NewClient(&redis.Options{ Addr: mr.Addr(), }) handler := newMockHandler() cleanup := func() { rdb.Close() mr.Close() } return mr, rdb, handler, cleanup } // --- BroadcastRelay 构造函数测试 --- func TestNewBroadcastRelay(t *testing.T) { _, rdb, handler, cleanup := setupBroadcastTest(t) defer cleanup() relay := NewBroadcastRelay(rdb, handler) assert.NotNil(t, relay, "BroadcastRelay 应成功创建") assert.Equal(t, rdb, relay.rdb, "Redis 客户端应正确设置") assert.Equal(t, handler, relay.hub, "MessageHandler 应正确设置") assert.NotNil(t, relay.subs, "subs map 应已初始化") assert.Empty(t, relay.subs, "初始状态下不应有订阅") } // --- Publish 测试 --- func TestPublish(t *testing.T) { _, rdb, handler, cleanup := setupBroadcastTest(t) defer cleanup() relay := NewBroadcastRelay(rdb, handler) ctx := context.Background() msg := &WSMessage{ Event: EventMessageCreated, Data: map[string]any{"content": "hello"}, AccountID: 1, } // 先订阅 room 频道,验证消息能被接收 channel := RedisPrefixRoom + "test_room" sub := rdb.Subscribe(ctx, channel) defer sub.Close() // 等待订阅生效 _, err := sub.Receive(ctx) require.NoError(t, err, "订阅应成功") // 发布消息到指定房间 err = relay.Publish(ctx, "test_room", msg) require.NoError(t, err, "Publish 应成功执行") // 从订阅接收消息(带超时) msgCh := sub.Channel() select { case redisMsg := <-msgCh: assert.Equal(t, channel, redisMsg.Channel, "频道名应匹配") var wsMsg WSMessage err := json.Unmarshal([]byte(redisMsg.Payload), &wsMsg) require.NoError(t, err, "应能反序列化接收到的消息") assert.Equal(t, EventMessageCreated, wsMsg.Event, "事件类型应匹配") assert.Equal(t, uint(1), wsMsg.AccountID, "AccountID 应匹配") case <-time.After(2 * time.Second): t.Fatal("订阅端应在2秒内收到消息") } } func TestPublishAccount(t *testing.T) { _, rdb, handler, cleanup := setupBroadcastTest(t) defer cleanup() relay := NewBroadcastRelay(rdb, handler) ctx := context.Background() msg := &WSMessage{ Event: EventAgentOnline, Data: map[string]any{"agent_id": 42}, AccountID: 5, } // 先订阅 account 频道 channel := RedisPrefixAccount + "5" sub := rdb.Subscribe(ctx, channel) defer sub.Close() // 等待订阅生效 _, err := sub.Receive(ctx) require.NoError(t, err, "订阅应成功") // 发布消息到 account 频道 err = relay.PublishAccount(ctx, 5, msg) require.NoError(t, err, "PublishAccount 应成功执行") // 接收消息 msgCh := sub.Channel() select { case redisMsg := <-msgCh: assert.Equal(t, channel, redisMsg.Channel, "频道名应匹配") var wsMsg WSMessage err := json.Unmarshal([]byte(redisMsg.Payload), &wsMsg) require.NoError(t, err, "应能反序列化接收到的消息") assert.Equal(t, EventAgentOnline, wsMsg.Event, "事件类型应匹配") assert.Equal(t, uint(5), wsMsg.AccountID, "AccountID 应匹配") case <-time.After(2 * time.Second): t.Fatal("应在2秒内收到发布的消息") } } // --- Start / Stop 测试 --- func TestStartAndStop(t *testing.T) { _, rdb, handler, cleanup := setupBroadcastTest(t) defer cleanup() relay := NewBroadcastRelay(rdb, handler) ctx, cancel := context.WithCancel(context.Background()) defer cancel() // 启动 relay err := relay.Start(ctx) require.NoError(t, err, "Start 应成功执行") assert.Len(t, relay.subs, 2, "Start 应订阅2个模式频道(room:* 和 account:*)") // 发布消息,验证 relay 能通过 receiveLoop 传递到 handler msg := &WSMessage{ Event: EventMessageCreated, Data: map[string]any{"content": "test"}, AccountID: 10, } // 发布到 account 频道 err = relay.PublishAccount(ctx, 10, msg) require.NoError(t, err, "PublishAccount 应成功") // 等待 receiveLoop 处理消息(需要一些时间让 goroutine 接收) time.Sleep(200 * time.Millisecond) // 验证 handler 是否收到了消息 data := handler.getAccountData(10) if data != nil { var wsMsg WSMessage err := json.Unmarshal(data, &wsMsg) require.NoError(t, err, "handler 收到的数据应可反序列化") assert.Equal(t, EventMessageCreated, wsMsg.Event, "事件类型应匹配") } // 停止 relay err = relay.Stop() require.NoError(t, err, "Stop 应成功执行") assert.Empty(t, relay.subs, "Stop 后应清空所有订阅") } func TestStopCleansUpSubscriptions(t *testing.T) { _, rdb, handler, cleanup := setupBroadcastTest(t) defer cleanup() relay := NewBroadcastRelay(rdb, handler) ctx := context.Background() err := relay.Start(ctx) require.NoError(t, err, "Start 应成功") // 记录订阅数量 assert.Len(t, relay.subs, 2, "应有2个订阅") // 停止 err = relay.Stop() require.NoError(t, err, "Stop 应成功") assert.Empty(t, relay.subs, "Stop 应清空订阅 map") } // --- extractAccountIDFromChannel / extractRoomFromChannel 测试 --- func TestExtractAccountIDFromChannel(t *testing.T) { tests := []struct { name string channel string expected uint }{ {"有效account频道", RedisPrefixAccount + "42", 42}, {"有效account频道大ID", RedisPrefixAccount + "99999", 99999}, {"仅前缀无ID", RedisPrefixAccount, 0}, {"不相关频道", "other:channel", 0}, {"空字符串", "", 0}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { result := extractAccountIDFromChannel(tt.channel) assert.Equal(t, tt.expected, result, "应正确提取 accountID") }) } } func TestExtractRoomFromChannel(t *testing.T) { tests := []struct { name string channel string expected string }{ {"有效room频道", RedisPrefixRoom + "account_1", "account_1"}, {"有效room频道复杂名称", RedisPrefixRoom + "conversation_42", "conversation_42"}, {"仅前缀无room", RedisPrefixRoom, ""}, {"不相关频道", "other:channel", ""}, {"空字符串", "", ""}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { result := extractRoomFromChannel(tt.channel) assert.Equal(t, tt.expected, result, "应正确提取 room 名称") }) } } // --- handleRedisMessage 测试 --- func TestHandleRedisMessage_AccountChannel(t *testing.T) { _, rdb, handler, cleanup := setupBroadcastTest(t) defer cleanup() relay := NewBroadcastRelay(rdb, handler) msg := &WSMessage{ Event: EventMessageCreated, Data: map[string]any{"content": "hello"}, AccountID: 10, } data, err := json.Marshal(msg) require.NoError(t, err) // 模拟 Redis Message 在 account 频道 redisMsg := &redis.Message{ Channel: RedisPrefixAccount + "10", Payload: string(data), } relay.handleRedisMessage(redisMsg) // 验证 handler 收到了消息 received := handler.getAccountData(10) assert.NotNil(t, received, "应通过 SendToAccount 发送消息") var wsMsg WSMessage err = json.Unmarshal(received, &wsMsg) require.NoError(t, err, "应能反序列化") assert.Equal(t, EventMessageCreated, wsMsg.Event, "事件应匹配") } func TestHandleRedisMessage_RoomChannel(t *testing.T) { _, rdb, handler, cleanup := setupBroadcastTest(t) defer cleanup() relay := NewBroadcastRelay(rdb, handler) msg := &WSMessage{ Event: EventConversationCreated, Data: map[string]any{"id": 1}, } data, err := json.Marshal(msg) require.NoError(t, err) redisMsg := &redis.Message{ Channel: RedisPrefixRoom + "account_1", Payload: string(data), } relay.handleRedisMessage(redisMsg) received := handler.getRoomData("account_1") assert.NotNil(t, received, "应通过 SendToRoom 发送消息") var wsMsg WSMessage err = json.Unmarshal(received, &wsMsg) require.NoError(t, err) assert.Equal(t, EventConversationCreated, wsMsg.Event) } func TestHandleRedisMessage_InvalidPayload(t *testing.T) { _, rdb, handler, cleanup := setupBroadcastTest(t) defer cleanup() relay := NewBroadcastRelay(rdb, handler) // 无法反序列化的 payload — 应被静默忽略 redisMsg := &redis.Message{ Channel: RedisPrefixAccount + "10", Payload: "not-valid-json{{{", } relay.handleRedisMessage(redisMsg) // handler 不应收到任何消息 assert.Empty(t, handler.accounts, "无效 payload 不应转发到 handler") } func TestHandleRedisMessage_UnrecognizedChannel(t *testing.T) { _, rdb, handler, cleanup := setupBroadcastTest(t) defer cleanup() relay := NewBroadcastRelay(rdb, handler) msg := &WSMessage{Event: "test.event"} data, err := json.Marshal(msg) require.NoError(t, err) // 使用短字符串频道 — 长度 <= RedisPrefixRoom 长度(15) // extractAccountIDFromChannel 和 extractRoomFromChannel 都不会解析出有效目标 redisMsg := &redis.Message{ Channel: "short:chan", Payload: string(data), } relay.handleRedisMessage(redisMsg) // handler 不应收到任何消息(频道太短,无法解析为 account 或 room) assert.Empty(t, handler.accounts, "短频道不应发送到 account handler") assert.Empty(t, handler.rooms, "短频道不应发送到 room handler") }