diff --git a/server/cmd/latencyprobe/main.go b/server/cmd/latencyprobe/main.go new file mode 100644 index 0000000..9922fc8 --- /dev/null +++ b/server/cmd/latencyprobe/main.go @@ -0,0 +1,267 @@ +// latencyprobe:12B 延迟验收工具——走本地网关全链路(WS 网关 → gummy)实测延迟分布。 +// +// 复刻桌面客户端行为:单条持久 WS 连接,串行多次会话,音频实时推流(100ms/帧)。 +// +// 测量口径对齐桌面端 dictation.rs: +// +// first_partial_ms = start 帧发出 → 本会话首个 partial 收到 +// (桌面端另含麦克风启动耗时 audio.start_ms,此处不含) +// release_to_commit_ms = stop 帧发出 → 本会话尾部 final 收到(文本完整落定) +// (桌面端另含文字注入耗时,通常 <10ms,此处不含; +// 若 stop 后先到 usage 帧,桌面端会提前以 partial 文本注入, +// 该情形单独计数为 usage_first) +// +// 用法(服务端需已以 gummy provider 启动): +// +// go run ./cmd/latencyprobe -wav /path/to/16k-mono-16bit.wav -n 12 +package main + +import ( + "encoding/binary" + "encoding/json" + "flag" + "fmt" + "net/http" + "os" + "sort" + "time" + + "github.com/gorilla/websocket" + + "dudu/server/internal/auth" + "dudu/server/internal/store" + "dudu/server/pkg/protocol" +) + +type event struct { + msg protocol.ServerMsg + at time.Time +} + +func main() { + addr := flag.String("addr", "localhost:8080", "服务端地址") + wavPath := flag.String("wav", "", "16kHz/16bit/mono wav 文件路径") + n := flag.Int("n", 12, "会话次数(注意 30 分钟滑动窗口上限 30 次/设备)") + dsn := flag.String("dsn", "host=localhost user=dudu password=dudu dbname=dudu port=5432 sslmode=disable TimeZone=Asia/Shanghai", "postgres dsn(用于确保探针用户存在且余额充足)") + secret := flag.String("secret", "dev-secret-change-me", "JWT secret(须与服务端一致)") + userID := flag.String("user", "latency-probe", "探针用户 ID") + deviceID := flag.String("device", "probe-dev", "探针设备 ID") + flag.Parse() + if *wavPath == "" { + fmt.Fprintln(os.Stderr, "用法: latencyprobe -wav test.wav [-n 12]") + os.Exit(2) + } + + pcm, err := readWavData(*wavPath) + if err != nil { + fmt.Fprintln(os.Stderr, "读 wav 失败:", err) + os.Exit(1) + } + audioSec := float64(len(pcm)) / float64(protocol.SampleRate*2) + fmt.Printf("音频 %.1fs · %d 次会话 · 单连接串行 · 实时推流(100ms/帧)\n", audioSec, *n) + + // 确保探针用户存在且余额充足(试用 180s/天 不够跑长测)。 + db, err := store.Open(*dsn) + if err != nil { + fmt.Fprintln(os.Stderr, "postgres 连接失败:", err) + os.Exit(1) + } + var u store.User + db.Where(store.User{ID: *userID}).FirstOrCreate(&u) + if err := db.Model(&store.User{}).Where("id = ?", *userID).Update("balance_seconds", 360000).Error; err != nil { + fmt.Fprintln(os.Stderr, "更新探针用户余额失败:", err) + os.Exit(1) + } + + token, err := auth.NewJWT(*secret, time.Hour, nil).Sign(*userID) + if err != nil { + fmt.Fprintln(os.Stderr, "签发 JWT 失败:", err) + os.Exit(1) + } + + // 单条持久连接(桌面端同款):会话串行复用,读循环常驻。 + url := fmt.Sprintf("ws://%s/v1/asr/stream?device_id=%s", *addr, *deviceID) + hdr := http.Header{"Authorization": {"Bearer " + token}} + conn, _, err := websocket.DefaultDialer.Dial(url, hdr) + if err != nil { + fmt.Fprintln(os.Stderr, "dial 失败:", err) + os.Exit(1) + } + defer conn.Close() + events := make(chan event, 256) + go func() { + defer close(events) + for { + var msg protocol.ServerMsg + if err := conn.ReadJSON(&msg); err != nil { + return + } + events <- event{msg, time.Now()} + } + }() + + var firstPartials, releaseToCommits []time.Duration + usageFirst := 0 + for i := 0; i < *n; i++ { + sid := fmt.Sprintf("probe-%d-%d", time.Now().UnixNano(), i) + r, err := runSession(conn, events, sid, pcm) + if err != nil { + fmt.Printf(" #%02d 失败: %v\n", i+1, err) + } else { + firstPartials = append(firstPartials, r.firstPartial) + releaseToCommits = append(releaseToCommits, r.releaseToCommit) + note := "" + if r.usageBeforeFinal { + usageFirst++ + note = " ⚠️usage先到" + } + fmt.Printf(" #%02d first_partial=%5dms release_to_commit=%4dms %q%s\n", + i+1, r.firstPartial.Milliseconds(), r.releaseToCommit.Milliseconds(), truncate(r.text, 24), note) + } + time.Sleep(800 * time.Millisecond) // 会话间隔(无论成败),模拟自然停顿 + } + + if len(firstPartials) == 0 { + fmt.Fprintln(os.Stderr, "无有效样本") + os.Exit(1) + } + fmt.Println() + report("first_partial ", firstPartials, 500*time.Millisecond) + report("release_to_commit ", releaseToCommits, 300*time.Millisecond) + if usageFirst > 0 { + fmt.Printf("⚠️ %d/%d 次会话 stop 后 usage 帧先于尾部 final 到达(桌面端将以 partial 文本提前注入,丢失 final 修正)\n", + usageFirst, len(firstPartials)) + } +} + +type sessionResult struct { + firstPartial time.Duration + releaseToCommit time.Duration + text string + usageBeforeFinal bool +} + +// runSession 在共享连接上执行一次会话:start → 实时推流 → stop → 等尾部 final。 +// 事件按 session_id 过滤,前一会话的迟到帧不串扰。 +func runSession(conn *websocket.Conn, events <-chan event, sid string, pcm []byte) (r sessionResult, err error) { + start, _ := json.Marshal(protocol.ClientMsg{Type: protocol.MsgStart, SessionID: sid, SampleRate: protocol.SampleRate, Format: "pcm16"}) + t0 := time.Now() + if err := conn.WriteMessage(websocket.TextMessage, start); err != nil { + return r, fmt.Errorf("start: %w", err) + } + + handle := func(ev event, tStop *time.Time) error { + if ev.msg.SessionID != sid { + return nil // 旧会话迟到帧 + } + switch ev.msg.Type { + case protocol.MsgPartial: + if r.firstPartial == 0 { + r.firstPartial = ev.at.Sub(t0) + } + case protocol.MsgFinal: + r.text += ev.msg.Text + if tStop != nil && r.releaseToCommit == 0 { + r.releaseToCommit = ev.at.Sub(*tStop) + } + case protocol.MsgUsage: + if tStop != nil && r.releaseToCommit == 0 { + r.usageBeforeFinal = true // 桌面端此刻已提前注入 partial 文本 + } + case protocol.MsgError: + return fmt.Errorf("server error: %s %s", ev.msg.Code, ev.msg.Message) + } + return nil + } + + // 实时推流 + 并行收包。 + ticker := time.NewTicker(protocol.FrameMillis * time.Millisecond) + defer ticker.Stop() + off := 0 + for off < len(pcm) { + select { + case <-ticker.C: + end := min(off+protocol.FrameBytes, len(pcm)) + if err := conn.WriteMessage(websocket.BinaryMessage, pcm[off:end]); err != nil { + return r, fmt.Errorf("audio: %w", err) + } + off = end + case ev, ok := <-events: + if !ok { + return r, fmt.Errorf("连接已断开") + } + if err := handle(ev, nil); err != nil { + return r, err + } + } + } + + stop, _ := json.Marshal(protocol.ClientMsg{Type: protocol.MsgStop, SessionID: sid}) + tStop := time.Now() + if err := conn.WriteMessage(websocket.TextMessage, stop); err != nil { + return r, fmt.Errorf("stop: %w", err) + } + + // 收尾:等尾部 final(文本完整落定),10s 兜底。 + deadline := time.After(10 * time.Second) + for r.releaseToCommit == 0 { + select { + case ev, ok := <-events: + if !ok { + return r, fmt.Errorf("连接已断开") + } + if err := handle(ev, &tStop); err != nil { + return r, err + } + case <-deadline: + return r, fmt.Errorf("收尾超时:stop 后 10s 未收到尾部 final") + } + } + if r.firstPartial == 0 { + return r, fmt.Errorf("全程未收到 partial") + } + return r, nil +} + +func report(name string, ds []time.Duration, target time.Duration) { + sorted := append([]time.Duration(nil), ds...) + sort.Slice(sorted, func(i, j int) bool { return sorted[i] < sorted[j] }) + p50 := sorted[len(sorted)/2] + p95 := sorted[(len(sorted)*95+99)/100-1] + verdict := "✅ 达标" + if p95 > target { + verdict = "❌ 超标" + } + fmt.Printf("%s n=%d P50=%dms P95=%dms min=%dms max=%dms 目标P95<%dms %s\n", + name, len(sorted), p50.Milliseconds(), p95.Milliseconds(), + sorted[0].Milliseconds(), sorted[len(sorted)-1].Milliseconds(), target.Milliseconds(), verdict) +} + +func truncate(s string, n int) string { + r := []rune(s) + if len(r) <= n { + return s + } + return string(r[:n]) + "…" +} + +// readWavData 提取 wav 的 data chunk(仅支持 PCM,同 gummycheck)。 +func readWavData(path string) ([]byte, error) { + b, err := os.ReadFile(path) + if err != nil { + return nil, err + } + if len(b) < 44 || string(b[0:4]) != "RIFF" || string(b[8:12]) != "WAVE" { + return nil, fmt.Errorf("不是 wav 文件") + } + off := 12 + for off+8 <= len(b) { + id := string(b[off : off+4]) + size := int(binary.LittleEndian.Uint32(b[off+4 : off+8])) + if id == "data" { + return b[off+8 : min(off+8+size, len(b))], nil + } + off += 8 + size + } + return nil, fmt.Errorf("未找到 data chunk") +} diff --git a/server/internal/asr/gummy.go b/server/internal/asr/gummy.go index 5e9fdfb..2fb7cdc 100644 --- a/server/internal/asr/gummy.go +++ b/server/internal/asr/gummy.go @@ -8,6 +8,7 @@ import ( "context" "encoding/json" "fmt" + "log/slog" "os" "sync" "time" @@ -34,10 +35,12 @@ func (g *GummyProvider) StartSession(ctx context.Context, cfg SessionConfig) (Se "Authorization": {"bearer " + g.APIKey}, "X-DashScope-DataInspection": {"enable"}, } + tDial := time.Now() conn, _, err := websocket.DefaultDialer.DialContext(ctx, dashscopeWS, header) if err != nil { return nil, fmt.Errorf("dashscope dial: %w", err) } + dialMs := time.Since(tDial).Milliseconds() taskID := uuid.NewString() runTask := map[string]any{ @@ -75,8 +78,13 @@ func (g *GummyProvider) StartSession(ctx context.Context, cfg SessionConfig) (Se go s.readLoop() // 协议要求:必须等 task-started 后才能推音频,否则服务端静默丢弃 + tWait := time.Now() select { case <-s.started: + // 会话建立耗时是 first_partial 延迟的组成部分(12B):dial 冷启(DNS+TLS) + // 可达 ~600ms、暖路径 ~60ms,task-started 通常 ~60ms。持续观测供延迟排查。 + slog.Info("gummy: session ready", "task", taskID, + "dial_ms", dialMs, "task_started_wait_ms", time.Since(tWait).Milliseconds()) return s, nil case <-s.done: return nil, fmt.Errorf("dashscope: closed before task-started") diff --git a/server/internal/gateway/gateway.go b/server/internal/gateway/gateway.go index a56c540..3979941 100644 --- a/server/internal/gateway/gateway.go +++ b/server/internal/gateway/gateway.go @@ -137,6 +137,12 @@ func (h *Handler) HandleWS(c *gin.Context) { // startSession 执行 start 预检并建立 Provider 会话;失败时返回 nil(已下发 error 帧)。 func (h *Handler) startSession(ctx context.Context, ws *wsConn, uid, deviceID string, msg protocol.ClientMsg) *session { now := time.Now() + // start 处理耗时(预检 + provider 握手)串行阻塞在读循环里,直接叠加进客户端 + // 可感知的 first_partial 延迟(12B),持续观测供延迟排查。 + defer func() { + slog.Info("gateway: start handled", "session", msg.SessionID, + "total_ms", time.Since(now).Milliseconds()) + }() // ① 设备 30min 窗口:次数 + 累计时长 if ok, err := store.AllowSession(ctx, h.RDB, deviceID, now); err != nil || !ok { diff --git a/server/internal/store/redis.go b/server/internal/store/redis.go index cbb6d37..b362bb4 100644 --- a/server/internal/store/redis.go +++ b/server/internal/store/redis.go @@ -46,10 +46,14 @@ else end end if sum + val > limit then return 0 end -local seq = redis.call('INCR', key .. ':seq') -redis.call('EXPIRE', key .. ':seq', window + 60) -- 17C:seq 计数器与 ZSET 同寿命,避免按设备永久泄漏 -redis.call('ZADD', key, now, now .. '-' .. seq .. ':' .. val) -redis.call('EXPIRE', key, window + 60) +-- val=0 为"只查不记"(AudioWindowExhausted 的 start 预检):不落 0 值成员, +-- 避免污染窗口 ZSET 并空耗 seq;count 模式 val 恒为 1,不受影响。 +if val > 0 then + local seq = redis.call('INCR', key .. ':seq') + redis.call('EXPIRE', key .. ':seq', window + 60) -- 17C:seq 计数器与 ZSET 同寿命,避免按设备永久泄漏 + redis.call('ZADD', key, now, now .. '-' .. seq .. ':' .. val) + redis.call('EXPIRE', key, window + 60) +end return 1 `)