Files
dudu/server/internal/gateway/gateway.go
T
wangjia 40760aa884
ci / server (push) Failing after 14s
ci / design-tokens (push) Failing after 11s
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>
2026-06-12 00:38:37 +08:00

369 lines
9.8 KiB
Go
Raw 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
}
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 mono32000 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()
}