// 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 } // wsWriteTimeout 单次写超时(16B):客户端不读时避免 WriteJSON 永久阻塞, // 致 pumpResults / 收尾路径卡死、goroutine 与连接泄漏。 const wsWriteTimeout = 5 * time.Second func (w *wsConn) sendJSON(v any) error { w.mu.Lock() defer w.mu.Unlock() _ = w.c.SetWriteDeadline(time.Now().Add(wsWriteTimeout)) 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 } 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,顺序显式(16B):先 conn.Close 拆连接——解开任何阻塞中的 // pumpResults 写,再 finish(其 Wait 已带超时)。避免 LIFO 让 finish 先跑、 // 而 Close 永不执行导致的 goroutine + 连接 + 设备槽泄漏。 defer func() { conn.Close() 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{}), } // Add 须 happen-before finish 的 Wait(WaitGroup 文档要求),放在启动 // goroutine 之前,避免 Add 与 Wait 竞态(16C)。 s.resultsDone.Add(1) 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 // consumeMu 串行化 Consume RPC 与 finish 取快照(16D):避免 finish 在 // 一笔在途 consumeDelta 已扣 Redis、尚未把 trialPart/balancePart 计入快照时 // 读到漏记的账。consumeMu 须在 mu 之外单独获取,且不可在持有 mu 时跨越 Redis 调用。 consumeMu sync.Mutex mu sync.Mutex audioBytes int // 计费音频字节累计;截断/死亡/扣穿后在 feed 处冻结不再增长 providerEndMs int64 // Provider 报告的最大句尾时间戳(ms),0=未提供 consumedSec int // 已增量扣减的秒数 trialPart int balancePart int truncated bool providerDead bool // provider 死亡(r.Err / Results 非正常关闭)后置位,停止计费(16A) quotaExhausted bool // 余额+试用扣穿后置位,已下发过一次 QUOTA_EXCEEDED(16E) finished bool done chan struct{} resultsDone sync.WaitGroup } // feed 转发音频帧并累计实收时长;超 180s 自动截断。 // 截断 / provider 死亡点之后到达的帧不再计入 audioBytes(不扣费): // 用户为"未识别的音频"付费是 bug——先判停止条件,再决定是否累计(16A)。 func (s *session) feed(pcm []byte) { s.mu.Lock() if s.finished || s.truncated || s.providerDead || s.quotaExhausted { // 已停止计费:丢弃此帧,不累计、不转发(扣穿后冻结,避免结算记录虚增秒数) s.mu.Unlock() return } // 先用"若计入本帧"的时长判是否超 180s overCap := (s.audioBytes+len(pcm)+31999)/32000 >= protocol.MaxSessionSeconds if overCap { // 本帧触发截断:冻结记账(本帧不计入),flush 定稿 s.truncated = true s.mu.Unlock() // 截断:flush 定稿,但保持会话记账状态直至客户端 stop / 连接收尾 _ = s.provider.Close() s.ws.sendErr(s.id, protocol.ErrSessionLimit) return } s.audioBytes += len(pcm) s.mu.Unlock() _ = s.provider.SendAudio(pcm) } func (s *session) audioSeconds() int { // 16kHz 16bit mono:32000 B/s,向上取整 return (s.audioBytes + 31999) / 32000 } // pumpResults Provider → 客户端下行泵。 // Add(1) 已移至 startSession(16C)。provider 死亡(r.Err 或 Results 非正常关闭) // 时置 providerDead,停止后续计费(16A)。 func (s *session) pumpResults() { defer s.resultsDone.Done() for r := range s.provider.Results() { if r.Err != nil { s.markProviderDead() 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}) } // Results 通道关闭:若并非由 finish/截断/扣穿主动 Close 触发,则 provider 自行 // 死亡(上游 EOF / 异常 task-finished),标记停止计费(16A)。 s.markProviderDead() } // markProviderDead 标记 provider 已死:feed 据此冻结 audioBytes(不再计入死亡点 // 之后的帧),consumeDelta 据此停止继续扣费;已识别部分仍由 finish 的 consumeFinal // 正常结算。由 finish/截断/扣穿主动 Close 触发的通道关闭也会走到这里,但此时计费已 // 另行冻结,置位无副作用。 func (s *session) markProviderDead() { s.mu.Lock() s.providerDead = true s.mu.Unlock() } // 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()) // 续期设备槽(17D):单会话墙钟可远超 4min TTL(低速发帧刷新读超时), // 不续期则 slot 先于会话过期,同设备第二路 start 会被错误放行。 _, _ = store.RefreshDeviceSlot(context.Background(), s.h.RDB, s.deviceID, s.id) } } } // consumeDelta 将"实收音频秒数 − 已扣秒数"差额扣减并广播余额。 // consumeMu 串行化与 finish 取快照(16D):持锁期间 finish 不会读到漏记的在途账。 func (s *session) consumeDelta(ctx context.Context) { s.consumeMu.Lock() defer s.consumeMu.Unlock() s.mu.Lock() // providerDead 后停止继续计费(16A):audioBytes 已在 feed 处冻结, // 已识别部分由 finish 的 consumeFinal 正常结算。 if s.finished || s.providerDead { s.mu.Unlock() return } delta := s.audioSeconds() - s.consumedSec if delta <= 0 { 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), }) // 扣穿检测(16E):余额与今日试用均已耗尽。优雅结束会话——不立刻掐断当句 // (Close 让 provider flush 当前句 final),下发一次 QUOTA_EXCEEDED 后由 finish 收尾。 if res.Exhausted() { s.mu.Lock() already := s.quotaExhausted s.quotaExhausted = true s.mu.Unlock() if !already { s.ws.sendErr(s.id, protocol.ErrQuotaExceeded) // 触发 provider flush 当前句尾 final 并结束;finish 串行收尾。 _ = s.provider.Close() } } } // finishWaitTimeout 收尾等待尾部 final 下发的最长时间(16B):客户端不读时 // pumpResults 可能卡在写上(已由 sendJSON 写超时兜底),此处再加一层超时, // 超时则继续收尾不无限等。 const finishWaitTimeout = 3 * time.Second // finish 结束会话:flush final → 补扣 → usage 帧 → 异步 settle → 记窗口 → 释放槽。 // 注意 done/consumeMu 顺序(16D):先停 usageLoop 并 drain 在途 consumeDelta, // 再读 trialPart/balancePart 快照,避免漏记在途扣费。 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.waitResults() // 尾部 final 全部下发后再收尾(带超时,16B) // 先标记 finished 并停 usageLoop,再 drain 在途 consumeDelta(16D): // 取 consumeMu 会等待任何已过 finished 检查、正等 Redis 返回的 consumeDelta // 完成并把 trialPart/balancePart 计入,从而快照不漏账。 s.mu.Lock() s.finished = true close(s.done) s.mu.Unlock() s.consumeMu.Lock() // drain:等当前在途 consume(若有)完成 s.consumeMu.Unlock() s.mu.Lock() 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) } }() } // waitResults 等 pumpResults 退出(尾部 final 下发完),最多 finishWaitTimeout(16B)。 // 超时则放弃等待继续收尾——避免客户端不读时永久阻塞收尾路径。 func (s *session) waitResults() { done := make(chan struct{}) go func() { s.resultsDone.Wait() close(done) }() select { case <-done: case <-time.After(finishWaitTimeout): slog.Warn("finish: results wait timed out", "session", s.id) } } // consumeFinal 结束时补扣差额(不足 2s 的短会话由此兜底)。 // 由 finish 在持有 consumeMu drain 后调用,且 usageLoop 已停,天然与增量扣减串行。 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() }