dudu MVP:五端语音输入法初始提交
- server:Go 网关(WS 流式识别中继/计费配额/微信登录支付 mock/反馈/埋点),gummy provider 已真实联调 - desktop:Tauri 2(全局快捷键 push-to-talk/浮层/托盘/设置/登录购买/反馈/首启引导) - android:Compose 主 App + IME(键盘内录音直传) - ios:App + 键盘扩展(1A spike 实证键盘内不可录音,走 deep link 听写) - design/design-pipeline:设计系统 + token 导出 iOS/Android 主题 - doc:前后端设计文档(HTML);web:官网宣传页;todo:任务看板 Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,368 @@
|
||||
// Package gateway WS 流式识别网关:客户端 ↔ ASR Provider 的中继。
|
||||
// 协议见 pkg/protocol(doc/backend-architecture.html 第四章),
|
||||
// 计费/限制策略见第五章与第九章:
|
||||
// - start 预检:配额(QUOTA_EXCEEDED)、设备 30min 窗口(RATE_LIMITED)、单设备单路
|
||||
// - 识别中每 2s 增量扣减并下发 usage 帧
|
||||
// - 单会话 ≤180s 自动截断定稿(SESSION_LIMIT,内容不丢)
|
||||
// - 结束:补扣、settle 落库(异步)、记录窗口用量、释放设备槽
|
||||
package gateway
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gorilla/websocket"
|
||||
"github.com/redis/go-redis/v9"
|
||||
|
||||
"dudu/server/internal/asr"
|
||||
"dudu/server/internal/auth"
|
||||
"dudu/server/internal/quota"
|
||||
"dudu/server/internal/store"
|
||||
"dudu/server/pkg/protocol"
|
||||
)
|
||||
|
||||
type Handler struct {
|
||||
Provider asr.Provider
|
||||
Quota *quota.Manager
|
||||
RDB *redis.Client
|
||||
}
|
||||
|
||||
// meterToleranceMs 计量交叉校验容差:provider 句尾时间戳超出实收音频时长
|
||||
// 该值即告警(少报正常——尾段静音无句尾)。
|
||||
const meterToleranceMs = 3000
|
||||
|
||||
var upgrader = websocket.Upgrader{
|
||||
ReadBufferSize: 8192,
|
||||
WriteBufferSize: 4096,
|
||||
CheckOrigin: func(*http.Request) bool { return true }, // 原生客户端,无浏览器 Origin
|
||||
}
|
||||
|
||||
// wsConn 串行化写(gorilla 不允许并发写)。
|
||||
type wsConn struct {
|
||||
mu sync.Mutex
|
||||
c *websocket.Conn
|
||||
}
|
||||
|
||||
func (w *wsConn) sendJSON(v any) error {
|
||||
w.mu.Lock()
|
||||
defer w.mu.Unlock()
|
||||
return w.c.WriteJSON(v)
|
||||
}
|
||||
|
||||
func (w *wsConn) sendErr(sessionID, code string) {
|
||||
_ = w.sendJSON(protocol.ServerMsg{
|
||||
Type: protocol.MsgError, SessionID: sessionID,
|
||||
Code: code, Message: protocol.ErrMessage[code],
|
||||
})
|
||||
}
|
||||
|
||||
// HandleWS GET /v1/asr/stream(需 JWT;设备标识取 X-Device-ID)。
|
||||
func (h *Handler) HandleWS(c *gin.Context) {
|
||||
uid := auth.UserID(c)
|
||||
deviceID := c.GetHeader("X-Device-ID")
|
||||
if deviceID == "" {
|
||||
deviceID = c.Query("device_id")
|
||||
}
|
||||
if deviceID == "" {
|
||||
c.JSON(http.StatusBadRequest, protocol.NewAPIError(protocol.ErrBadRequest))
|
||||
return
|
||||
}
|
||||
|
||||
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
defer conn.Close()
|
||||
ws := &wsConn{c: conn}
|
||||
|
||||
conn.SetPongHandler(func(string) error {
|
||||
return conn.SetReadDeadline(time.Now().Add(90 * time.Second))
|
||||
})
|
||||
_ = conn.SetReadDeadline(time.Now().Add(90 * time.Second))
|
||||
|
||||
var sess *session // 同连接串行多次会话
|
||||
defer func() {
|
||||
if sess != nil {
|
||||
sess.finish(context.Background(), false)
|
||||
}
|
||||
}()
|
||||
|
||||
for {
|
||||
mt, data, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
_ = conn.SetReadDeadline(time.Now().Add(90 * time.Second))
|
||||
|
||||
switch mt {
|
||||
case websocket.BinaryMessage:
|
||||
if sess != nil {
|
||||
sess.feed(data)
|
||||
}
|
||||
case websocket.TextMessage:
|
||||
var msg protocol.ClientMsg
|
||||
if err := json.Unmarshal(data, &msg); err != nil {
|
||||
ws.sendErr("", protocol.ErrBadRequest)
|
||||
continue
|
||||
}
|
||||
switch msg.Type {
|
||||
case protocol.MsgStart:
|
||||
if sess != nil { // 同连接重复 start:结束旧会话
|
||||
sess.finish(c.Request.Context(), false)
|
||||
}
|
||||
sess = h.startSession(c.Request.Context(), ws, uid, deviceID, msg)
|
||||
case protocol.MsgStop, protocol.MsgCancel:
|
||||
if sess != nil && sess.id == msg.SessionID {
|
||||
sess.finish(c.Request.Context(), msg.Type == protocol.MsgCancel)
|
||||
sess = nil
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// startSession 执行 start 预检并建立 Provider 会话;失败时返回 nil(已下发 error 帧)。
|
||||
func (h *Handler) startSession(ctx context.Context, ws *wsConn, uid, deviceID string, msg protocol.ClientMsg) *session {
|
||||
now := time.Now()
|
||||
|
||||
// ① 设备 30min 窗口:次数 + 累计时长
|
||||
if ok, err := store.AllowSession(ctx, h.RDB, deviceID, now); err != nil || !ok {
|
||||
ws.sendErr(msg.SessionID, protocol.ErrRateLimited)
|
||||
return nil
|
||||
}
|
||||
if full, err := store.AudioWindowExhausted(ctx, h.RDB, deviceID, now); err != nil || full {
|
||||
ws.sendErr(msg.SessionID, protocol.ErrRateLimited)
|
||||
return nil
|
||||
}
|
||||
// ② 单设备单路
|
||||
if ok, err := store.AcquireDeviceSlot(ctx, h.RDB, deviceID, msg.SessionID); err != nil || !ok {
|
||||
ws.sendErr(msg.SessionID, protocol.ErrRateLimited)
|
||||
return nil
|
||||
}
|
||||
// ③ 配额预检
|
||||
ok, _, err := h.Quota.Precheck(ctx, uid)
|
||||
if err != nil || !ok {
|
||||
_ = store.ReleaseDeviceSlot(ctx, h.RDB, deviceID, msg.SessionID)
|
||||
ws.sendErr(msg.SessionID, protocol.ErrQuotaExceeded)
|
||||
return nil
|
||||
}
|
||||
// ④ Provider 会话
|
||||
ps, err := h.Provider.StartSession(ctx, asr.SessionConfig{
|
||||
SampleRate: msg.SampleRate, SessionID: msg.SessionID,
|
||||
})
|
||||
if err != nil {
|
||||
_ = store.ReleaseDeviceSlot(ctx, h.RDB, deviceID, msg.SessionID)
|
||||
ws.sendErr(msg.SessionID, protocol.ErrASRUnavailable)
|
||||
return nil
|
||||
}
|
||||
|
||||
s := &session{
|
||||
h: h, ws: ws, id: msg.SessionID, uid: uid, deviceID: deviceID,
|
||||
provider: ps, started: now, done: make(chan struct{}),
|
||||
}
|
||||
go s.pumpResults()
|
||||
go s.usageLoop()
|
||||
return s
|
||||
}
|
||||
|
||||
type session struct {
|
||||
h *Handler
|
||||
ws *wsConn
|
||||
id string
|
||||
uid string
|
||||
deviceID string
|
||||
provider asr.Session
|
||||
started time.Time
|
||||
|
||||
mu sync.Mutex
|
||||
audioBytes int
|
||||
providerEndMs int64 // Provider 报告的最大句尾时间戳(ms),0=未提供
|
||||
consumedSec int // 已增量扣减的秒数
|
||||
trialPart int
|
||||
balancePart int
|
||||
truncated bool
|
||||
finished bool
|
||||
done chan struct{}
|
||||
resultsDone sync.WaitGroup
|
||||
}
|
||||
|
||||
// feed 转发音频帧并累计实收时长;超 180s 自动截断。
|
||||
func (s *session) feed(pcm []byte) {
|
||||
s.mu.Lock()
|
||||
if s.finished {
|
||||
s.mu.Unlock()
|
||||
return
|
||||
}
|
||||
s.audioBytes += len(pcm)
|
||||
overCap := s.audioSeconds() >= protocol.MaxSessionSeconds
|
||||
s.mu.Unlock()
|
||||
|
||||
if overCap {
|
||||
s.mu.Lock()
|
||||
already := s.truncated
|
||||
s.truncated = true
|
||||
s.mu.Unlock()
|
||||
if !already {
|
||||
// 截断:flush 定稿,但保持会话记账状态直至客户端 stop / 连接收尾
|
||||
_ = s.provider.Close()
|
||||
s.ws.sendErr(s.id, protocol.ErrSessionLimit)
|
||||
}
|
||||
return
|
||||
}
|
||||
_ = s.provider.SendAudio(pcm)
|
||||
}
|
||||
|
||||
func (s *session) audioSeconds() int {
|
||||
// 16kHz 16bit mono:32000 B/s,向上取整
|
||||
return (s.audioBytes + 31999) / 32000
|
||||
}
|
||||
|
||||
// pumpResults Provider → 客户端下行泵。
|
||||
func (s *session) pumpResults() {
|
||||
s.resultsDone.Add(1)
|
||||
defer s.resultsDone.Done()
|
||||
for r := range s.provider.Results() {
|
||||
if r.Err != nil {
|
||||
s.ws.sendErr(s.id, protocol.ErrASRUnavailable)
|
||||
return
|
||||
}
|
||||
if r.EndTimeMs > 0 {
|
||||
s.mu.Lock()
|
||||
if r.EndTimeMs > s.providerEndMs {
|
||||
s.providerEndMs = r.EndTimeMs
|
||||
}
|
||||
s.mu.Unlock()
|
||||
}
|
||||
typ := protocol.MsgPartial
|
||||
if r.IsFinal {
|
||||
typ = protocol.MsgFinal
|
||||
}
|
||||
_ = s.ws.sendJSON(protocol.ServerMsg{Type: typ, SessionID: s.id, Text: r.Text})
|
||||
}
|
||||
}
|
||||
|
||||
// usageLoop 每 2s 增量扣减并下发 usage 帧。
|
||||
func (s *session) usageLoop() {
|
||||
t := time.NewTicker(2 * time.Second)
|
||||
defer t.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-s.done:
|
||||
return
|
||||
case <-t.C:
|
||||
s.consumeDelta(context.Background())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// consumeDelta 将"实收音频秒数 − 已扣秒数"差额扣减并广播余额。
|
||||
func (s *session) consumeDelta(ctx context.Context) {
|
||||
s.mu.Lock()
|
||||
delta := s.audioSeconds() - s.consumedSec
|
||||
if delta <= 0 || s.finished {
|
||||
s.mu.Unlock()
|
||||
return
|
||||
}
|
||||
s.consumedSec += delta
|
||||
s.mu.Unlock()
|
||||
|
||||
res, err := s.h.Quota.Consume(ctx, s.uid, delta)
|
||||
if err != nil {
|
||||
slog.Error("quota consume failed", "err", err, "session", s.id)
|
||||
return
|
||||
}
|
||||
s.mu.Lock()
|
||||
s.trialPart += res.TrialPart
|
||||
s.balancePart += res.BalancePart
|
||||
s.mu.Unlock()
|
||||
|
||||
_ = s.ws.sendJSON(protocol.ServerMsg{
|
||||
Type: protocol.MsgUsage, SessionID: s.id,
|
||||
SessionSeconds: s.consumedSec,
|
||||
BalanceSeconds: res.BalanceSeconds,
|
||||
TrialRemaining: max(0, protocol.TrialDailySeconds-res.TrialUsedToday),
|
||||
})
|
||||
}
|
||||
|
||||
// finish 结束会话:flush final → 补扣 → usage 帧 → 异步 settle → 记窗口 → 释放槽。
|
||||
func (s *session) finish(ctx context.Context, canceled bool) {
|
||||
s.mu.Lock()
|
||||
if s.finished {
|
||||
s.mu.Unlock()
|
||||
return
|
||||
}
|
||||
s.mu.Unlock()
|
||||
|
||||
_ = s.provider.Close()
|
||||
s.resultsDone.Wait() // 尾部 final 全部下发后再收尾
|
||||
|
||||
s.mu.Lock()
|
||||
s.finished = true
|
||||
close(s.done)
|
||||
seconds := s.audioSeconds()
|
||||
audioMs := int64(s.audioBytes) * 1000 / 32000
|
||||
providerMs := s.providerEndMs
|
||||
s.mu.Unlock()
|
||||
|
||||
// 计量交叉校验:provider 句尾时间戳不应显著超出实收音频时长
|
||||
// (provider 只会少报——静音尾段无句尾;超出说明计量或对接异常)
|
||||
if providerMs > audioMs+meterToleranceMs {
|
||||
slog.Warn("metering cross-check divergence",
|
||||
"session", s.id, "provider", s.h.Provider.Name(),
|
||||
"audio_ms", audioMs, "provider_ms", providerMs)
|
||||
}
|
||||
|
||||
s.consumeFinal(ctx, seconds)
|
||||
|
||||
snap, err := s.h.Quota.Get(ctx, s.uid)
|
||||
if err == nil {
|
||||
_ = s.ws.sendJSON(protocol.ServerMsg{
|
||||
Type: protocol.MsgUsage, SessionID: s.id,
|
||||
SessionSeconds: seconds,
|
||||
BalanceSeconds: snap.BalanceSeconds,
|
||||
TrialRemaining: max(0, protocol.TrialDailySeconds-snap.TrialUsedToday),
|
||||
})
|
||||
}
|
||||
|
||||
_ = store.RecordAudioSeconds(ctx, s.h.RDB, s.deviceID, seconds, time.Now())
|
||||
_ = store.ReleaseDeviceSlot(ctx, s.h.RDB, s.deviceID, s.id)
|
||||
|
||||
s.mu.Lock()
|
||||
sess := store.ASRSession{
|
||||
ID: s.id, UserID: s.uid, DeviceID: s.deviceID,
|
||||
AudioSeconds: seconds, TrialPart: s.trialPart, BalancePart: s.balancePart,
|
||||
ProviderMs: providerMs,
|
||||
Provider: s.h.Provider.Name(), Canceled: canceled, CreatedAt: s.started,
|
||||
}
|
||||
s.mu.Unlock()
|
||||
go func() { // 异步落库,不阻塞下一次按键
|
||||
if err := s.h.Quota.SettleSession(context.Background(), sess); err != nil {
|
||||
slog.Error("settle session failed", "err", err, "session", sess.ID)
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
// consumeFinal 结束时补扣差额(不足 2s 的短会话由此兜底)。
|
||||
func (s *session) consumeFinal(ctx context.Context, seconds int) {
|
||||
s.mu.Lock()
|
||||
delta := seconds - s.consumedSec
|
||||
s.consumedSec = seconds
|
||||
s.mu.Unlock()
|
||||
if delta <= 0 {
|
||||
return
|
||||
}
|
||||
res, err := s.h.Quota.Consume(ctx, s.uid, delta)
|
||||
if err != nil {
|
||||
slog.Error("final consume failed", "err", err, "session", s.id)
|
||||
return
|
||||
}
|
||||
s.mu.Lock()
|
||||
s.trialPart += res.TrialPart
|
||||
s.balancePart += res.BalancePart
|
||||
s.mu.Unlock()
|
||||
}
|
||||
@@ -0,0 +1,239 @@
|
||||
package gateway
|
||||
|
||||
import (
|
||||
"context"
|
||||
"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 通道,保证键一致)。
|
||||
func mrConsumeAll(t *testing.T, srv *httptest.Server, db *gorm.DB) {
|
||||
t.Helper()
|
||||
conn := dial(t, srv)
|
||||
start, _ := json.Marshal(protocol.ClientMsg{Type: "start", SessionID: "s0", SampleRate: 16000})
|
||||
_ = conn.WriteMessage(websocket.TextMessage, start)
|
||||
frame := make([]byte, protocol.FrameBytes)
|
||||
// 180s 音频 = 1800 帧
|
||||
for i := 0; i < protocol.TrialDailySeconds*10; i++ {
|
||||
if err := conn.WriteMessage(websocket.BinaryMessage, frame); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
stop, _ := json.Marshal(protocol.ClientMsg{Type: "stop", SessionID: "s0"})
|
||||
_ = conn.WriteMessage(websocket.TextMessage, stop)
|
||||
// 读到连接收尾的 usage 帧为止
|
||||
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
|
||||
}
|
||||
if msg.Type == protocol.MsgUsage && msg.SessionSeconds >= protocol.TrialDailySeconds {
|
||||
break
|
||||
}
|
||||
}
|
||||
conn.Close()
|
||||
_ = context.Background()
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user