383 lines
10 KiB
Go
383 lines
10 KiB
Go
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")
|
||
} |