mirror of
https://github.com/zarzet/SpotiFLAC-Mobile.git
synced 2026-07-28 23:08:59 +02:00
fix(download): accelerate verified extension resumes
This commit is contained in:
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
+273
-110
@@ -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(),
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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"}
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user