diff --git a/go.mod b/go.mod index 63fbbba..7143ea7 100644 --- a/go.mod +++ b/go.mod @@ -3,6 +3,7 @@ module github.com/mathismqn/godeez go 1.25.0 require ( + github.com/mewkiz/flac v1.0.13 github.com/spf13/cobra v1.10.2 github.com/zalando/go-keyring v0.2.8 golang.org/x/mod v0.38.0 @@ -14,8 +15,11 @@ require ( github.com/danieljoos/wincred v1.2.3 // indirect github.com/fatih/color v1.19.0 // indirect github.com/godbus/dbus/v5 v5.2.2 // indirect + github.com/icza/bitio v1.1.0 // indirect github.com/mattn/go-colorable v0.1.15 // indirect github.com/mattn/go-isatty v0.0.24 // indirect + github.com/mewkiz/pkg v0.0.0-20250417130911-3f050ff8c56d // indirect + github.com/mewpkg/term v0.0.0-20241026122259-37a80af23985 // indirect golang.org/x/net v0.57.0 // indirect golang.org/x/sys v0.47.0 // indirect golang.org/x/text v0.40.0 // indirect diff --git a/go.sum b/go.sum index e5a3c5e..97842b3 100644 --- a/go.sum +++ b/go.sum @@ -23,12 +23,22 @@ github.com/go-flac/go-flac/v2 v2.0.4 h1:atf/kFa8U9idtkA//NO22XGr+MzQLeXZecnmP9sY github.com/go-flac/go-flac/v2 v2.0.4/go.mod h1:sYOlTKxutMW0RDYF+KlD6Zn+VOCZlIFQG/r/usPveCs= github.com/godbus/dbus/v5 v5.2.2 h1:TUR3TgtSVDmjiXOgAAyaZbYmIeP3DPkld3jgKGV8mXQ= github.com/godbus/dbus/v5 v5.2.2/go.mod h1:3AAv2+hPq5rdnr5txxxRwiGjPXamgoIHgz9FPBfOp3c= +github.com/icza/bitio v1.1.0 h1:ysX4vtldjdi3Ygai5m1cWy4oLkhWTAi+SyO6HC8L9T0= +github.com/icza/bitio v1.1.0/go.mod h1:0jGnlLAx8MKMr9VGnn/4YrvZiprkvBelsVIbA9Jjr9A= +github.com/icza/mighty v0.0.0-20180919140131-cfd07d671de6 h1:8UsGZ2rr2ksmEru6lToqnXgA8Mz1DP11X4zSJ159C3k= +github.com/icza/mighty v0.0.0-20180919140131-cfd07d671de6/go.mod h1:xQig96I1VNBDIWGCdTt54nHt6EeI639SmHycLYL7FkA= github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw= github.com/mattn/go-colorable v0.1.15 h1:+u9SLTRGnXv73cEsnsmoZBom+dMU88B2M0aDcWy0/jY= github.com/mattn/go-colorable v0.1.15/go.mod h1:6LmQG8QLFO4G5z1gPvYEzlUgJ2wF+stgPZH1UqBm1s8= github.com/mattn/go-isatty v0.0.24 h1:tGZZoVgT/KiqK1c8ocVLeDS8BSWMRd47J3Lbz7vsReI= github.com/mattn/go-isatty v0.0.24/go.mod h1:nMCL3Zebbrt45jsMDgnfIwz6ydEQApk5oEI3HqDio6A= +github.com/mewkiz/flac v1.0.13 h1:6wF8rRQKBFW159Daqx6Ro7K5ZnlVhHUKfS5aTsC4oXs= +github.com/mewkiz/flac v1.0.13/go.mod h1:HfPYDA+oxjyuqMu2V+cyKcxF51KM6incpw5eZXmfA6k= +github.com/mewkiz/pkg v0.0.0-20250417130911-3f050ff8c56d h1:IL2tii4jXLdhCeQN69HNzYYW1kl0meSG0wt5+sLwszU= +github.com/mewkiz/pkg v0.0.0-20250417130911-3f050ff8c56d/go.mod h1:SIpumAnUWSy0q9RzKD3pyH3g1t5vdawUAPcW5tQrUtI= +github.com/mewpkg/term v0.0.0-20241026122259-37a80af23985 h1:h8O1byDZ1uk6RUXMhj1QJU3VXFKXHDZxr4TXRPGeBa8= +github.com/mewpkg/term v0.0.0-20241026122259-37a80af23985/go.mod h1:uiPmbdUbdt1NkGApKl7htQjZ8S7XaGUAVulJUJ9v6q4= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM= diff --git a/internal/audio/wav.go b/internal/audio/wav.go new file mode 100644 index 0000000..4778568 --- /dev/null +++ b/internal/audio/wav.go @@ -0,0 +1,189 @@ +package audio + +import ( + "bufio" + "context" + "encoding/binary" + "errors" + "fmt" + "io" + "math" + "os" + "path/filepath" + + "github.com/mathismqn/godeez/internal/fsutil" + "github.com/mewkiz/flac" +) + +const ( + headerSize = 44 + formatPCM = 1 + ctxCheckInterval = 64 + maxDataSize = math.MaxUint32 - (headerSize - 8) +) + +func FLACToWAV(ctx context.Context, srcPath, dstPath string) error { + stream, err := flac.Open(srcPath) + if err != nil { + return err + } + defer stream.Close() + + info := stream.Info + bytesPerSample, err := bytesPerSample(info.BitsPerSample) + if err != nil { + return err + } + if info.NChannels < 1 || info.NChannels > 2 { + return fmt.Errorf("unsupported channel count: %d", info.NChannels) + } + if size := int64(info.NSamples) * int64(info.NChannels) * int64(bytesPerSample); size > maxDataSize { + return fmt.Errorf("audio data of %d bytes exceeds the wav format limit", size) + } + + file, err := os.CreateTemp(filepath.Dir(dstPath), fsutil.PartPattern) + if err != nil { + return err + } + tmpPath := file.Name() + done := false + defer func() { + if !done { + file.Close() + os.Remove(tmpPath) + } + }() + + w := bufio.NewWriter(file) + if err := writeHeader(w, info.SampleRate, info.NChannels, info.BitsPerSample, 0); err != nil { + return err + } + + dataSize, err := writeSamples(ctx, w, stream, int(info.NChannels), bytesPerSample) + if err != nil { + return err + } + if dataSize > maxDataSize { + return fmt.Errorf("audio data of %d bytes exceeds the wav format limit", dataSize) + } + if dataSize%2 != 0 { + if err := w.WriteByte(0); err != nil { + return err + } + } + if err := w.Flush(); err != nil { + return err + } + if err := patchSizes(file, dataSize); err != nil { + return err + } + + if err := file.Sync(); err != nil { + return err + } + if err := file.Close(); err != nil { + return err + } + if err := os.Rename(tmpPath, dstPath); err != nil { + return err + } + done = true + + return nil +} + +func bytesPerSample(bitsPerSample uint8) (int, error) { + switch bitsPerSample { + case 8, 16, 24: + return int(bitsPerSample) / 8, nil + default: + return 0, fmt.Errorf("unsupported bit depth: %d", bitsPerSample) + } +} + +func writeHeader(w io.Writer, sampleRate uint32, nChannels, bitsPerSample uint8, dataSize uint32) error { + blockAlign := uint32(nChannels) * uint32(bitsPerSample) / 8 + + header := make([]byte, 0, headerSize) + header = append(header, "RIFF"...) + header = binary.LittleEndian.AppendUint32(header, uint32(headerSize-8)+dataSize) + header = append(header, "WAVE"...) + header = append(header, "fmt "...) + header = binary.LittleEndian.AppendUint32(header, 16) + header = binary.LittleEndian.AppendUint16(header, formatPCM) + header = binary.LittleEndian.AppendUint16(header, uint16(nChannels)) + header = binary.LittleEndian.AppendUint32(header, sampleRate) + header = binary.LittleEndian.AppendUint32(header, sampleRate*blockAlign) + header = binary.LittleEndian.AppendUint16(header, uint16(blockAlign)) + header = binary.LittleEndian.AppendUint16(header, uint16(bitsPerSample)) + header = append(header, "data"...) + header = binary.LittleEndian.AppendUint32(header, dataSize) + + _, err := w.Write(header) + + return err +} + +func writeSamples(ctx context.Context, w io.Writer, stream *flac.Stream, nChannels, bytesPerSample int) (int64, error) { + var dataSize int64 + buf := make([]byte, 4) + + for i := 0; ; i++ { + if i%ctxCheckInterval == 0 { + select { + case <-ctx.Done(): + return dataSize, ctx.Err() + default: + } + } + + frame, err := stream.ParseNext() + if err != nil { + if errors.Is(err, io.EOF) { + break + } + return dataSize, err + } + if len(frame.Subframes) != nChannels { + return dataSize, fmt.Errorf("frame %d has %d channels, want %d", frame.Num, len(frame.Subframes), nChannels) + } + + for i := range frame.Subframes[0].Samples { + for _, subframe := range frame.Subframes { + putSample(buf, subframe.Samples[i], bytesPerSample) + if _, err := w.Write(buf[:bytesPerSample]); err != nil { + return dataSize, err + } + dataSize += int64(bytesPerSample) + } + } + } + + return dataSize, nil +} + +func putSample(buf []byte, sample int32, bytesPerSample int) { + if bytesPerSample == 1 { + buf[0] = byte(sample + 128) + return + } + + value := uint32(sample) + for i := range bytesPerSample { + buf[i] = byte(value >> (8 * i)) + } +} + +func patchSizes(file *os.File, dataSize int64) error { + buf := make([]byte, 4) + + binary.LittleEndian.PutUint32(buf, uint32(headerSize-8+dataSize+dataSize%2)) + if _, err := file.WriteAt(buf, 4); err != nil { + return err + } + + binary.LittleEndian.PutUint32(buf, uint32(dataSize)) + _, err := file.WriteAt(buf, headerSize-4) + + return err +} diff --git a/internal/audio/wav_test.go b/internal/audio/wav_test.go new file mode 100644 index 0000000..24643fe --- /dev/null +++ b/internal/audio/wav_test.go @@ -0,0 +1,291 @@ +package audio + +import ( + "bytes" + "context" + "encoding/binary" + "os" + "path/filepath" + "testing" + + "github.com/mewkiz/flac" + "github.com/mewkiz/flac/frame" + "github.com/mewkiz/flac/meta" +) + +func testSamples(nChannels, bitsPerSample, nSamples int) [][]int32 { + max := int32(1)<<(bitsPerSample-1) - 1 + min := -int32(1) << (bitsPerSample - 1) + + channels := make([][]int32, nChannels) + for c := range channels { + samples := make([]int32, nSamples) + for i := range samples { + switch i { + case 0: + samples[i] = min + case 1: + samples[i] = max + case 2: + samples[i] = 0 + default: + samples[i] = int32(i*(c+1)) % max + if i%3 == 0 { + samples[i] = -samples[i] + } + } + } + channels[c] = samples + } + + return channels +} + +func writeTestFLAC(t *testing.T, path string, sampleRate uint32, bitsPerSample uint8, channels [][]int32) { + t.Helper() + + nSamples := len(channels[0]) + info := &meta.StreamInfo{ + BlockSizeMin: uint16(nSamples), + BlockSizeMax: uint16(nSamples), + SampleRate: sampleRate, + NChannels: uint8(len(channels)), + BitsPerSample: bitsPerSample, + NSamples: uint64(nSamples), + } + + file, err := os.Create(path) + if err != nil { + t.Fatalf("create flac: %v", err) + } + defer file.Close() + + enc, err := flac.NewEncoder(file, info) + if err != nil { + t.Fatalf("new encoder: %v", err) + } + + subframes := make([]*frame.Subframe, len(channels)) + for c, samples := range channels { + subframes[c] = &frame.Subframe{ + SubHeader: frame.SubHeader{Pred: frame.PredVerbatim}, + Samples: samples, + NSamples: nSamples, + } + } + + channelsLayout := frame.ChannelsMono + if len(channels) == 2 { + channelsLayout = frame.ChannelsLR + } + + f := &frame.Frame{ + Header: frame.Header{ + HasFixedBlockSize: true, + BlockSize: uint16(nSamples), + SampleRate: sampleRate, + Channels: channelsLayout, + BitsPerSample: bitsPerSample, + }, + Subframes: subframes, + } + + if err := enc.WriteFrame(f); err != nil { + t.Fatalf("write frame: %v", err) + } + if err := enc.Close(); err != nil { + t.Fatalf("close encoder: %v", err) + } +} + +func expectedPCM(channels [][]int32, bytesPerSample int) []byte { + var buf bytes.Buffer + + for i := range channels[0] { + for _, samples := range channels { + sample := samples[i] + switch bytesPerSample { + case 1: + buf.WriteByte(byte(sample + 128)) + case 2: + buf.Write([]byte{byte(sample), byte(sample >> 8)}) + case 3: + buf.Write([]byte{byte(sample), byte(sample >> 8), byte(sample >> 16)}) + } + } + } + + return buf.Bytes() +} + +func TestFLACToWAV(t *testing.T) { + tests := []struct { + name string + sampleRate uint32 + bitsPerSample uint8 + nChannels int + nSamples int + }{ + {"16 bit stereo", 44100, 16, 2, 512}, + {"16 bit mono", 44100, 16, 1, 512}, + {"24 bit stereo", 48000, 24, 2, 333}, + {"8 bit mono", 22050, 8, 1, 128}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + dir := t.TempDir() + src := filepath.Join(dir, "in.flac") + dst := filepath.Join(dir, "out.wav") + + channels := testSamples(tt.nChannels, int(tt.bitsPerSample), tt.nSamples) + writeTestFLAC(t, src, tt.sampleRate, tt.bitsPerSample, channels) + + if err := FLACToWAV(context.Background(), src, dst); err != nil { + t.Fatalf("FLACToWAV() error = %v", err) + } + + got, err := os.ReadFile(dst) + if err != nil { + t.Fatalf("read wav: %v", err) + } + + bytesPerSample := int(tt.bitsPerSample) / 8 + blockAlign := uint16(tt.nChannels * bytesPerSample) + want := expectedPCM(channels, bytesPerSample) + dataSize := uint32(len(want)) + pad := dataSize % 2 + + if len(got) != headerSize+len(want)+int(pad) { + t.Fatalf("file size = %d, want %d", len(got), headerSize+len(want)+int(pad)) + } + if string(got[0:4]) != "RIFF" || string(got[8:12]) != "WAVE" { + t.Errorf("magic = %q %q, want \"RIFF\" \"WAVE\"", got[0:4], got[8:12]) + } + if size := binary.LittleEndian.Uint32(got[4:8]); size != headerSize-8+dataSize+pad { + t.Errorf("riff size = %d, want %d", size, headerSize-8+dataSize+pad) + } + if string(got[12:16]) != "fmt " { + t.Errorf("fmt chunk id = %q, want \"fmt \"", got[12:16]) + } + if size := binary.LittleEndian.Uint32(got[16:20]); size != 16 { + t.Errorf("fmt chunk size = %d, want 16", size) + } + if format := binary.LittleEndian.Uint16(got[20:22]); format != formatPCM { + t.Errorf("format = %d, want %d", format, formatPCM) + } + if n := binary.LittleEndian.Uint16(got[22:24]); n != uint16(tt.nChannels) { + t.Errorf("channels = %d, want %d", n, tt.nChannels) + } + if rate := binary.LittleEndian.Uint32(got[24:28]); rate != tt.sampleRate { + t.Errorf("sample rate = %d, want %d", rate, tt.sampleRate) + } + if rate := binary.LittleEndian.Uint32(got[28:32]); rate != tt.sampleRate*uint32(blockAlign) { + t.Errorf("byte rate = %d, want %d", rate, tt.sampleRate*uint32(blockAlign)) + } + if align := binary.LittleEndian.Uint16(got[32:34]); align != blockAlign { + t.Errorf("block align = %d, want %d", align, blockAlign) + } + if bits := binary.LittleEndian.Uint16(got[34:36]); bits != uint16(tt.bitsPerSample) { + t.Errorf("bits per sample = %d, want %d", bits, tt.bitsPerSample) + } + if string(got[36:40]) != "data" { + t.Errorf("data chunk id = %q, want \"data\"", got[36:40]) + } + if size := binary.LittleEndian.Uint32(got[40:44]); size != dataSize { + t.Errorf("data chunk size = %d, want %d", size, dataSize) + } + if !bytes.Equal(got[headerSize:headerSize+len(want)], want) { + t.Error("pcm payload does not match the source samples") + } + }) + } +} + +func TestFLACToWAVErrors(t *testing.T) { + dir := t.TempDir() + + notFLAC := filepath.Join(dir, "not.flac") + if err := os.WriteFile(notFLAC, []byte("this is not a flac file"), 0644); err != nil { + t.Fatalf("write file: %v", err) + } + + valid := filepath.Join(dir, "valid.flac") + writeTestFLAC(t, valid, 44100, 16, testSamples(2, 16, 64)) + data, err := os.ReadFile(valid) + if err != nil { + t.Fatalf("read flac: %v", err) + } + truncated := filepath.Join(dir, "truncated.flac") + if err := os.WriteFile(truncated, data[:len(data)/2], 0644); err != nil { + t.Fatalf("write file: %v", err) + } + + tests := []struct { + name string + path string + }{ + {"not a flac file", notFLAC}, + {"missing file", filepath.Join(dir, "missing.flac")}, + {"truncated file", truncated}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + dst := filepath.Join(t.TempDir(), "out.wav") + if err := FLACToWAV(context.Background(), tt.path, dst); err == nil { + t.Error("FLACToWAV() error = nil, want error") + } + if _, err := os.Stat(dst); err == nil { + t.Error("FLACToWAV() left an output file behind") + } + }) + } +} + +func TestFLACToWAVCanceled(t *testing.T) { + dir := t.TempDir() + src := filepath.Join(dir, "in.flac") + dst := filepath.Join(dir, "out.wav") + writeTestFLAC(t, src, 44100, 16, testSamples(2, 16, 512)) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + if err := FLACToWAV(ctx, src, dst); err == nil { + t.Error("FLACToWAV() error = nil, want context.Canceled") + } + if _, err := os.Stat(dst); err == nil { + t.Error("FLACToWAV() left an output file behind") + } + matches, _ := filepath.Glob(filepath.Join(dir, ".godeez-*.part")) + if len(matches) > 0 { + t.Errorf("FLACToWAV() left %d part files behind", len(matches)) + } +} + +func TestBytesPerSample(t *testing.T) { + tests := []struct { + bitsPerSample uint8 + want int + wantErr bool + }{ + {8, 1, false}, + {16, 2, false}, + {24, 3, false}, + {4, 0, true}, + {12, 0, true}, + {20, 0, true}, + {32, 0, true}, + } + + for _, tt := range tests { + got, err := bytesPerSample(tt.bitsPerSample) + if (err != nil) != tt.wantErr { + t.Errorf("bytesPerSample(%d) error = %v, wantErr %v", tt.bitsPerSample, err, tt.wantErr) + } + if got != tt.want { + t.Errorf("bytesPerSample(%d) = %d, want %d", tt.bitsPerSample, got, tt.want) + } + } +} diff --git a/internal/fsutil/fsutil.go b/internal/fsutil/fsutil.go index 3fc38c5..9d992c6 100644 --- a/internal/fsutil/fsutil.go +++ b/internal/fsutil/fsutil.go @@ -5,6 +5,8 @@ import ( "os" ) +const PartPattern = ".godeez-*.part" + func EnsureDir(path string) error { info, err := os.Stat(path) if os.IsNotExist(err) {