// 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") }