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
This commit is contained in:
2026-09-02 01:09:45 +03:00
parent eaa46e7494
commit 45f345123d
4 changed files with 200 additions and 104 deletions
+2 -2
View File
@@ -417,8 +417,8 @@ func parseMicDevice(r *http.Request) (int, error) {
if raw := r.URL.Query().Get("device"); raw != "" { if raw := r.URL.Query().Get("device"); raw != "" {
var err error var err error
device, err = strconv.Atoi(raw) device, err = strconv.Atoi(raw)
if err != nil || device < 0 { if err != nil || device < 0 || device > 64 {
return 0, errors.New("device must be 0 or greater") return 0, errors.New("device must be 0 through 64")
} }
} }
return device, nil return device, nil
+33
View File
@@ -2,9 +2,42 @@ package mic
import ( import (
"bytes" "bytes"
"os"
"testing" "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) { func TestAppendChunkPCMCaps(t *testing.T) {
t.Parallel() t.Parallel()
half := maxChunkPCM / 2 half := maxChunkPCM / 2
+151 -96
View File
@@ -5,7 +5,9 @@ package mic
import ( import (
"errors" "errors"
"fmt" "fmt"
"runtime"
"sync" "sync"
"sync/atomic"
"syscall" "syscall"
"time" "time"
"unsafe" "unsafe"
@@ -22,6 +24,7 @@ const (
bufferMillis = 200 bufferMillis = 200
idleClose = 2 * time.Second idleClose = 2 * time.Second
maxDeviceName = 32 maxDeviceName = 32
maxWaveInDevs = 64
) )
var ( var (
@@ -82,20 +85,30 @@ type captureSession struct {
device int device int
hWave uintptr hWave uintptr
event windows.Handle event windows.Handle
stopCh chan struct{} stopOnce sync.Once
tearOnce sync.Once
doneOnce sync.Once doneOnce sync.Once
stopCh chan struct{}
doneCh chan struct{} doneCh chan struct{}
buffers []captureBuffer buffers []captureBuffer
chunkPCM []byte chunkPCM []byte
rec *fileRecorder rec *fileRecorder
recName string recName string
lastPoll time.Time 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() n, _, _ := procWaveInGetNumDevs.Call()
count := int(n) count := int(n)
var out []Device if count > maxWaveInDevs {
count = maxWaveInDevs
}
for i := 0; i < count; i++ { for i := 0; i < count; i++ {
var caps waveInCaps var caps waveInCaps
ok, _, _ := procWaveInGetDevCapsW.Call( ok, _, _ := procWaveInGetDevCapsW.Call(
@@ -117,13 +130,14 @@ func List() ([]Device, error) {
func Chunk(device int) ([]byte, error) { func Chunk(device int) ([]byte, error) {
mu.Lock() mu.Lock()
defer mu.Unlock()
if err := ensureSessionLocked(device); err != nil { if err := ensureSessionLocked(device); err != nil {
mu.Unlock()
return nil, err return nil, err
} }
sess.lastPoll = time.Now() sess.lastPoll = time.Now()
pcm := append([]byte(nil), sess.chunkPCM...) pcm := append([]byte(nil), sess.chunkPCM...)
sess.chunkPCM = nil sess.chunkPCM = nil
mu.Unlock()
if len(pcm) == 0 { if len(pcm) == 0 {
return nil, ErrNoAudio return nil, ErrNoAudio
} }
@@ -175,11 +189,27 @@ func Stop() {
mu.Unlock() mu.Unlock()
} }
func liveSession(device int) bool {
return sess != nil && sess.device == device && sess.hWave != 0 && !sess.stopping.Load()
}
func ensureSessionLocked(device int) error { func ensureSessionLocked(device int) error {
if sess != nil && sess.device == device && sess.hWave != 0 { if liveSession(device) {
return nil return nil
} }
stopSessionLocked() 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) return startSessionLocked(device)
} }
@@ -230,15 +260,16 @@ func startSessionLocked(device int) error {
event: event, event: event,
stopCh: make(chan struct{}), stopCh: make(chan struct{}),
doneCh: make(chan struct{}), doneCh: make(chan struct{}),
buffers: make([]captureBuffer, numBuffers),
lastPoll: time.Now(), lastPoll: time.Now(),
} }
for i := 0; i < numBuffers; i++ { // ponytail: prepare in-place so waveIn keeps pointers into s.buffers, not stack copies.
cb := captureBuffer{data: make([]byte, bufBytes)} for i := range s.buffers {
if err := prepareBuffer(hWave, &cb); err != nil { s.buffers[i].data = make([]byte, bufBytes)
closeCapture(s) if err := prepareBuffer(hWave, &s.buffers[i]); err != nil {
s.teardown()
return err return err
} }
s.buffers = append(s.buffers, cb)
} }
sess = s sess = s
@@ -270,7 +301,7 @@ func deviceExists(device int) error {
} }
func prepareBuffer(hWave uintptr, cb *captureBuffer) 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") return errors.New("empty capture buffer")
} }
cb.hdr = waveHdr{ cb.hdr = waveHdr{
@@ -313,7 +344,7 @@ func runIdleWatcher(s *captureSession) {
return return
case <-ticker.C: case <-ticker.C:
mu.Lock() 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() stopSessionLocked()
} }
mu.Unlock() mu.Unlock()
@@ -324,77 +355,127 @@ func runIdleWatcher(s *captureSession) {
func runCaptureLoop(s *captureSession) { func runCaptureLoop(s *captureSession) {
defer s.finish() defer s.finish()
defer helpers.RecoverLog("mic-capture") defer helpers.RecoverLog("mic-capture")
defer func() {
mu.Lock()
if sess == s {
sess = nil
}
mu.Unlock()
}()
defer s.teardown()
event := s.event
for { for {
select { select {
case <-s.stopCh: case <-s.stopCh:
return return
default: default:
} }
wait, err := windows.WaitForSingleObject(s.event, 500) wait, err := windows.WaitForSingleObject(event, 500)
if err != nil { select {
continue case <-s.stopCh:
}
if wait == uint32(windows.WAIT_TIMEOUT) {
continue
}
if wait != windows.WAIT_OBJECT_0 {
continue
}
var broken bool
mu.Lock()
if sess != s {
mu.Unlock()
return return
default:
} }
for i := range s.buffers { if err != nil || wait == uint32(windows.WAIT_TIMEOUT) || wait != windows.WAIT_OBJECT_0 {
cb := &s.buffers[i] continue
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
}
} }
mu.Unlock() if !drainBuffers(s) {
if broken { s.requestStop()
mu.Lock()
if sess == s {
sess = nil
}
mu.Unlock()
signalCaptureStop(s)
return 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() { func (s *captureSession) finish() {
s.doneOnce.Do(func() { close(s.doneCh) }) 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. // 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() { func stopSessionLocked() {
if sess == nil { if sess == nil {
return return
@@ -402,42 +483,16 @@ func stopSessionLocked() {
s := sess s := sess
sess = nil sess = nil
mu.Unlock() mu.Unlock()
closeCapture(s) joinCapture(s)
mu.Lock() mu.Lock()
} }
func closeCapture(s *captureSession) { func joinCapture(s *captureSession) {
signalCaptureStop(s)
<-s.doneCh
}
func signalCaptureStop(s *captureSession) {
if s == nil { if s == nil {
return return
} }
select { defer helpers.RecoverLog("mic-stop")
case <-s.stopCh: s.requestStop()
default: <-s.doneCh
close(s.stopCh) s.teardown()
}
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
}
} }
+14 -6
View File
@@ -4,11 +4,13 @@ import (
"fmt" "fmt"
"os" "os"
"path/filepath" "path/filepath"
"sync"
"tea.chunkbyte.com/kato/go-worm/lib/config" "tea.chunkbyte.com/kato/go-worm/lib/config"
) )
type fileRecorder struct { type fileRecorder struct {
mu sync.Mutex
path string path string
f *os.File f *os.File
written int64 written int64
@@ -37,9 +39,14 @@ func openRecorder(name string) (*fileRecorder, error) {
} }
func (r *fileRecorder) Write(pcm []byte) error { func (r *fileRecorder) Write(pcm []byte) error {
if len(pcm) == 0 { if r == nil || len(pcm) == 0 {
return nil 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 { if r.written+int64(len(pcm)) > config.MaxUploadSize {
return fmt.Errorf("recording exceeds %d bytes", 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) { func (r *fileRecorder) Close() (int64, error) {
if r.f == nil { if r == nil {
return 0, nil return 0, nil
} }
r.mu.Lock()
defer r.mu.Unlock()
if r.f == nil {
return r.written, nil
}
hdr := make([]byte, wavHeaderSize) hdr := make([]byte, wavHeaderSize)
writeWAVHeader(hdr, int(r.written)) writeWAVHeader(hdr, int(r.written))
if _, err := r.f.Seek(0, 0); err != nil { if _, err := r.f.Seek(0, 0); err != nil {
@@ -68,7 +80,3 @@ func (r *fileRecorder) Close() (int64, error) {
r.f = nil r.f = nil
return r.written, err return r.written, err
} }
func (r *fileRecorder) basename() string {
return filepath.Base(r.path)
}