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

335 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"
"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 应生效")
}