Files
wangjia b5ab92a57e
ci / server (push) Failing after 13s
ci / design-tokens (push) Failing after 11s
fix: 应用 xhigh 代码评审的跨端修复
来自 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>
2026-06-13 11:50:08 +08:00

461 lines
15 KiB
Go
Raw Permalink 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 gateway WS 流式识别网关:客户端 ↔ ASR Provider 的中继。
// 协议见 pkg/protocoldoc/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 的 WaitWaitGroup 文档要求),放在启动
// 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_EXCEEDED16E
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 mono32000 B/s,向上取整
return (s.audioBytes + 31999) / 32000
}
// pumpResults Provider → 客户端下行泵。
// Add(1) 已移至 startSession16C)。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 在途 consumeDelta16D):
// 取 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 下发完),最多 finishWaitTimeout16B)。
// 超时则放弃等待继续收尾——避免客户端不读时永久阻塞收尾路径。
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()
}