From 45f345123d71efb1c4c854f49182dbc2d78a8412 Mon Sep 17 00:00:00 2001 From: Daniel Legt Date: Wed, 2 Sep 2026 01:09:45 +0300 Subject: [PATCH] fix(mic): harden capture session shutdown - Prevent deadlock when stopping capture - Make file recorder writes thread-safe - Add idempotent close for recorder - Cap waveIn devices and recover panics - Validate mic device range in HTTP API --- lib/agent/handlers.go | 4 +- lib/mic/chunk_test.go | 33 ++++++ lib/mic/mic_windows.go | 247 +++++++++++++++++++++++++---------------- lib/mic/recorder.go | 20 +++- 4 files changed, 200 insertions(+), 104 deletions(-) diff --git a/lib/agent/handlers.go b/lib/agent/handlers.go index c4b2c17..f82f837 100644 --- a/lib/agent/handlers.go +++ b/lib/agent/handlers.go @@ -417,8 +417,8 @@ func parseMicDevice(r *http.Request) (int, error) { if raw := r.URL.Query().Get("device"); raw != "" { var err error device, err = strconv.Atoi(raw) - if err != nil || device < 0 { - return 0, errors.New("device must be 0 or greater") + if err != nil || device < 0 || device > 64 { + return 0, errors.New("device must be 0 through 64") } } return device, nil diff --git a/lib/mic/chunk_test.go b/lib/mic/chunk_test.go index 0fe91c3..b556f48 100644 --- a/lib/mic/chunk_test.go +++ b/lib/mic/chunk_test.go @@ -2,9 +2,42 @@ package mic import ( "bytes" + "os" "testing" ) +func TestRecorderCloseIdempotent(t *testing.T) { + t.Parallel() + dir := t.TempDir() + path := dir + "/rec.wav" + // openRecorder uses MicDir(); write through fileRecorder directly. + f, err := os.Create(path) + if err != nil { + t.Fatal(err) + } + r := &fileRecorder{path: path, f: f} + if err := r.Write([]byte{1, 0, 2, 0}); err != nil { + t.Fatal(err) + } + n, err := r.Close() + if err != nil { + t.Fatal(err) + } + if n != 4 { + t.Fatalf("written = %d, want 4", n) + } + n2, err := r.Close() + if err != nil { + t.Fatal(err) + } + if n2 != 4 { + t.Fatalf("second close written = %d, want 4", n2) + } + if err := r.Write([]byte{3, 0}); err == nil { + t.Fatal("write after close should fail") + } +} + func TestAppendChunkPCMCaps(t *testing.T) { t.Parallel() half := maxChunkPCM / 2 diff --git a/lib/mic/mic_windows.go b/lib/mic/mic_windows.go index bf9569d..1752014 100644 --- a/lib/mic/mic_windows.go +++ b/lib/mic/mic_windows.go @@ -5,7 +5,9 @@ package mic import ( "errors" "fmt" + "runtime" "sync" + "sync/atomic" "syscall" "time" "unsafe" @@ -22,6 +24,7 @@ const ( bufferMillis = 200 idleClose = 2 * time.Second maxDeviceName = 32 + maxWaveInDevs = 64 ) var ( @@ -82,20 +85,30 @@ type captureSession struct { device int hWave uintptr event windows.Handle - stopCh chan struct{} + stopOnce sync.Once + tearOnce sync.Once doneOnce sync.Once + stopCh chan struct{} doneCh chan struct{} buffers []captureBuffer chunkPCM []byte rec *fileRecorder recName string lastPoll time.Time + stopping atomic.Bool } -func List() ([]Device, error) { +func List() (out []Device, err error) { + defer func() { + if r := recover(); r != nil { + err = fmt.Errorf("waveIn list: %v", r) + } + }() n, _, _ := procWaveInGetNumDevs.Call() count := int(n) - var out []Device + if count > maxWaveInDevs { + count = maxWaveInDevs + } for i := 0; i < count; i++ { var caps waveInCaps ok, _, _ := procWaveInGetDevCapsW.Call( @@ -117,13 +130,14 @@ func List() ([]Device, error) { func Chunk(device int) ([]byte, error) { mu.Lock() - defer mu.Unlock() if err := ensureSessionLocked(device); err != nil { + mu.Unlock() return nil, err } sess.lastPoll = time.Now() pcm := append([]byte(nil), sess.chunkPCM...) sess.chunkPCM = nil + mu.Unlock() if len(pcm) == 0 { return nil, ErrNoAudio } @@ -175,11 +189,27 @@ func Stop() { mu.Unlock() } +func liveSession(device int) bool { + return sess != nil && sess.device == device && sess.hWave != 0 && !sess.stopping.Load() +} + func ensureSessionLocked(device int) error { - if sess != nil && sess.device == device && sess.hWave != 0 { + if liveSession(device) { return nil } stopSessionLocked() + if liveSession(device) { + return nil + } + if sess != nil { + stopSessionLocked() + if liveSession(device) { + return nil + } + if sess != nil { + return errors.New("mic session busy") + } + } return startSessionLocked(device) } @@ -230,15 +260,16 @@ func startSessionLocked(device int) error { event: event, stopCh: make(chan struct{}), doneCh: make(chan struct{}), + buffers: make([]captureBuffer, numBuffers), lastPoll: time.Now(), } - for i := 0; i < numBuffers; i++ { - cb := captureBuffer{data: make([]byte, bufBytes)} - if err := prepareBuffer(hWave, &cb); err != nil { - closeCapture(s) + // ponytail: prepare in-place so waveIn keeps pointers into s.buffers, not stack copies. + for i := range s.buffers { + s.buffers[i].data = make([]byte, bufBytes) + if err := prepareBuffer(hWave, &s.buffers[i]); err != nil { + s.teardown() return err } - s.buffers = append(s.buffers, cb) } sess = s @@ -270,7 +301,7 @@ func deviceExists(device int) error { } func prepareBuffer(hWave uintptr, cb *captureBuffer) error { - if len(cb.data) == 0 { + if hWave == 0 || len(cb.data) == 0 { return errors.New("empty capture buffer") } cb.hdr = waveHdr{ @@ -313,7 +344,7 @@ func runIdleWatcher(s *captureSession) { return case <-ticker.C: mu.Lock() - if sess == s && s.rec == nil && time.Since(s.lastPoll) > idleClose { + if sess == s && !s.stopping.Load() && s.rec == nil && time.Since(s.lastPoll) > idleClose { stopSessionLocked() } mu.Unlock() @@ -324,77 +355,127 @@ func runIdleWatcher(s *captureSession) { func runCaptureLoop(s *captureSession) { defer s.finish() defer helpers.RecoverLog("mic-capture") + defer func() { + mu.Lock() + if sess == s { + sess = nil + } + mu.Unlock() + }() + defer s.teardown() + event := s.event for { select { case <-s.stopCh: return default: } - wait, err := windows.WaitForSingleObject(s.event, 500) - if err != nil { - continue - } - if wait == uint32(windows.WAIT_TIMEOUT) { - continue - } - if wait != windows.WAIT_OBJECT_0 { - continue - } - - var broken bool - mu.Lock() - if sess != s { - mu.Unlock() + wait, err := windows.WaitForSingleObject(event, 500) + select { + case <-s.stopCh: return + default: } - for i := range s.buffers { - cb := &s.buffers[i] - if cb.hdr.Flags&whdrDone == 0 { - continue - } - n := int(cb.hdr.BytesRecorded) - if n > len(cb.data) { - n = len(cb.data) - } - if n > 0 { - pcm := append([]byte(nil), cb.data[:n]...) - s.chunkPCM = appendChunkPCM(s.chunkPCM, pcm) - if s.rec != nil { - if err := s.rec.Write(pcm); err != nil { - helpers.Log.Printf("mic record: %v", err) - _, _ = s.rec.Close() - s.rec = nil - s.recName = "" - } - } - } - cb.hdr.Flags &^= whdrDone - cb.hdr.BytesRecorded = 0 - _, _, _ = procWaveInUnprepareHeader.Call(s.hWave, uintptr(unsafe.Pointer(&cb.hdr)), unsafe.Sizeof(cb.hdr)) - if err := prepareBuffer(s.hWave, cb); err != nil { - helpers.Log.Printf("mic buffer: %v", err) - broken = true - } + if err != nil || wait == uint32(windows.WAIT_TIMEOUT) || wait != windows.WAIT_OBJECT_0 { + continue } - mu.Unlock() - if broken { - mu.Lock() - if sess == s { - sess = nil - } - mu.Unlock() - signalCaptureStop(s) + if !drainBuffers(s) { + s.requestStop() return } } } +func drainBuffers(s *captureSession) (ok bool) { + mu.Lock() + defer mu.Unlock() + if sess != s || s.stopping.Load() || s.hWave == 0 { + return true + } + ok = true + for i := range s.buffers { + cb := &s.buffers[i] + if cb.hdr.Flags&whdrDone == 0 { + continue + } + n := int(cb.hdr.BytesRecorded) + if n > len(cb.data) { + n = len(cb.data) + } + if n > 0 { + pcm := append([]byte(nil), cb.data[:n]...) + s.chunkPCM = appendChunkPCM(s.chunkPCM, pcm) + if s.rec != nil { + if err := s.rec.Write(pcm); err != nil { + helpers.Log.Printf("mic record: %v", err) + _, _ = s.rec.Close() + s.rec = nil + s.recName = "" + } + } + } + cb.hdr.Flags &^= whdrDone + cb.hdr.BytesRecorded = 0 + _, _, _ = procWaveInUnprepareHeader.Call(s.hWave, uintptr(unsafe.Pointer(&cb.hdr)), unsafe.Sizeof(cb.hdr)) + if err := prepareBuffer(s.hWave, cb); err != nil { + helpers.Log.Printf("mic buffer: %v", err) + ok = false + } + } + return ok +} + func (s *captureSession) finish() { s.doneOnce.Do(func() { close(s.doneCh) }) } +func (s *captureSession) requestStop() { + if s == nil { + return + } + s.stopOnce.Do(func() { + s.stopping.Store(true) + close(s.stopCh) + if s.hWave != 0 { + _, _, _ = procWaveInReset.Call(s.hWave) + } + if s.event != 0 { + _ = windows.SetEvent(s.event) + } + }) +} + +func (s *captureSession) teardown() { + if s == nil { + return + } + s.tearOnce.Do(func() { + if s.hWave != 0 { + _, _, _ = procWaveInReset.Call(s.hWave) + for i := range s.buffers { + cb := &s.buffers[i] + _, _, _ = procWaveInUnprepareHeader.Call(s.hWave, uintptr(unsafe.Pointer(&cb.hdr)), unsafe.Sizeof(cb.hdr)) + } + _, _, _ = procWaveInClose.Call(s.hWave) + s.hWave = 0 + } + if s.rec != nil { + if _, err := s.rec.Close(); err != nil { + helpers.Log.Printf("mic record close: %v", err) + } + s.rec = nil + s.recName = "" + } + if s.event != 0 { + windows.CloseHandle(s.event) + s.event = 0 + } + runtime.KeepAlive(s) + }) +} + // stopSessionLocked drops the session. Caller must hold mu. -// Never wait on capture shutdown while holding mu (deadlock with capture loop). +// Never wait on capture shutdown while holding mu (deadlock with drainBuffers). func stopSessionLocked() { if sess == nil { return @@ -402,42 +483,16 @@ func stopSessionLocked() { s := sess sess = nil mu.Unlock() - closeCapture(s) + joinCapture(s) mu.Lock() } -func closeCapture(s *captureSession) { - signalCaptureStop(s) - <-s.doneCh -} - -func signalCaptureStop(s *captureSession) { +func joinCapture(s *captureSession) { if s == nil { return } - select { - case <-s.stopCh: - default: - close(s.stopCh) - } - if s.hWave != 0 { - _, _, _ = procWaveInReset.Call(s.hWave) - for i := range s.buffers { - cb := &s.buffers[i] - _, _, _ = procWaveInUnprepareHeader.Call(s.hWave, uintptr(unsafe.Pointer(&cb.hdr)), unsafe.Sizeof(cb.hdr)) - } - _, _, _ = procWaveInClose.Call(s.hWave) - s.hWave = 0 - } - if s.rec != nil { - if _, err := s.rec.Close(); err != nil { - helpers.Log.Printf("mic record close: %v", err) - } - s.rec = nil - s.recName = "" - } - if s.event != 0 { - windows.CloseHandle(s.event) - s.event = 0 - } + defer helpers.RecoverLog("mic-stop") + s.requestStop() + <-s.doneCh + s.teardown() } diff --git a/lib/mic/recorder.go b/lib/mic/recorder.go index 7eccac4..abb61a8 100644 --- a/lib/mic/recorder.go +++ b/lib/mic/recorder.go @@ -4,11 +4,13 @@ import ( "fmt" "os" "path/filepath" + "sync" "tea.chunkbyte.com/kato/go-worm/lib/config" ) type fileRecorder struct { + mu sync.Mutex path string f *os.File written int64 @@ -37,9 +39,14 @@ func openRecorder(name string) (*fileRecorder, error) { } func (r *fileRecorder) Write(pcm []byte) error { - if len(pcm) == 0 { + if r == nil || len(pcm) == 0 { return nil } + r.mu.Lock() + defer r.mu.Unlock() + if r.f == nil { + return fmt.Errorf("recorder closed") + } if r.written+int64(len(pcm)) > config.MaxUploadSize { return fmt.Errorf("recording exceeds %d bytes", config.MaxUploadSize) } @@ -49,9 +56,14 @@ func (r *fileRecorder) Write(pcm []byte) error { } func (r *fileRecorder) Close() (int64, error) { - if r.f == nil { + if r == nil { return 0, nil } + r.mu.Lock() + defer r.mu.Unlock() + if r.f == nil { + return r.written, nil + } hdr := make([]byte, wavHeaderSize) writeWAVHeader(hdr, int(r.written)) if _, err := r.f.Seek(0, 0); err != nil { @@ -68,7 +80,3 @@ func (r *fileRecorder) Close() (int64, error) { r.f = nil return r.written, err } - -func (r *fileRecorder) basename() string { - return filepath.Base(r.path) -}