335 lines
10 KiB
Go
335 lines
10 KiB
Go
package ws
|
||
|
||
import (
|
||
"context"
|
||
"fmt"
|
||
"testing"
|
||
"time"
|
||
|
||
"github.com/alicebob/miniredis/v2"
|
||
"github.com/redis/go-redis/v9"
|
||
"github.com/stretchr/testify/assert"
|
||
"github.com/stretchr/testify/require"
|
||
)
|
||
|
||
// --- 测试辅助函数 ---
|
||
// setupTypingTest 创建 miniredis 实例、redis.Client、BroadcastRelay 和 TypingTracker,
|
||
// 测试结束后需要调用 cleanup。
|
||
|
||
func setupTypingTest(t *testing.T) (*miniredis.Miniredis, *redis.Client, *TypingTracker, *mockHandler, func()) {
|
||
t.Helper()
|
||
mr := miniredis.RunT(t)
|
||
rdb := redis.NewClient(&redis.Options{
|
||
Addr: mr.Addr(),
|
||
})
|
||
handler := newMockHandler()
|
||
relay := NewBroadcastRelay(rdb, handler)
|
||
tracker := NewTypingTracker(rdb, relay)
|
||
cleanup := func() {
|
||
rdb.Close()
|
||
mr.Close()
|
||
}
|
||
return mr, rdb, tracker, handler, cleanup
|
||
}
|
||
|
||
// --- TypingTracker 构造函数测试 ---
|
||
|
||
func TestNewTypingTracker(t *testing.T) {
|
||
_, rdb, _, _, cleanup := setupTypingTest(t)
|
||
defer cleanup()
|
||
|
||
handler := newMockHandler()
|
||
relay := NewBroadcastRelay(rdb, handler)
|
||
tracker := NewTypingTracker(rdb, relay)
|
||
|
||
assert.NotNil(t, tracker, "TypingTracker 应成功创建")
|
||
assert.Equal(t, rdb, tracker.rdb, "Redis 客户端应正确设置")
|
||
assert.Equal(t, relay, tracker.relay, "BroadcastRelay 应正确设置")
|
||
}
|
||
|
||
// --- SetTypingOn 测试 ---
|
||
|
||
func TestSetTypingOn(t *testing.T) {
|
||
_, rdb, tracker, _, cleanup := setupTypingTest(t)
|
||
defer cleanup()
|
||
|
||
ctx := context.Background()
|
||
accountID := uint(1)
|
||
conversationID := uint(42)
|
||
performer := &Performer{
|
||
ID: 10,
|
||
Name: "Agent Zhang",
|
||
Type: "user",
|
||
AvatarURL: "https://example.com/avatar.jpg",
|
||
}
|
||
|
||
err := tracker.SetTypingOn(ctx, accountID, conversationID, performer)
|
||
require.NoError(t, err, "SetTypingOn 应成功执行")
|
||
|
||
// 验证 Redis 中存在 typing key
|
||
key := fmt.Sprintf(RedisKeyTyping, accountID, conversationID)
|
||
exists, err := rdb.Exists(ctx, key).Result()
|
||
require.NoError(t, err)
|
||
assert.Equal(t, int64(1), exists, "typing key 应存在")
|
||
|
||
// 验证 TTL 设置正确
|
||
ttl, err := rdb.TTL(ctx, key).Result()
|
||
require.NoError(t, err)
|
||
expectedTTL := time.Duration(TypingTTLSec) * time.Second
|
||
// TTL 应接近 TypingTTLSec(允许1秒误差)
|
||
assert.WithinDuration(t, time.Now().Add(expectedTTL), time.Now().Add(ttl), time.Second,
|
||
"TTL 应接近 TypingTTLSec")
|
||
}
|
||
|
||
func TestSetTypingOn_MultipleTyping(t *testing.T) {
|
||
_, rdb, tracker, _, cleanup := setupTypingTest(t)
|
||
defer cleanup()
|
||
|
||
ctx := context.Background()
|
||
|
||
// 在同一 account 下设置多个会话的 typing
|
||
err := tracker.SetTypingOn(ctx, 1, 10, &Performer{ID: 5, Name: "User A", Type: "user"})
|
||
require.NoError(t, err)
|
||
err = tracker.SetTypingOn(ctx, 1, 20, &Performer{ID: 6, Name: "User B", Type: "user"})
|
||
require.NoError(t, err)
|
||
err = tracker.SetTypingOn(ctx, 2, 10, &Performer{ID: 7, Name: "User C", Type: "user"})
|
||
require.NoError(t, err)
|
||
|
||
// 验证3个 typing key 都存在
|
||
for _, tc := range []struct{ acct, conv uint }{{1, 10}, {1, 20}, {2, 10}} {
|
||
key := fmt.Sprintf(RedisKeyTyping, tc.acct, tc.conv)
|
||
exists, err := rdb.Exists(ctx, key).Result()
|
||
require.NoError(t, err)
|
||
assert.Equal(t, int64(1), exists, "typing key 应存在")
|
||
}
|
||
}
|
||
|
||
func TestSetTypingOn_OverwriteExisting(t *testing.T) {
|
||
mr, _, tracker, _, cleanup := setupTypingTest(t)
|
||
defer cleanup()
|
||
|
||
ctx := context.Background()
|
||
|
||
// 第一次设置 typing
|
||
err := tracker.SetTypingOn(ctx, 1, 10, &Performer{ID: 5, Name: "First User", Type: "user"})
|
||
require.NoError(t, err)
|
||
|
||
// 推进一小段时间后再次设置(新的 performer)
|
||
mr.FastForward(1 * time.Second)
|
||
|
||
err = tracker.SetTypingOn(ctx, 1, 10, &Performer{ID: 6, Name: "Second User", Type: "user"})
|
||
require.NoError(t, err, "覆盖已有 typing 应成功")
|
||
|
||
// 验证 typing state 已更新为新的 performer
|
||
state, err := tracker.GetTypingState(ctx, 1, 10)
|
||
require.NoError(t, err)
|
||
require.NotNil(t, state, "typing state 应存在")
|
||
assert.Equal(t, uint(6), state.Performer.ID, "performer 应更新为第二个用户")
|
||
}
|
||
|
||
// --- SetTypingOff 测试 ---
|
||
|
||
func TestSetTypingOff(t *testing.T) {
|
||
_, rdb, tracker, _, cleanup := setupTypingTest(t)
|
||
defer cleanup()
|
||
|
||
ctx := context.Background()
|
||
accountID := uint(1)
|
||
conversationID := uint(42)
|
||
performer := &Performer{ID: 10, Name: "Agent Zhang", Type: "user"}
|
||
|
||
// 先设置 typing_on
|
||
err := tracker.SetTypingOn(ctx, accountID, conversationID, performer)
|
||
require.NoError(t, err)
|
||
|
||
// 设置 typing_off
|
||
err = tracker.SetTypingOff(ctx, accountID, conversationID, performer)
|
||
require.NoError(t, err, "SetTypingOff 应成功执行")
|
||
|
||
// 验证 Redis 中不再有 typing key
|
||
key := fmt.Sprintf(RedisKeyTyping, accountID, conversationID)
|
||
exists, err := rdb.Exists(ctx, key).Result()
|
||
require.NoError(t, err)
|
||
assert.Equal(t, int64(0), exists, "typing_off 后 key 应被删除")
|
||
}
|
||
|
||
func TestSetTypingOff_NotTyping(t *testing.T) {
|
||
_, _, tracker, _, cleanup := setupTypingTest(t)
|
||
defer cleanup()
|
||
|
||
ctx := context.Background()
|
||
|
||
// 对从未 typing_on 的会话设置 typing_off — 应不会出错
|
||
err := tracker.SetTypingOff(ctx, 1, 999, &Performer{ID: 5, Name: "User", Type: "user"})
|
||
require.NoError(t, err, "对不存在的 typing key 设置 off 不应报错")
|
||
}
|
||
|
||
// --- IsTyping 测试 ---
|
||
|
||
func TestIsTyping(t *testing.T) {
|
||
_, _, tracker, _, cleanup := setupTypingTest(t)
|
||
defer cleanup()
|
||
|
||
ctx := context.Background()
|
||
|
||
// 未设置 typing — 应返回 false
|
||
typing, err := tracker.IsTyping(ctx, 1, 10)
|
||
require.NoError(t, err)
|
||
assert.False(t, typing, "未设置 typing 应返回 false")
|
||
|
||
// 设置 typing_on
|
||
err = tracker.SetTypingOn(ctx, 1, 10, &Performer{ID: 5, Name: "User A", Type: "user"})
|
||
require.NoError(t, err)
|
||
|
||
// 现在应返回 true
|
||
typing, err = tracker.IsTyping(ctx, 1, 10)
|
||
require.NoError(t, err)
|
||
assert.True(t, typing, "设置 typing_on 后应返回 true")
|
||
|
||
// 设置 typing_off
|
||
err = tracker.SetTypingOff(ctx, 1, 10, &Performer{ID: 5, Name: "User A", Type: "user"})
|
||
require.NoError(t, err)
|
||
|
||
// 又应返回 false
|
||
typing, err = tracker.IsTyping(ctx, 1, 10)
|
||
require.NoError(t, err)
|
||
assert.False(t, typing, "设置 typing_off 后应返回 false")
|
||
}
|
||
|
||
func TestIsTyping_TTLExpiry(t *testing.T) {
|
||
mr, _, tracker, _, cleanup := setupTypingTest(t)
|
||
defer cleanup()
|
||
|
||
ctx := context.Background()
|
||
|
||
// 设置 typing_on
|
||
err := tracker.SetTypingOn(ctx, 1, 10, &Performer{ID: 5, Name: "User", Type: "user"})
|
||
require.NoError(t, err)
|
||
|
||
// 推进时间超过 TTL — typing 应自动过期
|
||
mr.FastForward(time.Duration(TypingTTLSec+2) * time.Second)
|
||
|
||
// 过期后 IsTyping 应返回 false
|
||
typing, err := tracker.IsTyping(ctx, 1, 10)
|
||
require.NoError(t, err)
|
||
assert.False(t, typing, "TTL 过期后 typing 应自动消失")
|
||
}
|
||
|
||
// --- GetTypingState 测试 ---
|
||
|
||
func TestGetTypingState(t *testing.T) {
|
||
_, _, tracker, _, cleanup := setupTypingTest(t)
|
||
defer cleanup()
|
||
|
||
ctx := context.Background()
|
||
accountID := uint(1)
|
||
conversationID := uint(42)
|
||
performer := &Performer{
|
||
ID: 10,
|
||
Name: "Agent Zhang",
|
||
Type: "user",
|
||
AvatarURL: "https://example.com/avatar.jpg",
|
||
}
|
||
|
||
// 设置 typing_on
|
||
err := tracker.SetTypingOn(ctx, accountID, conversationID, performer)
|
||
require.NoError(t, err)
|
||
|
||
// 获取 typing state
|
||
state, err := tracker.GetTypingState(ctx, accountID, conversationID)
|
||
require.NoError(t, err, "GetTypingState 应成功执行")
|
||
require.NotNil(t, state, "应返回 typingState")
|
||
|
||
assert.Equal(t, accountID, state.AccountID, "AccountID 应匹配")
|
||
assert.Equal(t, conversationID, state.ConversationID, "ConversationID 应匹配")
|
||
assert.Equal(t, performer.ID, state.Performer.ID, "Performer ID 应匹配")
|
||
assert.Equal(t, performer.Name, state.Performer.Name, "Performer Name 应匹配")
|
||
assert.Equal(t, performer.Type, state.Performer.Type, "Performer Type 应匹配")
|
||
// StartedAt 应接近当前时间
|
||
assert.WithinDuration(t, time.Now(), time.Unix(state.StartedAt, 0), 2*time.Second,
|
||
"StartedAt 应接近当前时间")
|
||
}
|
||
|
||
func TestGetTypingState_NotActive(t *testing.T) {
|
||
_, _, tracker, _, cleanup := setupTypingTest(t)
|
||
defer cleanup()
|
||
|
||
ctx := context.Background()
|
||
|
||
// 没有 active typing — 应返回 nil
|
||
state, err := tracker.GetTypingState(ctx, 1, 999)
|
||
require.NoError(t, err, "对不存在的 key 应不报错")
|
||
assert.Nil(t, state, "不存在的 typing state 应返回 nil")
|
||
}
|
||
|
||
func TestGetTypingState_TTLExpiry(t *testing.T) {
|
||
mr, _, tracker, _, cleanup := setupTypingTest(t)
|
||
defer cleanup()
|
||
|
||
ctx := context.Background()
|
||
|
||
// 设置 typing_on
|
||
err := tracker.SetTypingOn(ctx, 1, 10, &Performer{ID: 5, Name: "User", Type: "user"})
|
||
require.NoError(t, err)
|
||
|
||
// 推进时间超过 TTL
|
||
mr.FastForward(time.Duration(TypingTTLSec+2) * time.Second)
|
||
|
||
// 过期后 GetTypingState 应返回 nil
|
||
state, err := tracker.GetTypingState(ctx, 1, 10)
|
||
require.NoError(t, err)
|
||
assert.Nil(t, state, "TTL 过期后 typing state 应返回 nil")
|
||
}
|
||
|
||
// --- 集成测试:完整的 typing on/off 生命周期 ---
|
||
|
||
func TestTypingLifecycle(t *testing.T) {
|
||
mr, _, tracker, _, cleanup := setupTypingTest(t)
|
||
defer cleanup()
|
||
|
||
ctx := context.Background()
|
||
|
||
performer := &Performer{ID: 5, Name: "Agent Li", Type: "user"}
|
||
|
||
// 1. 初始状态:无 typing
|
||
typing, err := tracker.IsTyping(ctx, 1, 10)
|
||
require.NoError(t, err)
|
||
assert.False(t, typing, "初始状态不应有 typing")
|
||
|
||
// 2. 设置 typing_on
|
||
err = tracker.SetTypingOn(ctx, 1, 10, performer)
|
||
require.NoError(t, err)
|
||
|
||
typing, err = tracker.IsTyping(ctx, 1, 10)
|
||
require.NoError(t, err)
|
||
assert.True(t, typing, "typing_on 后应显示正在输入")
|
||
|
||
state, err := tracker.GetTypingState(ctx, 1, 10)
|
||
require.NoError(t, err)
|
||
require.NotNil(t, state)
|
||
assert.Equal(t, performer.ID, state.Performer.ID)
|
||
|
||
// 3. 等待一段时间但不超过 TTL
|
||
mr.FastForward(time.Duration(TypingTTLSec-2) * time.Second)
|
||
|
||
// typing 应仍然存在
|
||
typing, err = tracker.IsTyping(ctx, 1, 10)
|
||
require.NoError(t, err)
|
||
assert.True(t, typing, "TTL 未过期前 typing 应存在")
|
||
|
||
// 4. 设置 typing_off
|
||
err = tracker.SetTypingOff(ctx, 1, 10, performer)
|
||
require.NoError(t, err)
|
||
|
||
typing, err = tracker.IsTyping(ctx, 1, 10)
|
||
require.NoError(t, err)
|
||
assert.False(t, typing, "typing_off 后不应显示正在输入")
|
||
|
||
// 5. 再次设置 typing_on(模拟用户重新开始输入)
|
||
err = tracker.SetTypingOn(ctx, 1, 10, performer)
|
||
require.NoError(t, err)
|
||
|
||
typing, err = tracker.IsTyping(ctx, 1, 10)
|
||
require.NoError(t, err)
|
||
assert.True(t, typing, "重新设置 typing_on 应生效")
|
||
} |