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:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
+151
-96
@@ -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()
|
||||
}
|
||||
|
||||
+14
-6
@@ -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)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user