diff --git a/.github/workflows/docker.yml b/.github/workflows/docker.yml index 5aa0818..f3d53d8 100644 --- a/.github/workflows/docker.yml +++ b/.github/workflows/docker.yml @@ -60,6 +60,11 @@ jobs: if [[ "$GITHUB_REF_NAME" == "main" ]]; then version="$base_version" elif [[ "$GITHUB_REF_NAME" == "dev" ]]; then + base_version="$(tr -d '[:space:]' < VERSION)" + if [[ ! "$base_version" =~ ^[0-9]+\.[0-9]+\.[0-9]+$ ]]; then + echo "Invalid development version: $base_version" + exit 1 + fi version="$base_version-dev.$GITHUB_RUN_NUMBER" elif [[ "$GITHUB_REF" != refs/tags/v* ]]; then channel="${GITHUB_REF_NAME//[^0-9A-Za-z-]/-}" diff --git a/VERSION b/VERSION new file mode 100644 index 0000000..88c5fb8 --- /dev/null +++ b/VERSION @@ -0,0 +1 @@ +1.4.0 diff --git a/internal/sabr/multipart.go b/internal/sabr/multipart.go index 3015cf6..59a4242 100644 --- a/internal/sabr/multipart.go +++ b/internal/sabr/multipart.go @@ -18,17 +18,28 @@ func downloadParts( ctx, cancel := context.WithCancel(ctx) defer cancel() failures := make(chan error, parts) - var workers sync.WaitGroup + partJobs := make(chan int, parts) for part := range parts { + partJobs <- part + } + close(partJobs) + var workers sync.WaitGroup + for range min(parts, maxConcurrentPartDownloads) { workers.Add(1) go func() { defer workers.Done() - if err := downloadPart(ctx, client, options, part, parts, progress); err != nil { - select { - case failures <- err: - default: + for part := range partJobs { + if ctx.Err() != nil { + return + } + if err := downloadPart(ctx, client, options, part, parts, progress); err != nil { + select { + case failures <- err: + default: + } + cancel() + return } - cancel() } }() } @@ -81,3 +92,5 @@ func downloadPart( } return closeTracks(tracks) } + +const maxConcurrentPartDownloads = 3 diff --git a/internal/sabr/multipart_test.go b/internal/sabr/multipart_test.go new file mode 100644 index 0000000..e196124 --- /dev/null +++ b/internal/sabr/multipart_test.go @@ -0,0 +1,56 @@ +package sabr + +import ( + "context" + "net/http" + "path/filepath" + "strconv" + "sync/atomic" + "testing" + "time" +) + +func TestDownloadBoundsConcurrentSABRParts(t *testing.T) { + var active atomic.Int64 + var peak atomic.Int64 + var requests atomic.Int64 + server := downloadTestServer(t, func(writer http.ResponseWriter, request *http.Request) { + current := active.Add(1) + defer active.Add(-1) + requests.Add(1) + for { + observed := peak.Load() + if current <= observed || peak.CompareAndSwap(observed, current) { + break + } + } + time.Sleep(20 * time.Millisecond) + part, err := strconv.Atoi(request.URL.Query().Get("part")) + if err != nil { + t.Error(err) + return + } + writeTestFrame(t, writer, frameInitialization, 140, 0, []byte("init")) + writeTestFrame(t, writer, frameMedia, 140, part+1, []byte("media")) + writeTestFrame(t, writer, frameComplete, 0, 0, nil) + }) + defer server.Close() + + output := filepath.Join(t.TempDir(), "audio.m4a") + err := Download(context.Background(), server.Client(), Options{ + ManifestURL: server.URL, + AudioItag: 140, + AudioOnly: true, + AudioPath: output, + Parts: maxDownloadParts, + }, nil) + if err != nil { + t.Fatal(err) + } + if requests.Load() != maxDownloadParts { + t.Fatalf("requests = %d, want %d", requests.Load(), maxDownloadParts) + } + if peak.Load() != maxConcurrentPartDownloads { + t.Fatalf("peak concurrency = %d, want %d", peak.Load(), maxConcurrentPartDownloads) + } +}