From 685392f0730b0cc2c0a0f0e20be317c3437a77d1 Mon Sep 17 00:00:00 2001 From: zarzet Date: Thu, 23 Jul 2026 09:56:39 +0700 Subject: [PATCH] fix(download): accelerate verified extension resumes --- go_backend/download_preparation_cache.go | 129 ++++++ go_backend/download_preparation_cache_test.go | 104 +++++ go_backend/exports_extensions.go | 61 ++- go_backend/extension_fallback.go | 383 +++++++++++++----- go_backend/extension_fallback_resume_test.go | 122 ++++++ go_backend/extension_signed_session.go | 63 ++- go_backend/extension_signed_session_test.go | 190 +++++++++ go_backend/progress.go | 16 + go_backend/progress_test.go | 26 ++ 9 files changed, 979 insertions(+), 115 deletions(-) create mode 100644 go_backend/download_preparation_cache.go create mode 100644 go_backend/download_preparation_cache_test.go create mode 100644 go_backend/extension_fallback_resume_test.go diff --git a/go_backend/download_preparation_cache.go b/go_backend/download_preparation_cache.go new file mode 100644 index 00000000..0175c21f --- /dev/null +++ b/go_backend/download_preparation_cache.go @@ -0,0 +1,129 @@ +package gobackend + +import ( + "strings" + "sync" + "time" +) + +const ( + downloadPreparationCacheTTL = 5 * time.Minute + downloadPreparationCacheMax = 128 +) + +type preparedDownloadRequestEntry struct { + key string + request DownloadRequest + metadataPrepared bool + createdAt time.Time +} + +var ( + preparedDownloadRequests = make(map[string]preparedDownloadRequestEntry) + preparedDownloadRequestsMu sync.Mutex +) + +func downloadPreparationKey(req DownloadRequest) string { + return strings.Join([]string{ + strings.TrimSpace(req.ItemID), + strings.ToLower(strings.TrimSpace(req.Service)), + strings.ToLower(strings.TrimSpace(req.Source)), + strings.TrimSpace(req.SpotifyID), + strings.TrimSpace(req.TidalID), + strings.TrimSpace(req.QobuzID), + strings.TrimSpace(req.DeezerID), + strings.ToLower(strings.TrimSpace(req.TrackName)), + strings.ToLower(strings.TrimSpace(req.ArtistName)), + }, "\n") +} + +func prunePreparedDownloadRequestsLocked(now time.Time) { + for itemID, entry := range preparedDownloadRequests { + if now.Sub(entry.createdAt) >= downloadPreparationCacheTTL { + delete(preparedDownloadRequests, itemID) + } + } + for len(preparedDownloadRequests) >= downloadPreparationCacheMax { + var oldestID string + var oldestAt time.Time + for itemID, entry := range preparedDownloadRequests { + if oldestID == "" || entry.createdAt.Before(oldestAt) { + oldestID = itemID + oldestAt = entry.createdAt + } + } + if oldestID == "" { + break + } + delete(preparedDownloadRequests, oldestID) + } +} + +func cacheDownloadRequestForVerification(key string, req DownloadRequest, metadataPrepared bool) { + itemID := strings.TrimSpace(req.ItemID) + if itemID == "" || strings.TrimSpace(key) == "" { + return + } + + preparedDownloadRequestsMu.Lock() + defer preparedDownloadRequestsMu.Unlock() + now := time.Now() + prunePreparedDownloadRequestsLocked(now) + preparedDownloadRequests[itemID] = preparedDownloadRequestEntry{ + key: key, + request: req, + metadataPrepared: metadataPrepared, + createdAt: now, + } +} + +func cachePreparedDownloadRequest(key string, req DownloadRequest) { + cacheDownloadRequestForVerification(key, req, true) +} + +func cacheUnpreparedDownloadRequest(key string, req DownloadRequest) { + cacheDownloadRequestForVerification(key, req, false) +} + +func takePreparedDownloadRequest(key string, fresh DownloadRequest) (DownloadRequest, bool, bool) { + itemID := strings.TrimSpace(fresh.ItemID) + if itemID == "" || strings.TrimSpace(key) == "" { + return fresh, false, false + } + + preparedDownloadRequestsMu.Lock() + defer preparedDownloadRequestsMu.Unlock() + now := time.Now() + prunePreparedDownloadRequestsLocked(now) + entry, ok := preparedDownloadRequests[itemID] + if !ok { + return fresh, false, false + } + delete(preparedDownloadRequests, itemID) + if entry.key != key { + return fresh, false, false + } + + prepared := entry.request + fresh.ISRC = prepared.ISRC + fresh.SpotifyID = prepared.SpotifyID + fresh.TrackName = prepared.TrackName + fresh.ArtistName = prepared.ArtistName + fresh.AlbumName = prepared.AlbumName + fresh.AlbumArtist = prepared.AlbumArtist + fresh.CoverURL = prepared.CoverURL + fresh.TrackNumber = prepared.TrackNumber + fresh.DiscNumber = prepared.DiscNumber + fresh.TotalTracks = prepared.TotalTracks + fresh.TotalDiscs = prepared.TotalDiscs + fresh.ReleaseDate = prepared.ReleaseDate + fresh.DurationMS = prepared.DurationMS + fresh.Genre = prepared.Genre + fresh.Label = prepared.Label + fresh.Copyright = prepared.Copyright + fresh.Composer = prepared.Composer + fresh.TidalID = prepared.TidalID + fresh.QobuzID = prepared.QobuzID + fresh.DeezerID = prepared.DeezerID + return fresh, entry.metadataPrepared, true +} diff --git a/go_backend/download_preparation_cache_test.go b/go_backend/download_preparation_cache_test.go new file mode 100644 index 00000000..c97d3b18 --- /dev/null +++ b/go_backend/download_preparation_cache_test.go @@ -0,0 +1,104 @@ +package gobackend + +import ( + "testing" + "time" +) + +func resetPreparedDownloadRequestCacheForTest() { + preparedDownloadRequestsMu.Lock() + preparedDownloadRequests = make(map[string]preparedDownloadRequestEntry) + preparedDownloadRequestsMu.Unlock() +} + +func TestPreparedDownloadRequestCache(t *testing.T) { + t.Cleanup(resetPreparedDownloadRequestCacheForTest) + resetPreparedDownloadRequestCacheForTest() + + fresh := DownloadRequest{ + ItemID: "item-1", + Service: "provider-a", + Source: "source-a", + SpotifyID: "spotify-1", + TrackName: "Track", + ArtistName: "Artist", + OutputDir: "/new/output", + OutputPath: "/new/output/current.flac", + OutputFD: 42, + Quality: "lossless", + EmbedMetadata: true, + } + key := downloadPreparationKey(fresh) + prepared := fresh + prepared.ISRC = "USRC17607839" + prepared.AlbumName = "Resolved Album" + prepared.AlbumArtist = "Resolved Album Artist" + prepared.DeezerID = "deezer-1" + prepared.Genre = "Pop" + prepared.OutputDir = "/stale/output" + prepared.OutputPath = "/stale/output/old.flac" + prepared.OutputFD = 7 + prepared.Quality = "stale-quality" + prepared.EmbedMetadata = false + cachePreparedDownloadRequest(key, prepared) + + got, metadataPrepared, ok := takePreparedDownloadRequest(key, fresh) + if !ok { + t.Fatal("expected prepared request cache hit") + } + if !metadataPrepared { + t.Fatal("prepared request should be marked as metadata-prepared") + } + if got.ISRC != prepared.ISRC || got.AlbumName != prepared.AlbumName || got.DeezerID != prepared.DeezerID || got.Genre != prepared.Genre { + t.Fatalf("prepared metadata was not restored: %#v", got) + } + if got.OutputDir != fresh.OutputDir || got.OutputPath != fresh.OutputPath || got.OutputFD != fresh.OutputFD || got.Quality != fresh.Quality || got.EmbedMetadata != fresh.EmbedMetadata { + t.Fatalf("fresh output/settings fields were overwritten: %#v", got) + } + if _, _, ok := takePreparedDownloadRequest(key, fresh); ok { + t.Fatal("prepared request should be consumed after one retry") + } + + cacheUnpreparedDownloadRequest(key, fresh) + _, metadataPrepared, ok = takePreparedDownloadRequest(key, fresh) + if !ok || metadataPrepared { + t.Fatalf("unprepared verification request = hit:%v metadataPrepared:%v", ok, metadataPrepared) + } +} + +func TestPreparedDownloadRequestCacheRejectsChangedTrackAndExpiry(t *testing.T) { + t.Cleanup(resetPreparedDownloadRequestCacheForTest) + resetPreparedDownloadRequestCacheForTest() + + req := DownloadRequest{ + ItemID: "item-2", + Service: "provider-a", + SpotifyID: "spotify-2", + TrackName: "Track", + ArtistName: "Artist", + } + key := downloadPreparationKey(req) + cachePreparedDownloadRequest(key, req) + + changed := req + changed.SpotifyID = "spotify-other" + if _, _, ok := takePreparedDownloadRequest(downloadPreparationKey(changed), changed); ok { + t.Fatal("changed track must not reuse another track's prepared metadata") + } + preparedDownloadRequestsMu.Lock() + _, staleEntryExists := preparedDownloadRequests[req.ItemID] + preparedDownloadRequestsMu.Unlock() + if staleEntryExists { + t.Fatal("mismatched prepared request should be discarded") + } + + cachePreparedDownloadRequest(key, req) + preparedDownloadRequestsMu.Lock() + entry := preparedDownloadRequests[req.ItemID] + entry.createdAt = time.Now().Add(-downloadPreparationCacheTTL) + preparedDownloadRequests[req.ItemID] = entry + preparedDownloadRequestsMu.Unlock() + if _, _, ok := takePreparedDownloadRequest(key, req); ok { + t.Fatal("expired prepared request must not be reused") + } +} diff --git a/go_backend/exports_extensions.go b/go_backend/exports_extensions.go index 861cb09e..daf75c2b 100644 --- a/go_backend/exports_extensions.go +++ b/go_backend/exports_extensions.go @@ -451,6 +451,29 @@ func SearchTracksWithMetadataProvidersJSON(query string, limit int, includeExten return marshalJSONString(tracks) } +func preflightExtensionDownloadSession(extensionID string) (bool, error) { + extensionID = strings.TrimSpace(extensionID) + if extensionID == "" { + return false, nil + } + + ext, err := getExtensionManager().GetExtension(extensionID) + if err != nil || ext == nil || !ext.Enabled || ext.Manifest == nil || + !ext.Manifest.IsDownloadProvider() || ext.Manifest.SignedSession == nil { + return false, nil + } + + if _, err := ext.lockReadyVM(); err != nil { + return false, err + } + defer ext.VMMu.Unlock() + if ext.runtime == nil { + return false, fmt.Errorf("extension '%s' runtime is unavailable", extensionID) + } + + return ext.runtime.preflightSignedSession() +} + func DownloadWithExtensionsJSON(requestJSON string) (string, error) { var req DownloadRequest if err := json.Unmarshal([]byte(requestJSON), &req); err != nil { @@ -477,15 +500,51 @@ func DownloadWithExtensionsJSON(requestJSON string) (string, error) { AddAllowedDownloadDir(req.OutputDir) } - enrichRequestExtendedMetadata(&req) + sessionProvider := strings.TrimSpace(req.Service) + if sessionProvider == "" { + sessionProvider = strings.TrimSpace(req.Source) + } + if req.ItemID != "" { + StartItemProgress(req.ItemID) + SetItemPreparingStage(req.ItemID, "checking_session") + } + preflightStartedAt := time.Now() + verificationRequired, preflightErr := preflightExtensionDownloadSession(sessionProvider) + if preflightErr != nil { + GoLog("[DownloadWithExtensions] Signed-session preflight for %s was inconclusive after %s: %v\n", sessionProvider, time.Since(preflightStartedAt).Round(time.Millisecond), preflightErr) + } else if verificationRequired { + GoLog("[DownloadWithExtensions] Signed-session verification required for %s after %s; skipping metadata preparation\n", sessionProvider, time.Since(preflightStartedAt).Round(time.Millisecond)) + cacheUnpreparedDownloadRequest(downloadPreparationKey(req), req) + if req.ItemID != "" { + RemoveItemProgress(req.ItemID) + } + return marshalJSONString(&DownloadResponse{ + Success: false, + Error: "Verification required before download", + ErrorType: "verification_required", + Service: sessionProvider, + }) + } else if sessionProvider != "" { + LogDebug("DownloadWithExtensions", "Signed-session preflight ready for %s in %s", sessionProvider, time.Since(preflightStartedAt).Round(time.Millisecond)) + } + if isDownloadCancelled(req.ItemID) { + if req.ItemID != "" { + RemoveItemProgress(req.ItemID) + } return "", ErrDownloadCancelled } result, err := DownloadWithExtensionFallback(req) if err != nil { + if req.ItemID != "" { + RemoveItemProgress(req.ItemID) + } return "", err } + if req.ItemID != "" && (result == nil || !result.Success) { + RemoveItemProgress(req.ItemID) + } return marshalJSONString(result) } diff --git a/go_backend/extension_fallback.go b/go_backend/extension_fallback.go index 9bffd340..d65d63fe 100644 --- a/go_backend/extension_fallback.go +++ b/go_backend/extension_fallback.go @@ -6,6 +6,7 @@ import ( "os" "path/filepath" "strings" + "time" ) func manifestCapabilityStringList(manifest *ExtensionManifest, key string) []string { @@ -391,7 +392,7 @@ func attemptExtensionDownload( ) (resp *DownloadResponse, cancelledOuter bool) { outputPath := buildOutputPathForExtension(req, ext) if req.ItemID != "" { - StartItemProgress(req.ItemID) + SetItemPreparingStage(req.ItemID, "resolving_stream") } result, err := provider.Download(trackID, quality, outputPath, req.ItemID, func(percent int) { @@ -406,18 +407,25 @@ func attemptExtensionDownload( SetItemProgress(req.ItemID, normalized, 0, 0) } }) - if req.ItemID != "" { - if err == nil && result != nil && result.Success { - CompleteItemProgress(req.ItemID) - } else { - RemoveItemProgress(req.ItemID) - } + downloadSucceeded := err == nil && result != nil && result.Success + if req.ItemID != "" && downloadSucceeded { + SetItemFinalizing(req.ItemID) } if shouldAbortCancelledFallback(req.ItemID, err) { return nil, true } - if err == nil && result.Success { + if downloadSucceeded { + metadataStartedAt := time.Now() + enrichRequestExtendedMetadata(&req) + LogDebug( + "DownloadPipeline", + "item=%s provider=%s post-transfer metadataMs=%.1f", + req.ItemID, + providerLabel, + extensionDurationMs(time.Since(metadataStartedAt)), + ) + normalizedResult, alreadyExists := normalizeExtensionDownloadResult(result) message := "Downloaded from " + providerLabel if alreadyExists { @@ -462,6 +470,9 @@ func attemptExtensionDownload( } } + if req.ItemID != "" { + CompleteItemProgress(req.ItemID) + } return &built, false } @@ -476,15 +487,153 @@ func attemptExtensionDownload( } *lastErr = err *lastErrType = "" - } else if result.ErrorMessage != "" { + } else if result != nil && result.ErrorMessage != "" { *lastErr = fmt.Errorf("%s", result.ErrorMessage) *lastErrType = normalizeExtensionDownloadErrorType(result.ErrorType, result.ErrorMessage) *lastRetryAfterSeconds = result.RetryAfterSeconds + } else if result == nil { + *lastErr = fmt.Errorf("extension returned no download result") + *lastErrType = "extension_error" } return nil, false } +// attemptVerifiedResumeBeforeMetadata gives a just-verified request one direct +// chance against its selected provider before optional metadata providers run. +// Download-provider sources keep their normal precedence because they may +// dynamically lock fallback from checkAvailability. +func attemptVerifiedResumeBeforeMetadata( + req DownloadRequest, + selectedProvider string, + extManager *extensionManager, +) (*DownloadResponse, bool) { + selectedProvider = strings.TrimSpace(selectedProvider) + if selectedProvider == "" || extManager == nil { + return nil, false + } + + sourceProvider := strings.TrimSpace(req.Source) + if sourceProvider != "" && !strings.EqualFold(sourceProvider, selectedProvider) { + if sourceExt, err := extManager.GetExtension(sourceProvider); err == nil && + sourceExt != nil && sourceExt.Enabled && sourceExt.Error == "" && + sourceExt.Manifest != nil && sourceExt.Manifest.IsDownloadProvider() { + return nil, false + } + } + + ext, err := extManager.GetExtension(selectedProvider) + if err != nil || ext == nil || !ext.Enabled || ext.Error != "" || + ext.Manifest == nil || !ext.Manifest.IsDownloadProvider() { + return nil, false + } + + provider := newExtensionProviderWrapper(ext) + var availability *ExtAvailabilityResult + trackID := "" + if strings.EqualFold(sourceProvider, selectedProvider) { + trackID = resolvePreferredTrackIDForExtension(ext, req, "") + } else { + availability, err = provider.CheckAvailabilityForItemID( + req.ISRC, + req.TrackName, + req.ArtistName, + req.SpotifyID, + req.DeezerID, + req.TidalID, + req.QobuzID, + req.DurationMS, + req.ItemID, + ) + if shouldAbortCancelledFallback(req.ItemID, err) { + return nil, true + } + if err != nil { + if strings.EqualFold(classifyDownloadErrorType(err.Error()), "verification_required") { + return &DownloadResponse{ + Success: false, + Error: "Download failed: " + err.Error(), + ErrorType: "verification_required", + Service: selectedProvider, + }, false + } + return nil, false + } + if availability == nil || !availability.Available { + if shouldStopProviderFallback(availability) { + return buildExtensionFallbackStoppedResponse(selectedProvider, availability, nil), false + } + return nil, false + } + trackID = resolvePreferredTrackIDForExtension(ext, req, availability.TrackID) + } + + var lastErr error + var lastErrType string + var lastRetryAfterSeconds int + resp, cancelled := attemptExtensionDownload( + req, + ext, + provider, + trackID, + req.Quality, + selectedProvider, + strings.EqualFold(sourceProvider, selectedProvider), + &lastErr, + &lastErrType, + &lastRetryAfterSeconds, + ) + if cancelled || resp != nil { + return resp, cancelled + } + + errorType := lastErrType + if errorType == "" && lastErr != nil { + errorType = classifyDownloadErrorType(lastErr.Error()) + } + if strings.EqualFold(errorType, "verification_required") { + errorMessage := "Verification required" + if lastErr != nil { + errorMessage = "Download failed: " + lastErr.Error() + } + return &DownloadResponse{ + Success: false, + Error: errorMessage, + ErrorType: "verification_required", + RetryAfterSeconds: lastRetryAfterSeconds, + Service: selectedProvider, + }, false + } + // A failed fast attempt may still succeed after source enrichment resolves a + // better provider-native ID. The normal path below retains strict/stop- + // fallback semantics for that prepared retry. + return nil, false +} + func DownloadWithExtensionFallback(req DownloadRequest) (*DownloadResponse, error) { + pipelineStartedAt := time.Now() + defer func() { + LogDebug( + "DownloadPipeline", + "item=%s service=%s totalMs=%.1f", + req.ItemID, + req.Service, + extensionDurationMs(time.Since(pipelineStartedAt)), + ) + }() + preparationKey := downloadPreparationKey(req) + metadataPrepared := false + resumedAfterVerification := false + if prepared, preparedMetadata, ok := takePreparedDownloadRequest(preparationKey, req); ok { + req = prepared + metadataPrepared = preparedMetadata + resumedAfterVerification = true + GoLog("[DownloadWithExtensionFallback] Resuming item %s after verification (metadata prepared: %v)\n", req.ItemID, metadataPrepared) + } + if req.ItemID != "" { + StartItemProgress(req.ItemID) + SetItemPreparingStage(req.ItemID, "resolving_metadata") + } + priority := GetProviderPriority() extManager := getExtensionManager() strictMode := !req.UseFallback @@ -535,6 +684,20 @@ func DownloadWithExtensionFallback(req DownloadRequest) (*DownloadResponse, erro var sourceExtensionAvailability *ExtAvailabilityResult var sourceExtensionTrackID string + if resumedAfterVerification && !metadataPrepared { + GoLog("[DownloadWithExtensionFallback] Trying verified provider %s before optional metadata enrichment\n", selectedProvider) + resp, cancelled := attemptVerifiedResumeBeforeMetadata(req, selectedProvider, extManager) + if cancelled { + return nil, ErrDownloadCancelled + } + if resp != nil { + return resp, nil + } + if req.ItemID != "" { + SetItemPreparingStage(req.ItemID, "resolving_metadata") + } + } + if req.Source != "" && selectedProvider != req.Source { ext, err := extManager.GetExtension(req.Source) if err == nil && ext.Enabled && ext.Error == "" && ext.Manifest.IsDownloadProvider() { @@ -555,115 +718,112 @@ func DownloadWithExtensionFallback(req DownloadRequest) (*DownloadResponse, erro } } - if req.Source != "" { - ext, err := extManager.GetExtension(req.Source) - if err == nil && ext.Enabled && ext.Error == "" && ext.Manifest.IsMetadataProvider() { - GoLog("[DownloadWithExtensionFallback] Enriching track from extension '%s'...\n", req.Source) + if !metadataPrepared { + if req.Source != "" { + ext, err := extManager.GetExtension(req.Source) + if err == nil && ext.Enabled && ext.Error == "" && ext.Manifest.IsMetadataProvider() { + GoLog("[DownloadWithExtensionFallback] Enriching track from extension '%s'...\n", req.Source) - provider := newExtensionProviderWrapper(ext) - trackMeta := &ExtTrackMetadata{ - ID: req.SpotifyID, - Name: req.TrackName, - Artists: req.ArtistName, - AlbumName: req.AlbumName, - DurationMS: req.DurationMS, - ISRC: req.ISRC, - ReleaseDate: req.ReleaseDate, - TrackNumber: req.TrackNumber, - TotalTracks: req.TotalTracks, - DiscNumber: req.DiscNumber, - TotalDiscs: req.TotalDiscs, - ProviderID: req.Source, - Composer: req.Composer, + provider := newExtensionProviderWrapper(ext) + trackMeta := &ExtTrackMetadata{ + ID: req.SpotifyID, + Name: req.TrackName, + Artists: req.ArtistName, + AlbumName: req.AlbumName, + DurationMS: req.DurationMS, + ISRC: req.ISRC, + ReleaseDate: req.ReleaseDate, + TrackNumber: req.TrackNumber, + TotalTracks: req.TotalTracks, + DiscNumber: req.DiscNumber, + TotalDiscs: req.TotalDiscs, + ProviderID: req.Source, + Composer: req.Composer, + } + + enrichedTrack, err := provider.EnrichTrackForItemID(trackMeta, req.ItemID) + if shouldAbortCancelledFallback(req.ItemID, err) { + return nil, ErrDownloadCancelled + } + if err == nil && enrichedTrack != nil { + if enrichedTrack.ISRC != "" && enrichedTrack.ISRC != req.ISRC { + GoLog("[DownloadWithExtensionFallback] ISRC enriched: %s -> %s\n", req.ISRC, enrichedTrack.ISRC) + req.ISRC = enrichedTrack.ISRC + } + if enrichedTrack.TidalID != "" { + GoLog("[DownloadWithExtensionFallback] Tidal ID from Odesli: %s\n", enrichedTrack.TidalID) + req.TidalID = enrichedTrack.TidalID + } + if enrichedTrack.QobuzID != "" { + GoLog("[DownloadWithExtensionFallback] Qobuz ID from Odesli: %s\n", enrichedTrack.QobuzID) + req.QobuzID = enrichedTrack.QobuzID + } + if enrichedTrack.DeezerID != "" { + GoLog("[DownloadWithExtensionFallback] Deezer ID from Odesli: %s\n", enrichedTrack.DeezerID) + req.DeezerID = enrichedTrack.DeezerID + } + if enrichedTrack.Name != "" { + req.TrackName = enrichedTrack.Name + } + if enrichedTrack.Artists != "" { + req.ArtistName = enrichedTrack.Artists + } + overlayStr(&req.AlbumName, enrichedTrack.AlbumName, "AlbumName") + overlayStr(&req.AlbumArtist, enrichedTrack.AlbumArtist, "") + overlayInt(&req.DurationMS, enrichedTrack.DurationMS, "DurationMS") + overlayStr(&req.CoverURL, enrichedTrack.CoverURL, "") + overlayStr(&req.SpotifyID, enrichedTrack.ID, "Track ID") + overlayStr(&req.Label, enrichedTrack.Label, "Label") + overlayStr(&req.Copyright, enrichedTrack.Copyright, "Copyright") + overlayStr(&req.Genre, enrichedTrack.Genre, "Genre") + overlayStr(&req.ReleaseDate, enrichedTrack.ReleaseDate, "ReleaseDate") + overlayInt(&req.TrackNumber, enrichedTrack.TrackNumber, "TrackNumber") + overlayInt(&req.TotalTracks, enrichedTrack.TotalTracks, "TotalTracks") + overlayInt(&req.DiscNumber, enrichedTrack.DiscNumber, "DiscNumber") + overlayInt(&req.TotalDiscs, enrichedTrack.TotalDiscs, "TotalDiscs") + overlayStr(&req.Composer, enrichedTrack.Composer, "Composer") + } } + } - enrichedTrack, err := provider.EnrichTrackForItemID(trackMeta, req.ItemID) - if shouldAbortCancelledFallback(req.ItemID, err) { + if req.Source != "" && + req.TrackName != "" && req.ArtistName != "" && + (req.AlbumName == "" || req.ReleaseDate == "" || req.ISRC == "") { + + searchQuery := req.TrackName + " " + req.ArtistName + GoLog("[DownloadWithExtensionFallback] Metadata incomplete, searching providers for: %s\n", searchQuery) + + // Only the first match is consumed below. Asking for five made the manager + // continue through additional providers even after it already had a usable + // match, multiplying the per-provider timeout on slow networks. + tracks, searchErr := extManager.SearchTracksWithMetadataProvidersForItemID(searchQuery, 1, true, req.ItemID) + if shouldAbortCancelledFallback(req.ItemID, searchErr) { return nil, ErrDownloadCancelled } - if err == nil && enrichedTrack != nil { - if enrichedTrack.ISRC != "" && enrichedTrack.ISRC != req.ISRC { - GoLog("[DownloadWithExtensionFallback] ISRC enriched: %s -> %s\n", req.ISRC, enrichedTrack.ISRC) - req.ISRC = enrichedTrack.ISRC - } - if enrichedTrack.TidalID != "" { - GoLog("[DownloadWithExtensionFallback] Tidal ID from Odesli: %s\n", enrichedTrack.TidalID) - req.TidalID = enrichedTrack.TidalID - } - if enrichedTrack.QobuzID != "" { - GoLog("[DownloadWithExtensionFallback] Qobuz ID from Odesli: %s\n", enrichedTrack.QobuzID) - req.QobuzID = enrichedTrack.QobuzID - } - if enrichedTrack.DeezerID != "" { - GoLog("[DownloadWithExtensionFallback] Deezer ID from Odesli: %s\n", enrichedTrack.DeezerID) - req.DeezerID = enrichedTrack.DeezerID - } - if enrichedTrack.Name != "" { - req.TrackName = enrichedTrack.Name - } - if enrichedTrack.Artists != "" { - req.ArtistName = enrichedTrack.Artists - } - overlayStr(&req.AlbumName, enrichedTrack.AlbumName, "AlbumName") - overlayStr(&req.AlbumArtist, enrichedTrack.AlbumArtist, "") - overlayInt(&req.DurationMS, enrichedTrack.DurationMS, "DurationMS") - overlayStr(&req.CoverURL, enrichedTrack.CoverURL, "") - overlayStr(&req.SpotifyID, enrichedTrack.ID, "Track ID") - overlayStr(&req.Label, enrichedTrack.Label, "Label") - overlayStr(&req.Copyright, enrichedTrack.Copyright, "Copyright") - overlayStr(&req.Genre, enrichedTrack.Genre, "Genre") - overlayStr(&req.ReleaseDate, enrichedTrack.ReleaseDate, "ReleaseDate") - overlayInt(&req.TrackNumber, enrichedTrack.TrackNumber, "TrackNumber") - overlayInt(&req.TotalTracks, enrichedTrack.TotalTracks, "TotalTracks") - overlayInt(&req.DiscNumber, enrichedTrack.DiscNumber, "DiscNumber") - overlayInt(&req.TotalDiscs, enrichedTrack.TotalDiscs, "TotalDiscs") - overlayStr(&req.Composer, enrichedTrack.Composer, "Composer") + if searchErr == nil && len(tracks) > 0 { + track := tracks[0] + GoLog("[DownloadWithExtensionFallback] Metadata match (%s): %s - %s (album: %s, date: %s, isrc: %s)\n", + track.ProviderID, track.Name, track.Artists, track.AlbumName, track.ReleaseDate, track.ISRC) + + overlayStr(&req.AlbumName, track.AlbumName, "") + overlayStr(&req.AlbumArtist, track.AlbumArtist, "") + overlayStr(&req.ReleaseDate, track.ReleaseDate, "") + overlayStr(&req.ISRC, track.ISRC, "") + overlayInt(&req.TrackNumber, track.TrackNumber, "") + overlayInt(&req.TotalTracks, track.TotalTracks, "") + overlayInt(&req.DiscNumber, track.DiscNumber, "") + overlayInt(&req.TotalDiscs, track.TotalDiscs, "") + overlayStr(&req.Composer, track.Composer, "") + overlayStr(&req.CoverURL, track.CoverURL, "") + overlayStr(&req.Genre, track.Genre, "") + overlayStr(&req.Label, track.Label, "") + overlayStr(&req.Copyright, track.Copyright, "") + } else if searchErr != nil { + GoLog("[DownloadWithExtensionFallback] Metadata provider search failed (non-fatal): %v\n", searchErr) } + } } - - if req.Source != "" && - req.TrackName != "" && req.ArtistName != "" && - (req.AlbumName == "" || req.ReleaseDate == "" || req.ISRC == "") { - - searchQuery := req.TrackName + " " + req.ArtistName - GoLog("[DownloadWithExtensionFallback] Metadata incomplete, searching providers for: %s\n", searchQuery) - - tracks, searchErr := extManager.SearchTracksWithMetadataProvidersForItemID(searchQuery, 5, true, req.ItemID) - if shouldAbortCancelledFallback(req.ItemID, searchErr) { - return nil, ErrDownloadCancelled - } - if searchErr == nil && len(tracks) > 0 { - track := tracks[0] - GoLog("[DownloadWithExtensionFallback] Metadata match (%s): %s - %s (album: %s, date: %s, isrc: %s)\n", - track.ProviderID, track.Name, track.Artists, track.AlbumName, track.ReleaseDate, track.ISRC) - - overlayStr(&req.AlbumName, track.AlbumName, "") - overlayStr(&req.AlbumArtist, track.AlbumArtist, "") - overlayStr(&req.ReleaseDate, track.ReleaseDate, "") - overlayStr(&req.ISRC, track.ISRC, "") - overlayInt(&req.TrackNumber, track.TrackNumber, "") - overlayInt(&req.TotalTracks, track.TotalTracks, "") - overlayInt(&req.DiscNumber, track.DiscNumber, "") - overlayInt(&req.TotalDiscs, track.TotalDiscs, "") - overlayStr(&req.Composer, track.Composer, "") - overlayStr(&req.CoverURL, track.CoverURL, "") - overlayStr(&req.Genre, track.Genre, "") - overlayStr(&req.Label, track.Label, "") - overlayStr(&req.Copyright, track.Copyright, "") - } else if searchErr != nil { - GoLog("[DownloadWithExtensionFallback] Metadata provider search failed (non-fatal): %v\n", searchErr) - } - - if req.ISRC != "" && - (req.Genre == "" || req.Label == "" || req.Copyright == "") { - enrichExtraMetadataByISRC("DownloadWithExtensionFallback", req.ISRC, &req.Genre, &req.Label, &req.Copyright) - if isDownloadCancelled(req.ItemID) { - return nil, ErrDownloadCancelled - } - } - } - if req.Source != "" && selectedProvider == req.Source { if isDownloadCancelled(req.ItemID) { return nil, ErrDownloadCancelled @@ -701,6 +861,7 @@ func DownloadWithExtensionFallback(req DownloadRequest) (*DownloadResponse, erro } if strings.EqualFold(sourceErrType, "verification_required") { GoLog("[DownloadWithExtensionFallback] Source extension %s requires verification, not trying other providers\n", req.Source) + cachePreparedDownloadRequest(preparationKey, req) return &DownloadResponse{ Success: false, Error: "Download failed: " + lastErr.Error(), @@ -776,6 +937,7 @@ func DownloadWithExtensionFallback(req DownloadRequest) (*DownloadResponse, erro lastErr = err if strings.EqualFold(classifyDownloadErrorType(err.Error()), "verification_required") { GoLog("[DownloadWithExtensionFallback] %s requires verification (availability); pausing fallback to open the challenge\n", providerID) + cachePreparedDownloadRequest(preparationKey, req) return &DownloadResponse{ Success: false, Error: "Download failed: " + err.Error(), @@ -832,6 +994,7 @@ func DownloadWithExtensionFallback(req DownloadRequest) (*DownloadResponse, erro } if strings.EqualFold(effType, "verification_required") { GoLog("[DownloadWithExtensionFallback] %s requires verification; pausing fallback to open the challenge\n", providerID) + cachePreparedDownloadRequest(preparationKey, req) return &DownloadResponse{ Success: false, Error: "Download failed: " + lastErr.Error(), diff --git a/go_backend/extension_fallback_resume_test.go b/go_backend/extension_fallback_resume_test.go new file mode 100644 index 00000000..0d0588f3 --- /dev/null +++ b/go_backend/extension_fallback_resume_test.go @@ -0,0 +1,122 @@ +package gobackend + +import ( + "path/filepath" + "testing" +) + +func TestVerifiedDownloadResumeTriesSelectedProviderBeforeMetadata(t *testing.T) { + resetPreparedDownloadRequestCacheForTest() + t.Cleanup(resetPreparedDownloadRequestCacheForTest) + + metadataExt := newTestLoadedExtension(t, ExtensionTypeMetadataProvider) + metadataExt.ID = "resume-metadata" + metadataExt.Manifest.Name = metadataExt.ID + downloadExt := newTestLoadedExtension(t, ExtensionTypeDownloadProvider) + downloadExt.ID = "resume-download" + downloadExt.Manifest.Name = downloadExt.ID + + manager := getExtensionManager() + manager.mu.Lock() + previousExtensions := manager.extensions + manager.extensions = map[string]*loadedExtension{ + metadataExt.ID: metadataExt, + downloadExt.ID: downloadExt, + } + manager.mu.Unlock() + t.Cleanup(func() { + teardownExtension(metadataExt) + teardownExtension(downloadExt) + manager.mu.Lock() + manager.extensions = previousExtensions + manager.mu.Unlock() + }) + + req := DownloadRequest{ + ItemID: "resume-item", + Service: downloadExt.ID, + Source: metadataExt.ID, + SpotifyID: "spotify:track:1", + TrackName: "Original Song", + ArtistName: "Artist", + AlbumName: "Album", + ReleaseDate: "2026-05-04", + OutputDir: t.TempDir(), + OutputExt: ".flac", + FilenameFormat: "{title}", + Quality: "LOSSLESS", + UseFallback: false, + } + key := downloadPreparationKey(req) + cacheUnpreparedDownloadRequest(key, req) + + resp, err := DownloadWithExtensionFallback(req) + if err != nil { + t.Fatalf("DownloadWithExtensionFallback: %v", err) + } + if resp == nil || !resp.Success { + t.Fatalf("resume response = %#v", resp) + } + if got := filepath.Base(resp.FilePath); got != "Original Song.flac" { + t.Fatalf("resume ran metadata enrichment before download: file = %q", got) + } +} + +func TestVerifiedDownloadResumeReusesPreparedMetadata(t *testing.T) { + resetPreparedDownloadRequestCacheForTest() + t.Cleanup(resetPreparedDownloadRequestCacheForTest) + + metadataExt := newTestLoadedExtension(t, ExtensionTypeMetadataProvider) + metadataExt.ID = "prepared-metadata" + metadataExt.Manifest.Name = metadataExt.ID + downloadExt := newTestLoadedExtension(t, ExtensionTypeDownloadProvider) + downloadExt.ID = "prepared-download" + downloadExt.Manifest.Name = downloadExt.ID + + manager := getExtensionManager() + manager.mu.Lock() + previousExtensions := manager.extensions + manager.extensions = map[string]*loadedExtension{ + metadataExt.ID: metadataExt, + downloadExt.ID: downloadExt, + } + manager.mu.Unlock() + t.Cleanup(func() { + teardownExtension(metadataExt) + teardownExtension(downloadExt) + manager.mu.Lock() + manager.extensions = previousExtensions + manager.mu.Unlock() + }) + + req := DownloadRequest{ + ItemID: "prepared-item", + Service: downloadExt.ID, + Source: metadataExt.ID, + SpotifyID: "spotify:track:2", + TrackName: "Original Song", + ArtistName: "Artist", + AlbumName: "Album", + ReleaseDate: "2026-05-04", + OutputDir: t.TempDir(), + OutputExt: ".flac", + FilenameFormat: "{title}", + Quality: "LOSSLESS", + UseFallback: false, + } + key := downloadPreparationKey(req) + prepared := req + prepared.TrackName = "Prepared Song" + cachePreparedDownloadRequest(key, prepared) + + resp, err := DownloadWithExtensionFallback(req) + if err != nil { + t.Fatalf("DownloadWithExtensionFallback: %v", err) + } + if resp == nil || !resp.Success { + t.Fatalf("prepared response = %#v", resp) + } + if got := filepath.Base(resp.FilePath); got != "Prepared Song.flac" { + t.Fatalf("prepared metadata was not reused: file = %q", got) + } +} diff --git a/go_backend/extension_signed_session.go b/go_backend/extension_signed_session.go index 8c8ab9c4..4cab0564 100644 --- a/go_backend/extension_signed_session.go +++ b/go_backend/extension_signed_session.go @@ -201,6 +201,64 @@ func parseSignedSessionTime(value string) (time.Time, bool) { return time.Time{}, false } +func signedSessionRecordIsUsable(record *signedSessionRecord) bool { + if record == nil || strings.TrimSpace(record.SessionID) == "" || strings.TrimSpace(record.SessionSecret) == "" { + return false + } + if expiresAt, ok := parseSignedSessionTime(record.ExpiresAt); ok { + return time.Now().Before(expiresAt) + } + return true +} + +// preflightSignedSession prepares a signed session before download metadata +// enrichment starts. A fresh pending challenge is reused, while bootstrap +// responses that can issue a session silently are accepted without prompting +// the user. Bootstrap failures remain non-fatal to the caller so the normal +// provider path can still surface its more specific error. +func (r *extensionRuntime) preflightSignedSession() (bool, error) { + if r == nil || r.manifest == nil || r.manifest.SignedSession == nil { + return false, nil + } + + config := signedSessionConfigWithDefaults(r.manifest.SignedSession) + if config.Namespace == "" || config.BaseURL == "" { + return false, fmt.Errorf("signedSession is not configured") + } + + record, err := r.loadSignedSession(config) + if err != nil { + return false, err + } + if signedSessionRecordIsUsable(record) { + return false, nil + } + + if pending := GetPendingAuthRequest(r.extensionID); pending != nil { + if time.Since(pending.CreatedAt) < pendingAuthRequestTTL && + strings.TrimSpace(pending.AuthURL) != "" { + return true, nil + } + ClearPendingAuthRequest(r.extensionID) + } + + if authURL := r.startSignedSessionVerification(config, "download-preflight"); authURL != "" { + return true, nil + } + + // Bootstrap may provision a session directly instead of returning a + // challenge. Reload the record before treating the empty URL as a failure. + record, err = r.loadSignedSession(config) + if err != nil { + return false, err + } + if signedSessionRecordIsUsable(record) { + return false, nil + } + + return false, fmt.Errorf("signed-session bootstrap did not return a session or verification challenge") +} + func (r *extensionRuntime) signedSessionStatus(call goja.FunctionCall) goja.Value { config := signedSessionConfigWithDefaults(r.manifest.SignedSession) if config.Namespace == "" || config.BaseURL == "" { @@ -210,10 +268,7 @@ func (r *extensionRuntime) signedSessionStatus(call goja.FunctionCall) goja.Valu if err != nil { return r.vm.ToValue(map[string]any{"authenticated": false, "error": err.Error()}) } - authenticated := record.SessionID != "" && record.SessionSecret != "" - if expiresAt, ok := parseSignedSessionTime(record.ExpiresAt); ok && time.Now().After(expiresAt) { - authenticated = false - } + authenticated := signedSessionRecordIsUsable(record) return r.vm.ToValue(map[string]any{ "authenticated": authenticated, "expires_at": record.ExpiresAt, diff --git a/go_backend/extension_signed_session_test.go b/go_backend/extension_signed_session_test.go index 9f44ab6d..f5427190 100644 --- a/go_backend/extension_signed_session_test.go +++ b/go_backend/extension_signed_session_test.go @@ -126,6 +126,196 @@ func TestParseSignedSessionTime(t *testing.T) { } } +func TestPreflightSignedSession(t *testing.T) { + t.Run("reuses a valid session without network", func(t *testing.T) { + calls := 0 + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + calls++ + return nil, fmt.Errorf("unexpected request: %s", req.URL) + }) + runtime := newSignedSessionTestRuntime(t, "preflight-valid", transport) + runtime.manifest.SignedSession = &SignedSessionConfig{ + Namespace: "preflight-valid", + BaseURL: "https://auth.example.com", + } + config := signedSessionConfigWithDefaults(runtime.manifest.SignedSession) + record, err := runtime.loadSignedSession(config) + if err != nil { + t.Fatalf("load session: %v", err) + } + record.SessionID = "session" + record.SessionSecret = "secret" + record.ExpiresAt = time.Now().Add(time.Hour).UTC().Format(time.RFC3339) + if err := runtime.saveSignedSession(config, record); err != nil { + t.Fatalf("save session: %v", err) + } + + verificationRequired, err := runtime.preflightSignedSession() + if err != nil || verificationRequired { + t.Fatalf("preflight = verification:%v error:%v", verificationRequired, err) + } + if calls != 0 { + t.Fatalf("valid session made %d network request(s)", calls) + } + }) + + t.Run("reuses a fresh pending challenge", func(t *testing.T) { + runtime := newSignedSessionTestRuntime(t, "preflight-pending", roundTripFunc(func(req *http.Request) (*http.Response, error) { + return nil, fmt.Errorf("unexpected request: %s", req.URL) + })) + runtime.manifest.SignedSession = &SignedSessionConfig{ + Namespace: "preflight-pending", + BaseURL: "https://auth.example.com", + } + pendingAuthRequestsMu.Lock() + pendingAuthRequests[runtime.extensionID] = &PendingAuthRequest{ + ExtensionID: runtime.extensionID, + AuthURL: "https://auth.example.com/challenge", + CreatedAt: time.Now(), + } + pendingAuthRequestsMu.Unlock() + t.Cleanup(func() { ClearPendingAuthRequest(runtime.extensionID) }) + + verificationRequired, err := runtime.preflightSignedSession() + if err != nil || !verificationRequired { + t.Fatalf("preflight = verification:%v error:%v", verificationRequired, err) + } + }) + + t.Run("bootstraps a challenge for an unauthenticated session", func(t *testing.T) { + calls := 0 + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + calls++ + payload, _ := json.Marshal(signedSessionExchangeResponse{ + ChallengeURL: "https://auth.example.com/challenge", + }) + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(string(payload))), + Request: req, + }, nil + }) + runtime := newSignedSessionTestRuntime(t, "preflight-challenge", transport) + runtime.manifest.SignedSession = &SignedSessionConfig{ + Namespace: "preflight-challenge", + BaseURL: "https://auth.example.com", + } + t.Cleanup(func() { ClearPendingAuthRequest(runtime.extensionID) }) + + verificationRequired, err := runtime.preflightSignedSession() + if err != nil || !verificationRequired { + t.Fatalf("preflight = verification:%v error:%v", verificationRequired, err) + } + if calls != 1 { + t.Fatalf("bootstrap calls = %d, want 1", calls) + } + if pending := GetPendingAuthRequest(runtime.extensionID); pending == nil || pending.AuthURL == "" { + t.Fatalf("pending challenge = %#v", pending) + } + }) + + t.Run("accepts a session issued directly by bootstrap", func(t *testing.T) { + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + payload, _ := json.Marshal(signedSessionExchangeResponse{ + SessionID: "boot-session", + SessionSecret: "boot-secret", + ExpiresAt: time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }) + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(string(payload))), + Request: req, + }, nil + }) + runtime := newSignedSessionTestRuntime(t, "preflight-direct", transport) + runtime.manifest.SignedSession = &SignedSessionConfig{ + Namespace: "preflight-direct", + BaseURL: "https://auth.example.com", + } + + verificationRequired, err := runtime.preflightSignedSession() + if err != nil || verificationRequired { + t.Fatalf("preflight = verification:%v error:%v", verificationRequired, err) + } + config := signedSessionConfigWithDefaults(runtime.manifest.SignedSession) + record, err := runtime.loadSignedSession(config) + if err != nil || record.SessionID != "boot-session" { + t.Fatalf("bootstrapped session = %#v error:%v", record, err) + } + }) +} + +func TestDownloadWithExtensionsPreflightsBeforeMetadataEnrichment(t *testing.T) { + extensionID := "preflight-download" + itemID := "preflight-item" + RemoveItemProgress(itemID) + t.Cleanup(func() { RemoveItemProgress(itemID) }) + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + payload, _ := json.Marshal(signedSessionExchangeResponse{ + ChallengeURL: "https://auth.example.com/challenge", + }) + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(string(payload))), + Request: req, + }, nil + }) + runtime := newSignedSessionTestRuntime(t, extensionID, transport) + manifest := &ExtensionManifest{ + Name: extensionID, + Types: []ExtensionType{ExtensionTypeDownloadProvider}, + SignedSession: &SignedSessionConfig{ + Namespace: extensionID, + BaseURL: "https://auth.example.com", + }, + } + runtime.manifest = manifest + ext := &loadedExtension{ + ID: extensionID, + Manifest: manifest, + VM: runtime.vm, + runtime: runtime, + initialized: true, + Enabled: true, + DataDir: runtime.dataDir, + } + + manager := getExtensionManager() + manager.mu.Lock() + previous, hadPrevious := manager.extensions[extensionID] + manager.extensions[extensionID] = ext + manager.mu.Unlock() + t.Cleanup(func() { + ClearPendingAuthRequest(extensionID) + manager.mu.Lock() + if hadPrevious { + manager.extensions[extensionID] = previous + } else { + delete(manager.extensions, extensionID) + } + manager.mu.Unlock() + }) + + requestJSON := `{"service":"preflight-download","item_id":"preflight-item","isrc":"USRC17607839"}` + responseJSON, err := DownloadWithExtensionsJSON(requestJSON) + if err != nil { + t.Fatalf("DownloadWithExtensionsJSON: %v", err) + } + var response DownloadResponse + if err := json.Unmarshal([]byte(responseJSON), &response); err != nil { + t.Fatalf("decode response: %v", err) + } + if response.ErrorType != "verification_required" || response.Service != extensionID { + t.Fatalf("response = %#v", response) + } + if got := GetItemProgress(itemID); got != "{}" { + t.Fatalf("verification response left stale progress: %s", got) + } +} + func TestSignedSessionURL(t *testing.T) { base := SignedSessionConfig{BaseURL: "https://auth.example.com/api"} diff --git a/go_backend/progress.go b/go_backend/progress.go index b697828c..c6a0cc43 100644 --- a/go_backend/progress.go +++ b/go_backend/progress.go @@ -15,6 +15,7 @@ type DownloadProgress struct { BytesReceived int64 `json:"bytes_received"` IsDownloading bool `json:"is_downloading"` Status string `json:"status"` + Stage string `json:"stage,omitempty"` } type ItemProgress struct { @@ -25,6 +26,7 @@ type ItemProgress struct { SpeedMBps float64 `json:"speed_mbps"` IsDownloading bool `json:"is_downloading"` Status string `json:"status"` + Stage string `json:"stage,omitempty"` revision int64 } @@ -53,6 +55,7 @@ type progressBridgeState struct { speedDeciMBps int64 downloading bool status string + stage string } var ( @@ -94,6 +97,7 @@ func itemProgressBridgeState(item *ItemProgress) progressBridgeState { speedDeciMBps: int64(math.Round(speed * 10)), downloading: item.IsDownloading, status: item.Status, + stage: item.Stage, } } @@ -116,6 +120,7 @@ func getProgress() DownloadProgress { BytesReceived: item.BytesReceived, IsDownloading: item.IsDownloading, Status: item.Status, + Stage: item.Stage, } } @@ -211,6 +216,10 @@ func StartItemProgress(itemID string) { } func SetItemPreparing(itemID string) { + SetItemPreparingStage(itemID, "") +} + +func SetItemPreparingStage(itemID, stage string) { multiMu.Lock() defer multiMu.Unlock() @@ -222,6 +231,7 @@ func SetItemPreparing(itemID string) { item.SpeedMBps = 0 item.IsDownloading = true item.Status = itemProgressStatusPreparing + item.Stage = stage markMultiProgressDirtyIfChangedLocked(item, before) } } @@ -234,6 +244,7 @@ func SetItemDownloading(itemID string) { before := itemProgressBridgeState(item) item.IsDownloading = true item.Status = itemProgressStatusDownloading + item.Stage = "" markMultiProgressDirtyIfChangedLocked(item, before) } } @@ -262,6 +273,7 @@ func SetItemBytesReceived(itemID string, received int64) { if received > 0 { item.IsDownloading = true item.Status = itemProgressStatusDownloading + item.Stage = "" } markMultiProgressDirtyIfChangedLocked(item, before) } @@ -281,6 +293,7 @@ func SetItemBytesReceivedWithSpeed(itemID string, received int64, speedMBps floa if received > 0 { item.IsDownloading = true item.Status = itemProgressStatusDownloading + item.Stage = "" } markMultiProgressDirtyIfChangedLocked(item, before) } @@ -295,6 +308,7 @@ func CompleteItemProgress(itemID string) { item.Progress = 1.0 item.IsDownloading = false item.Status = itemProgressStatusCompleted + item.Stage = "" markMultiProgressDirtyIfChangedLocked(item, before) } } @@ -320,6 +334,7 @@ func SetItemProgress(itemID string, progress float64, bytesReceived, bytesTotal if hasByteProgress || progress >= 1 || item.Status == itemProgressStatusDownloading { item.IsDownloading = true item.Status = itemProgressStatusDownloading + item.Stage = "" } markMultiProgressDirtyIfChangedLocked(item, before) } @@ -333,6 +348,7 @@ func SetItemFinalizing(itemID string) { before := itemProgressBridgeState(item) item.Progress = 1.0 item.Status = itemProgressStatusFinalizing + item.Stage = "" markMultiProgressDirtyIfChangedLocked(item, before) } } diff --git a/go_backend/progress_test.go b/go_backend/progress_test.go index 214509aa..38166d20 100644 --- a/go_backend/progress_test.go +++ b/go_backend/progress_test.go @@ -50,6 +50,32 @@ func TestItemProgressPreparingAndDownloadingStatuses(t *testing.T) { } } +func TestItemProgressPreparationStageIsObservable(t *testing.T) { + ClearAllItemProgress() + defer ClearAllItemProgress() + + itemID := "stage-item" + StartItemProgress(itemID) + SetItemPreparingStage(itemID, "resolving_metadata") + + var progress ItemProgress + if err := json.Unmarshal([]byte(GetItemProgress(itemID)), &progress); err != nil { + t.Fatalf("decode progress: %v", err) + } + if progress.Status != itemProgressStatusPreparing || progress.Stage != "resolving_metadata" { + t.Fatalf("unexpected preparation progress: %#v", progress) + } + + SetItemDownloading(itemID) + progress = ItemProgress{} + if err := json.Unmarshal([]byte(GetItemProgress(itemID)), &progress); err != nil { + t.Fatalf("decode downloading progress: %v", err) + } + if progress.Stage != "" { + t.Fatalf("download stage was not cleared: %#v", progress) + } +} + func TestItemProgressFinalizingAndCompletedStatuses(t *testing.T) { const itemID = "progress-finalizing-item" RemoveItemProgress(itemID)