b5ab92a57e
来自 xhigh code review 的正确性/健壮性修复,覆盖全部五端: - server:鉴权 fail-closed、计量交叉校验与配额扣穿处理、WS 网关并发与关闭顺序、 billing 行锁、redis Lua 过期与设备槽刷新、config 解析 - desktop:会话 epoch 防串话、WS 重连与 401 处理、api 客户端复用、统一 usePoll 轮询 - android:握手时序、请求头封装、账户状态派生、按需重组 - ios:finalize 宽限、串行采集、错误文案服务端优先、删除死代码 CommitController Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
249 lines
7.7 KiB
Go
249 lines
7.7 KiB
Go
package gateway
|
|
|
|
import (
|
|
"encoding/json"
|
|
"net/http/httptest"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/alicebob/miniredis/v2"
|
|
"github.com/gin-gonic/gin"
|
|
"github.com/glebarez/sqlite"
|
|
"github.com/gorilla/websocket"
|
|
"github.com/redis/go-redis/v9"
|
|
"gorm.io/gorm"
|
|
|
|
"dudu/server/internal/asr"
|
|
"dudu/server/internal/quota"
|
|
"dudu/server/internal/store"
|
|
"dudu/server/pkg/protocol"
|
|
)
|
|
|
|
func setup(t *testing.T) (*httptest.Server, *gorm.DB, *redis.Client) {
|
|
t.Helper()
|
|
gin.SetMode(gin.TestMode)
|
|
mr := miniredis.RunT(t)
|
|
rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()})
|
|
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := db.AutoMigrate(store.AllModels()...); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
db.Create(&store.User{ID: "u1", BalanceSeconds: 600})
|
|
|
|
mock := asr.NewMock()
|
|
mock.Script = []string{"你好,世界。"}
|
|
h := &Handler{Provider: mock, Quota: quota.New(rdb, db), RDB: rdb}
|
|
|
|
r := gin.New()
|
|
// 测试中跳过 JWT,直接注入 uid
|
|
r.GET("/v1/asr/stream", func(c *gin.Context) {
|
|
c.Set("auth.user_id", "u1")
|
|
h.HandleWS(c)
|
|
})
|
|
srv := httptest.NewServer(r)
|
|
t.Cleanup(srv.Close)
|
|
return srv, db, rdb
|
|
}
|
|
|
|
func dial(t *testing.T, srv *httptest.Server) *websocket.Conn {
|
|
t.Helper()
|
|
url := "ws" + strings.TrimPrefix(srv.URL, "http") + "/v1/asr/stream?device_id=dev1"
|
|
conn, _, err := websocket.DefaultDialer.Dial(url, nil)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
t.Cleanup(func() { conn.Close() })
|
|
return conn
|
|
}
|
|
|
|
func TestStreamE2E(t *testing.T) {
|
|
srv, db, _ := setup(t)
|
|
conn := dial(t, srv)
|
|
|
|
start, _ := json.Marshal(protocol.ClientMsg{Type: "start", SessionID: "s1", SampleRate: 16000, Format: "pcm16"})
|
|
if err := conn.WriteMessage(websocket.TextMessage, start); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
// 推 10 帧(1s 音频)→ mock 识别完 6 字全句
|
|
frame := make([]byte, protocol.FrameBytes)
|
|
for i := 0; i < 10; i++ {
|
|
if err := conn.WriteMessage(websocket.BinaryMessage, frame); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
stop, _ := json.Marshal(protocol.ClientMsg{Type: "stop", SessionID: "s1"})
|
|
if err := conn.WriteMessage(websocket.TextMessage, stop); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
var partials, finals []string
|
|
var lastUsage protocol.ServerMsg
|
|
deadline := time.Now().Add(5 * time.Second)
|
|
for time.Now().Before(deadline) {
|
|
_ = conn.SetReadDeadline(time.Now().Add(2 * time.Second))
|
|
var msg protocol.ServerMsg
|
|
if err := conn.ReadJSON(&msg); err != nil {
|
|
break
|
|
}
|
|
switch msg.Type {
|
|
case protocol.MsgPartial:
|
|
partials = append(partials, msg.Text)
|
|
case protocol.MsgFinal:
|
|
finals = append(finals, msg.Text)
|
|
case protocol.MsgUsage:
|
|
lastUsage = msg
|
|
case protocol.MsgError:
|
|
t.Fatalf("unexpected error frame: %+v", msg)
|
|
}
|
|
if lastUsage.SessionSeconds > 0 && len(finals) >= 2 {
|
|
break
|
|
}
|
|
}
|
|
|
|
if len(partials) == 0 {
|
|
t.Fatal("expect partial frames")
|
|
}
|
|
if got := strings.Join(finals, ""); got != "你好,世界。" {
|
|
t.Fatalf("final mismatch: %q", got)
|
|
}
|
|
if lastUsage.SessionSeconds != 1 {
|
|
t.Fatalf("want usage 1s, got %d", lastUsage.SessionSeconds)
|
|
}
|
|
// trial 先扣:余额应不变(600),试用剩 179
|
|
if lastUsage.BalanceSeconds != 600 || lastUsage.TrialRemaining != protocol.TrialDailySeconds-1 {
|
|
t.Fatalf("unexpected usage: %+v", lastUsage)
|
|
}
|
|
|
|
// 异步 settle 落库
|
|
var sess store.ASRSession
|
|
for i := 0; i < 50; i++ {
|
|
if err := db.First(&sess, "id = ?", "s1").Error; err == nil {
|
|
break
|
|
}
|
|
time.Sleep(20 * time.Millisecond)
|
|
}
|
|
if sess.AudioSeconds != 1 || sess.TrialPart != 1 || sess.BalancePart != 0 {
|
|
t.Fatalf("settle mismatch: %+v", sess)
|
|
}
|
|
// 计量交叉校验:provider 时间戳已落库,且不超过实收音频时长
|
|
if sess.ProviderMs <= 0 || sess.ProviderMs > int64(sess.AudioSeconds)*1000 {
|
|
t.Fatalf("provider ms cross-check mismatch: %+v", sess)
|
|
}
|
|
}
|
|
|
|
func TestQuotaExceededOnStart(t *testing.T) {
|
|
srv, db, _ := setup(t)
|
|
// 用尽试用 + 余额
|
|
db.Model(&store.User{}).Where("id = ?", "u1").Update("balance_seconds", 0)
|
|
mrConsumeAll(t, srv, db)
|
|
|
|
conn := dial(t, srv)
|
|
start, _ := json.Marshal(protocol.ClientMsg{Type: "start", SessionID: "s9", SampleRate: 16000})
|
|
_ = conn.WriteMessage(websocket.TextMessage, start)
|
|
var msg protocol.ServerMsg
|
|
_ = conn.SetReadDeadline(time.Now().Add(2 * time.Second))
|
|
if err := conn.ReadJSON(&msg); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if msg.Type != protocol.MsgError || msg.Code != protocol.ErrQuotaExceeded {
|
|
t.Fatalf("want QUOTA_EXCEEDED, got %+v", msg)
|
|
}
|
|
}
|
|
|
|
// mrConsumeAll 把 u1 的当日试用耗尽(经 quota 通道,保证键一致)。
|
|
// 单会话被截断在略低于 180s 处(截断帧不再计费,16A),故一会话不足以扣满 180s 试用;
|
|
// 此处循环开会话推音频,直到某个 usage 帧报告 TrialRemaining<=0 为止。
|
|
func mrConsumeAll(t *testing.T, srv *httptest.Server, db *gorm.DB) {
|
|
t.Helper()
|
|
frame := make([]byte, protocol.FrameBytes)
|
|
deadline := time.Now().Add(15 * time.Second)
|
|
for sess := 0; sess < 5 && time.Now().Before(deadline); sess++ {
|
|
conn := dial(t, srv)
|
|
sid := "s0_" + string(rune('a'+sess))
|
|
start, _ := json.Marshal(protocol.ClientMsg{Type: "start", SessionID: sid, SampleRate: 16000})
|
|
_ = conn.WriteMessage(websocket.TextMessage, start)
|
|
// 推 180s 音频(1800 帧)→ 会话在临界处截断
|
|
for i := 0; i < protocol.TrialDailySeconds*10; i++ {
|
|
if err := conn.WriteMessage(websocket.BinaryMessage, frame); err != nil {
|
|
break
|
|
}
|
|
}
|
|
stop, _ := json.Marshal(protocol.ClientMsg{Type: "stop", SessionID: sid})
|
|
_ = conn.WriteMessage(websocket.TextMessage, stop)
|
|
|
|
exhausted := false
|
|
for time.Now().Before(deadline) {
|
|
_ = conn.SetReadDeadline(time.Now().Add(2 * time.Second))
|
|
var msg protocol.ServerMsg
|
|
if err := conn.ReadJSON(&msg); err != nil {
|
|
break
|
|
}
|
|
if msg.Type == protocol.MsgUsage && msg.TrialRemaining <= 0 {
|
|
exhausted = true
|
|
break
|
|
}
|
|
}
|
|
conn.Close()
|
|
if exhausted {
|
|
return
|
|
}
|
|
}
|
|
t.Fatal("failed to exhaust trial via repeated sessions")
|
|
}
|
|
|
|
func TestSessionLimitTruncates(t *testing.T) {
|
|
srv, _, _ := setup(t)
|
|
conn := dial(t, srv)
|
|
start, _ := json.Marshal(protocol.ClientMsg{Type: "start", SessionID: "s2", SampleRate: 16000})
|
|
_ = conn.WriteMessage(websocket.TextMessage, start)
|
|
frame := make([]byte, protocol.FrameBytes)
|
|
// 181s 音频 → 触发截断
|
|
for i := 0; i < 1810; i++ {
|
|
if err := conn.WriteMessage(websocket.BinaryMessage, frame); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
sawLimit := false
|
|
deadline := time.Now().Add(5 * time.Second)
|
|
for time.Now().Before(deadline) && !sawLimit {
|
|
_ = conn.SetReadDeadline(time.Now().Add(2 * time.Second))
|
|
var msg protocol.ServerMsg
|
|
if err := conn.ReadJSON(&msg); err != nil {
|
|
break
|
|
}
|
|
if msg.Type == protocol.MsgError && msg.Code == protocol.ErrSessionLimit {
|
|
sawLimit = true
|
|
}
|
|
}
|
|
if !sawLimit {
|
|
t.Fatal("expect SESSION_LIMIT after 180s audio")
|
|
}
|
|
}
|
|
|
|
func TestConcurrentSecondSessionRejected(t *testing.T) {
|
|
srv, _, _ := setup(t)
|
|
c1 := dial(t, srv)
|
|
start1, _ := json.Marshal(protocol.ClientMsg{Type: "start", SessionID: "a1", SampleRate: 16000})
|
|
_ = c1.WriteMessage(websocket.TextMessage, start1)
|
|
// 等会话真正建立(收到任意帧前先推一帧拿 partial)
|
|
frame := make([]byte, protocol.FrameBytes)
|
|
_ = c1.WriteMessage(websocket.BinaryMessage, frame)
|
|
time.Sleep(100 * time.Millisecond)
|
|
|
|
c2 := dial(t, srv) // 同 device_id=dev1
|
|
start2, _ := json.Marshal(protocol.ClientMsg{Type: "start", SessionID: "a2", SampleRate: 16000})
|
|
_ = c2.WriteMessage(websocket.TextMessage, start2)
|
|
var msg protocol.ServerMsg
|
|
_ = c2.SetReadDeadline(time.Now().Add(2 * time.Second))
|
|
if err := c2.ReadJSON(&msg); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if msg.Type != protocol.MsgError || msg.Code != protocol.ErrRateLimited {
|
|
t.Fatalf("want RATE_LIMITED for concurrent session, got %+v", msg)
|
|
}
|
|
}
|