Files
gochat/internal/ws/broadcast_test.go
T
2026-06-04 15:44:48 +08:00

383 lines
10 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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")
}