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) } }