dudu MVP:五端语音输入法初始提交
ci / server (push) Failing after 14s
ci / design-tokens (push) Failing after 11s

- 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:
wangjia
2026-06-12 00:38:37 +08:00
commit 40760aa884
252 changed files with 40789 additions and 0 deletions
+368
View File
@@ -0,0 +1,368 @@
// 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()
}
+239
View File
@@ -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)
}
}