diff --git a/external/ytmusic/cookie_provider.go b/external/ytmusic/cookie_provider.go index b817d19..b164fc8 100644 --- a/external/ytmusic/cookie_provider.go +++ b/external/ytmusic/cookie_provider.go @@ -5,6 +5,7 @@ import ( "fmt" "strings" "sync" + "time" "github.com/bjarneo/cliamp/playlist" "github.com/bjarneo/cliamp/provider" @@ -29,21 +30,27 @@ const ( ) type cookieBase struct { - browser string - fetchFn func(browser string) ([]playlist.PlaylistInfo, error) - resolveFn func(pageURL string, start, count int, browser ...string) ([]playlist.Track, int, error) - mu sync.Mutex - playlists []playlist.PlaylistInfo - trackCache map[string][]playlist.Track + browser string + fetchFn func(browser string) ([]playlist.PlaylistInfo, error) + resolveFn func(ctx context.Context, pageURL string, start, count int, browser ...string) ([]playlist.Track, int, error) + mu sync.Mutex + playlists []playlist.PlaylistInfo + trackCache map[string][]playlist.Track + nextLoad uint64 + loadCancels map[uint64]context.CancelFunc } -const cookiePlaylistBatchSize = 100 +const ( + cookiePlaylistBatchSize = 100 + cookiePlaylistLoadTimeout = 5 * time.Minute +) func newCookieBase(browser string) *cookieBase { return &cookieBase{ - browser: browser, - fetchFn: resolve.FetchUserPlaylists, - trackCache: make(map[string][]playlist.Track), + browser: browser, + fetchFn: resolve.FetchUserPlaylists, + trackCache: make(map[string][]playlist.Track), + loadCancels: make(map[uint64]context.CancelFunc), } } @@ -82,15 +89,28 @@ func (b *cookieBase) fetchTracks(target string) ([]playlist.Track, error) { return cached, nil } + ctx, cancel := context.WithTimeout(context.Background(), cookiePlaylistLoadTimeout) + b.nextLoad++ + loadID := b.nextLoad + if b.loadCancels == nil { + b.loadCancels = make(map[uint64]context.CancelFunc) + } + b.loadCancels[loadID] = cancel b.mu.Unlock() + defer func() { + b.mu.Lock() + delete(b.loadCancels, loadID) + b.mu.Unlock() + cancel() + }() resolveBatch := b.resolveFn if resolveBatch == nil { - resolveBatch = resolve.ResolveYTDLBatchPage + resolveBatch = resolve.ResolveYTDLBatchPageContext } var tracks []playlist.Track for start := 0; ; { - batch, entries, err := resolveBatch(target, start, cookiePlaylistBatchSize, b.browser) + batch, entries, err := resolveBatch(ctx, target, start, cookiePlaylistBatchSize, b.browser) if err != nil { return nil, fmt.Errorf("ytmusic: resolve playlist tracks: %w", err) } @@ -109,11 +129,24 @@ func (b *cookieBase) fetchTracks(target string) ([]playlist.Track, error) { func (b *cookieBase) refresh() { b.mu.Lock() + for _, cancel := range b.loadCancels { + cancel() + } + clear(b.loadCancels) b.playlists = nil clear(b.trackCache) b.mu.Unlock() } +func (b *cookieBase) close() { + b.mu.Lock() + for _, cancel := range b.loadCancels { + cancel() + } + clear(b.loadCancels) + b.mu.Unlock() +} + // CookieProvider provides YouTube and YouTube Music playlist access using // browser cookies via yt-dlp, without requiring Google Cloud OAuth credentials. type CookieProvider struct { @@ -267,4 +300,4 @@ func (p *CookieProvider) Refresh() { } // Close releases any held resources. -func (p *CookieProvider) Close() {} +func (p *CookieProvider) Close() { p.base.close() } diff --git a/external/ytmusic/cookie_provider_test.go b/external/ytmusic/cookie_provider_test.go index 8452101..2043285 100644 --- a/external/ytmusic/cookie_provider_test.go +++ b/external/ytmusic/cookie_provider_test.go @@ -10,6 +10,7 @@ import ( "slices" "strings" "testing" + "time" "github.com/bjarneo/cliamp/playlist" "github.com/bjarneo/cliamp/provider" @@ -255,7 +256,7 @@ func TestCookieProviderTracksLoadsInBatches(t *testing.T) { t.Setenv("HOME", t.TempDir()) base := newCookieBase("firefox") var starts []int - base.resolveFn = func(_ string, start, count int, browser ...string) ([]playlist.Track, int, error) { + base.resolveFn = func(_ context.Context, _ string, start, count int, browser ...string) ([]playlist.Track, int, error) { starts = append(starts, start) if count != cookiePlaylistBatchSize { t.Fatalf("count = %d, want %d", count, cookiePlaylistBatchSize) @@ -295,6 +296,45 @@ func TestCookieProviderTracksLoadsInBatches(t *testing.T) { } } +func TestCookieProviderStopsTrackLoad(t *testing.T) { + tests := []struct { + name string + stop func(*CookieProvider) + }{ + {name: "refresh", stop: func(p *CookieProvider) { p.Refresh() }}, + {name: "close", stop: func(p *CookieProvider) { p.Close() }}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + base := newCookieBase("firefox") + started := make(chan struct{}) + base.resolveFn = func(ctx context.Context, _ string, _, _ int, _ ...string) ([]playlist.Track, int, error) { + close(started) + <-ctx.Done() + return nil, 0, ctx.Err() + } + prov := &CookieProvider{base: base, kind: KindMusic} + done := make(chan error, 1) + go func() { + _, err := prov.Tracks("PL123") + done <- err + }() + + <-started + tt.stop(prov) + select { + case err := <-done: + if !errors.Is(err, context.Canceled) { + t.Fatalf("Tracks() error = %v, want context.Canceled", err) + } + case <-time.After(time.Second): + t.Fatal("track load did not stop") + } + }) + } +} + func TestCookieProviderSearchTracksHonorsCancellation(t *testing.T) { if runtime.GOOS == "windows" { t.Skip("skipping Unix shell script test on Windows") diff --git a/resolve/resolve.go b/resolve/resolve.go index 1d008fa..7b03a09 100644 --- a/resolve/resolve.go +++ b/resolve/resolve.go @@ -621,6 +621,12 @@ func ResolveYTDLBatch(pageURL string, start, count int, browser ...string) ([]pl func ResolveYTDLBatchPage(pageURL string, start, count int, browser ...string) ([]playlist.Track, int, error) { ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() + return ResolveYTDLBatchPageContext(ctx, pageURL, start, count, browser...) +} + +// ResolveYTDLBatchPageContext is ResolveYTDLBatchPage with caller-controlled +// cancellation and timeout. +func ResolveYTDLBatchPageContext(ctx context.Context, pageURL string, start, count int, browser ...string) ([]playlist.Track, int, error) { end := 0 if count > 0 { end = start + count