From bab954c6d99d256a2ec6d3ad8d7a84ebce9c3b9e Mon Sep 17 00:00:00 2001 From: zarzet Date: Wed, 15 Jul 2026 21:31:16 +0700 Subject: [PATCH] fix(extensions): harden package and runtime lifecycle --- go_backend/cancel.go | 24 ++ go_backend/exports_extensions.go | 3 + go_backend/extension_manager.go | 348 ++++++++++++------ .../extension_manager_supplement_test.go | 22 +- go_backend/extension_manifest.go | 27 ++ go_backend/extension_provider_wrapper.go | 10 +- go_backend/extension_providers.go | 38 ++ go_backend/extension_runtime.go | 132 +++---- go_backend/extension_runtime_ffmpeg.go | 36 +- go_backend/extension_runtime_storage.go | 339 ++++++----------- go_backend/extension_runtime_storage_test.go | 74 +++- .../extension_runtime_supplement_test.go | 7 +- go_backend/extension_test.go | 108 ++++++ go_backend/extension_timeout.go | 37 +- .../log_progress_timeout_supplement_test.go | 49 +++ 15 files changed, 815 insertions(+), 439 deletions(-) diff --git a/go_backend/cancel.go b/go_backend/cancel.go index e9a3a247..0c0cdcfb 100644 --- a/go_backend/cancel.go +++ b/go_backend/cancel.go @@ -59,6 +59,18 @@ func initDownloadCancel(itemID string) context.Context { return ctx } +func downloadCancelContext(itemID string) context.Context { + if itemID == "" { + return context.Background() + } + cancelMu.Lock() + defer cancelMu.Unlock() + if entry, ok := cancelMap[itemID]; ok && entry.ctx != nil { + return entry.ctx + } + return context.Background() +} + func cancelDownload(itemID string) { if itemID == "" { return @@ -153,6 +165,18 @@ func initExtensionRequestCancel(requestID string) context.Context { return ctx } +func extensionRequestCancelContext(requestID string) context.Context { + if requestID == "" { + return context.Background() + } + extensionRequestCancelMu.Lock() + defer extensionRequestCancelMu.Unlock() + if entry, ok := extensionRequestCancelMap[requestID]; ok && entry.ctx != nil { + return entry.ctx + } + return context.Background() +} + func cancelExtensionRequest(requestID string) { if requestID == "" { return diff --git a/go_backend/exports_extensions.go b/go_backend/exports_extensions.go index 328876ef..6af3c95a 100644 --- a/go_backend/exports_extensions.go +++ b/go_backend/exports_extensions.go @@ -1077,6 +1077,9 @@ func callExtensionFunctionJSONWithRequestID(extensionID, functionName string, ti result, err := RunWithTimeoutContextAndRecover(requestCtx, vm, script, timeout) perf.recordJS(time.Since(jsStartedAt)) if err != nil { + if IsRuntimeUnsafeError(err) { + quarantineRuntimeLocked(ext, vm) + } if isExtensionRequestCancelled(requestID) || errors.Is(err, ErrExtensionRequestCancelled) { return "", ErrExtensionRequestCancelled } diff --git a/go_backend/extension_manager.go b/go_backend/extension_manager.go index fbdff897..27886cc2 100644 --- a/go_backend/extension_manager.go +++ b/go_backend/extension_manager.go @@ -6,6 +6,7 @@ import ( "fmt" "io" "os" + "path" "path/filepath" "strconv" "strings" @@ -49,6 +50,75 @@ func isExtensionPackagePath(filePath string) bool { return strings.HasSuffix(lowerPath, ".spotiflac-ext") || strings.HasSuffix(lowerPath, ".sflx") } +func managedExtensionPath(root, extensionID string) (string, error) { + if root == "" { + return "", fmt.Errorf("extension directory is not configured") + } + if !extensionIDPattern.MatchString(extensionID) { + return "", fmt.Errorf("invalid extension ID %q", extensionID) + } + fullPath := filepath.Join(root, extensionID) + if !isPathWithinBase(root, fullPath) { + return "", fmt.Errorf("extension path escapes its managed directory") + } + return fullPath, nil +} + +func safeExtensionAssetPath(root, assetPath string) (string, bool) { + if root == "" || assetPath == "" || filepath.IsAbs(assetPath) || strings.Contains(assetPath, `\`) { + return "", false + } + cleaned := path.Clean(assetPath) + if cleaned == "." || cleaned == ".." || strings.HasPrefix(cleaned, "../") { + return "", false + } + fullPath := filepath.Join(root, filepath.FromSlash(cleaned)) + return fullPath, isPathWithinBase(root, fullPath) +} + +func extractExtensionArchive(zipReader *zip.ReadCloser, destination string) error { + for _, file := range zipReader.File { + if file.FileInfo().IsDir() { + continue + } + if file.FileInfo().Mode()&os.ModeSymlink != 0 || strings.Contains(file.Name, `\`) { + return fmt.Errorf("unsafe path in extension archive: %s", file.Name) + } + + relPath := path.Clean(file.Name) + if relPath == "." || relPath == ".." || strings.HasPrefix(relPath, "../") || path.IsAbs(relPath) { + return fmt.Errorf("unsafe path in extension archive: %s", file.Name) + } + destPath := filepath.Join(destination, filepath.FromSlash(relPath)) + if !isPathWithinBase(destination, destPath) { + return fmt.Errorf("unsafe path in extension archive: %s", file.Name) + } + + if err := os.MkdirAll(filepath.Dir(destPath), 0755); err != nil { + return fmt.Errorf("failed to create extension directory: %w", err) + } + destFile, err := os.OpenFile(destPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0600) + if err != nil { + return fmt.Errorf("failed to create extension file: %w", err) + } + srcFile, err := file.Open() + if err != nil { + destFile.Close() + return fmt.Errorf("failed to open file in archive: %w", err) + } + _, copyErr := io.Copy(destFile, srcFile) + closeSrcErr := srcFile.Close() + closeDestErr := destFile.Close() + if copyErr != nil { + return fmt.Errorf("failed to extract extension file: %w", copyErr) + } + if closeSrcErr != nil || closeDestErr != nil { + return fmt.Errorf("failed to close extracted extension file") + } + } + return nil +} + type loadedExtension struct { ID string `json:"id"` Manifest *ExtensionManifest `json:"manifest"` @@ -251,48 +321,33 @@ func (m *extensionManager) loadExtensionFromFileLocked(filePath string) (*loaded return nil, fmt.Errorf("extension '%s' was installed by another process", manifest.DisplayName) } - extDir := filepath.Join(m.extensionsDir, manifest.Name) - if err := os.MkdirAll(extDir, 0755); err != nil { - return nil, fmt.Errorf("failed to create extension directory: %w", err) + extDir, err := managedExtensionPath(m.extensionsDir, manifest.Name) + if err != nil { + return nil, err + } + if _, err := os.Lstat(extDir); err == nil { + return nil, fmt.Errorf("extension directory already exists for %q", manifest.Name) + } else if !os.IsNotExist(err) { + return nil, fmt.Errorf("failed to inspect extension directory: %w", err) + } + stagingDir, err := os.MkdirTemp(m.extensionsDir, "."+manifest.Name+"-install-*") + if err != nil { + return nil, fmt.Errorf("failed to create extension staging directory: %w", err) + } + stagingCommitted := false + defer func() { + if !stagingCommitted { + _ = os.RemoveAll(stagingDir) + } + }() + if err := extractExtensionArchive(zipReader, stagingDir); err != nil { + return nil, err } - for _, file := range zipReader.File { - if file.FileInfo().IsDir() { - continue - } - - relPath := filepath.Clean(file.Name) - if strings.HasPrefix(relPath, "..") || filepath.IsAbs(relPath) { - GoLog("[Extension] Skipping unsafe path in archive: %s\n", file.Name) - continue - } - destPath := filepath.Join(extDir, relPath) - - destDir := filepath.Dir(destPath) - if err := os.MkdirAll(destDir, 0755); err != nil { - return nil, fmt.Errorf("failed to create directory %s: %w", destDir, err) - } - - destFile, err := os.Create(destPath) - if err != nil { - return nil, fmt.Errorf("failed to create file %s: %w", destPath, err) - } - - srcFile, err := file.Open() - if err != nil { - destFile.Close() - return nil, fmt.Errorf("failed to open file in archive: %w", err) - } - - _, err = io.Copy(destFile, srcFile) - srcFile.Close() - destFile.Close() - if err != nil { - return nil, fmt.Errorf("failed to extract file: %w", err) - } + extDataDir, err := managedExtensionPath(m.dataDir, manifest.Name) + if err != nil { + return nil, err } - - extDataDir := filepath.Join(m.dataDir, manifest.Name) if err := os.MkdirAll(extDataDir, 0755); err != nil { return nil, fmt.Errorf("failed to create extension data directory: %w", err) } @@ -302,7 +357,7 @@ func (m *extensionManager) loadExtensionFromFileLocked(filePath string) (*loaded Manifest: manifest, Enabled: false, // New extensions start disabled DataDir: extDataDir, - SourceDir: extDir, + SourceDir: stagingDir, } if err := validateExtensionLoad(ext); err != nil { @@ -310,6 +365,11 @@ func (m *extensionManager) loadExtensionFromFileLocked(filePath string) (*loaded ext.Enabled = false GoLog("[Extension] Failed to validate extension %s: %v\n", manifest.Name, err) } + if err := os.Rename(stagingDir, extDir); err != nil { + return nil, fmt.Errorf("failed to activate extension: %w", err) + } + stagingCommitted = true + ext.SourceDir = extDir m.extensions[manifest.Name] = ext GoLog("[Extension] Loaded extension: %s v%s\n", manifest.DisplayName, manifest.Version) @@ -390,13 +450,12 @@ func newIsolatedExtensionRuntime(ext *loadedExtension) (*goja.Runtime, *extensio } runtime := &extensionRuntime{ - extensionID: ext.ID, - manifest: ext.Manifest, - settings: make(map[string]any), - cookieJar: nil, - dataDir: ext.DataDir, - vm: vm, - storageFlushDelay: defaultStorageFlushDelay, + extensionID: ext.ID, + manifest: ext.Manifest, + settings: make(map[string]any), + cookieJar: nil, + dataDir: ext.DataDir, + vm: vm, } if ext.runtime != nil && ext.runtime.cookieJar != nil { runtime.cookieJar = ext.runtime.cookieJar @@ -476,7 +535,7 @@ func acquireIsolatedExtensionRuntime(ext *loadedExtension) (*goja.Runtime, *exte // releaseIsolatedExtensionRuntime pools a healthy runtime for reuse or tears // it down. Pass healthy=false after an interrupt/timeout/script error, whose // VM state can't be trusted for reuse. -func releaseIsolatedExtensionRuntime(ext *loadedExtension, vm *goja.Runtime, runtime *extensionRuntime, healthy bool) { +func releaseIsolatedExtensionRuntime(ext *loadedExtension, vm *goja.Runtime, runtime *extensionRuntime, healthy, cleanupSafe bool) { if runtime != nil { if err := runtime.flushStorageNow(); err != nil { GoLog("[Extension:%s] isolated download storage flush failed: %v\n", ext.ID, err) @@ -493,14 +552,29 @@ func releaseIsolatedExtensionRuntime(ext *loadedExtension, vm *goja.Runtime, run ext.isolatedPoolMu.Unlock() } - if cleanupErr := runCleanupOnVM(vm); cleanupErr != nil { - GoLog("[Extension:%s] isolated download cleanup failed: %v\n", ext.ID, cleanupErr) + if cleanupSafe { + if cleanupErr := runCleanupOnVM(vm); cleanupErr != nil { + GoLog("[Extension:%s] isolated download cleanup failed: %v\n", ext.ID, cleanupErr) + } } if runtime != nil { runtime.closeStorageFlusher() } } +// quarantineRuntimeLocked detaches a VM that remained busy after interrupt. +// The caller holds VMMu. Touching or cleaning up that VM would race its stuck +// goroutine; a later call will build a fresh runtime from indexProgram. +func quarantineRuntimeLocked(ext *loadedExtension, vm *goja.Runtime) { + if ext == nil || ext.VM != vm { + return + } + ext.VM = nil + ext.runtime = nil + ext.initialized = false + ext.Error = "extension runtime was quarantined after an unresponsive script" +} + // drainIsolatedRuntimePool tears down idle isolated runtimes. Called on // extension teardown and on app-wide memory release. func drainIsolatedRuntimePool(ext *loadedExtension) { @@ -690,6 +764,16 @@ func validateExtensionLoad(ext *loadedExtension) error { return nil } +func teardownExtension(ext *loadedExtension) { + if ext == nil { + return + } + ext.Enabled = false + ext.VMMu.Lock() + teardownVMLocked(ext) + ext.VMMu.Unlock() +} + func (m *extensionManager) UnloadExtension(extensionID string) error { m.mu.Lock() defer m.mu.Unlock() @@ -699,6 +783,7 @@ func (m *extensionManager) UnloadExtension(extensionID string) error { return fmt.Errorf("extension not found") } + ext.Enabled = false ext.VMMu.Lock() teardownVMLocked(ext) ext.VMMu.Unlock() @@ -828,7 +913,14 @@ func (m *extensionManager) loadExtensionFromDirectory(dirPath string) (*loadedEx return existing, nil } - extDataDir := filepath.Join(m.dataDir, manifest.Name) + expectedSourceDir, err := managedExtensionPath(m.extensionsDir, manifest.Name) + if err != nil || filepath.Clean(dirPath) != filepath.Clean(expectedSourceDir) { + return nil, fmt.Errorf("extension directory name must match manifest name %q", manifest.Name) + } + extDataDir, err := managedExtensionPath(m.dataDir, manifest.Name) + if err != nil { + return nil, err + } if err := os.MkdirAll(extDataDir, 0755); err != nil { return nil, fmt.Errorf("failed to create extension data directory: %w", err) } @@ -870,22 +962,27 @@ func (m *extensionManager) RemoveExtension(extensionID string) error { return err } + sourceDir, err := managedExtensionPath(m.extensionsDir, ext.ID) + if err != nil || !isPathWithinBase(m.extensionsDir, ext.SourceDir) || filepath.Clean(ext.SourceDir) != filepath.Clean(sourceDir) { + return fmt.Errorf("refusing to remove extension outside the managed source directory") + } + dataDir, err := managedExtensionPath(m.dataDir, ext.ID) + if err != nil || !isPathWithinBase(m.dataDir, ext.DataDir) || filepath.Clean(ext.DataDir) != filepath.Clean(dataDir) { + return fmt.Errorf("refusing to remove extension outside the managed data directory") + } + if err := m.UnloadExtension(extensionID); err != nil { return err } - if ext.SourceDir != "" { - if err := os.RemoveAll(ext.SourceDir); err != nil { - GoLog("[Extension] Warning: failed to remove source dir: %v\n", err) - } + if err := os.RemoveAll(sourceDir); err != nil { + GoLog("[Extension] Warning: failed to remove source dir: %v\n", err) } // Uninstall means gone: storage.json and encrypted credentials must not // linger on disk after the extension is removed. - if ext.DataDir != "" { - if err := os.RemoveAll(ext.DataDir); err != nil { - GoLog("[Extension] Warning: failed to remove data dir: %v\n", err) - } + if err := os.RemoveAll(dataDir); err != nil { + GoLog("[Extension] Warning: failed to remove data dir: %v\n", err) } return nil @@ -960,56 +1057,28 @@ func (m *extensionManager) upgradeExtensionLocked(filePath string) (*loadedExten GoLog("[Extension] Upgrading %s from v%s to v%s\n", newManifest.DisplayName, existing.Manifest.Version, newManifest.Version) - extDataDir := existing.DataDir - extDir := existing.SourceDir + extDataDir, err := managedExtensionPath(m.dataDir, newManifest.Name) + if err != nil || filepath.Clean(existing.DataDir) != filepath.Clean(extDataDir) { + return nil, fmt.Errorf("installed extension has an invalid data directory") + } + extDir, err := managedExtensionPath(m.extensionsDir, newManifest.Name) + if err != nil || filepath.Clean(existing.SourceDir) != filepath.Clean(extDir) { + return nil, fmt.Errorf("installed extension has an invalid source directory") + } wasEnabled := existing.Enabled - m.UnloadExtension(existing.ID) - - if extDir != "" { - if err := os.RemoveAll(extDir); err != nil { - GoLog("[Extension] Warning: failed to remove old source dir: %v\n", err) - } + stagingDir, err := os.MkdirTemp(m.extensionsDir, "."+newManifest.Name+"-upgrade-*") + if err != nil { + return nil, fmt.Errorf("failed to create upgrade staging directory: %w", err) } - - if err := os.MkdirAll(extDir, 0755); err != nil { - return nil, fmt.Errorf("failed to create extension directory: %w", err) - } - - for _, file := range zipReader.File { - if file.FileInfo().IsDir() { - continue - } - - relPath := filepath.Clean(file.Name) - if strings.HasPrefix(relPath, "..") || filepath.IsAbs(relPath) { - GoLog("[Extension] Skipping unsafe path in archive: %s\n", file.Name) - continue - } - destPath := filepath.Join(extDir, relPath) - - destDir := filepath.Dir(destPath) - if err := os.MkdirAll(destDir, 0755); err != nil { - return nil, fmt.Errorf("failed to create directory %s: %w", destDir, err) - } - - destFile, err := os.Create(destPath) - if err != nil { - return nil, fmt.Errorf("failed to create file %s: %w", destPath, err) - } - - srcFile, err := file.Open() - if err != nil { - destFile.Close() - return nil, fmt.Errorf("failed to open file in archive: %w", err) - } - - _, err = io.Copy(destFile, srcFile) - srcFile.Close() - destFile.Close() - if err != nil { - return nil, fmt.Errorf("failed to extract file: %w", err) + stagingActive := true + defer func() { + if stagingActive { + _ = os.RemoveAll(stagingDir) } + }() + if err := extractExtensionArchive(zipReader, stagingDir); err != nil { + return nil, err } ext := &loadedExtension{ @@ -1017,22 +1086,53 @@ func (m *extensionManager) upgradeExtensionLocked(filePath string) (*loadedExten Manifest: newManifest, Enabled: wasEnabled, // Preserve enabled state from before upgrade DataDir: extDataDir, - SourceDir: extDir, + SourceDir: stagingDir, } if wasEnabled { if err := ext.ensureRuntimeReady(); err != nil { - GoLog("[Extension] Failed to initialize upgraded extension %s: %v\n", newManifest.Name, err) + return nil, fmt.Errorf("upgraded extension failed validation: %w", err) } } else if err := validateExtensionLoad(ext); err != nil { - ext.Error = err.Error() - ext.Enabled = false - GoLog("[Extension] Failed to validate upgraded extension %s: %v\n", newManifest.Name, err) + return nil, fmt.Errorf("upgraded extension failed validation: %w", err) + } + + backupDir, err := os.MkdirTemp(m.extensionsDir, "."+newManifest.Name+"-backup-*") + if err != nil { + teardownExtension(ext) + return nil, fmt.Errorf("failed to prepare upgrade backup: %w", err) + } + if err := os.Remove(backupDir); err != nil { + teardownExtension(ext) + return nil, fmt.Errorf("failed to prepare upgrade backup: %w", err) + } + if err := os.Rename(extDir, backupDir); err != nil { + teardownExtension(ext) + return nil, fmt.Errorf("failed to preserve current extension: %w", err) + } + if err := os.Rename(stagingDir, extDir); err != nil { + _ = os.Rename(backupDir, extDir) + teardownExtension(ext) + return nil, fmt.Errorf("failed to activate upgraded extension: %w", err) + } + stagingActive = false + ext.SourceDir = extDir + + existing.Enabled = false + if err := m.UnloadExtension(existing.ID); err != nil { + _ = os.RemoveAll(extDir) + _ = os.Rename(backupDir, extDir) + existing.Enabled = wasEnabled + teardownExtension(ext) + return nil, fmt.Errorf("failed to unload current extension: %w", err) } m.mu.Lock() m.extensions[newManifest.Name] = ext m.mu.Unlock() + if err := os.RemoveAll(backupDir); err != nil { + GoLog("[Extension] Warning: failed to remove upgrade backup: %v\n", err) + } GoLog("[Extension] Upgraded extension: %s to v%s\n", newManifest.DisplayName, newManifest.Version) @@ -1160,6 +1260,15 @@ func (m *extensionManager) GetInstalledExtensionsJSON() (string, error) { if ext.Manifest.Permissions.Storage { permissions = append(permissions, "storage:enabled") } + if ext.Manifest.Permissions.File { + permissions = append(permissions, "file:enabled") + } + if ext.Manifest.Permissions.AllowHTTP { + permissions = append(permissions, "network:http") + } + if ext.Manifest.HasCapability("rawFfmpeg") { + permissions = append(permissions, "ffmpeg:raw") + } status := "loaded" if ext.Error != "" { @@ -1170,15 +1279,19 @@ func (m *extensionManager) GetInstalledExtensionsJSON() (string, error) { iconPath := "" if ext.Manifest.Icon != "" && ext.SourceDir != "" { - possibleIcon := filepath.Join(ext.SourceDir, ext.Manifest.Icon) - if _, err := os.Stat(possibleIcon); err == nil { - iconPath = possibleIcon + possibleIcon, safe := safeExtensionAssetPath(ext.SourceDir, ext.Manifest.Icon) + if safe { + if _, err := os.Stat(possibleIcon); err == nil { + iconPath = possibleIcon + } } } if iconPath == "" && ext.SourceDir != "" { - possibleIcon := filepath.Join(ext.SourceDir, "icon.png") - if _, err := os.Stat(possibleIcon); err == nil { - iconPath = possibleIcon + possibleIcon, safe := safeExtensionAssetPath(ext.SourceDir, "icon.png") + if safe { + if _, err := os.Stat(possibleIcon); err == nil { + iconPath = possibleIcon + } } } @@ -1334,6 +1447,9 @@ func (m *extensionManager) InvokeAction(extensionID string, actionName string) ( result, err := RunWithTimeoutAndRecover(vm, script, DefaultJSTimeout) if err != nil { + if IsRuntimeUnsafeError(err) { + quarantineRuntimeLocked(ext, vm) + } GoLog("[Extension] InvokeAction error for %s.%s: %v\n", extensionID, actionName, err) return nil, fmt.Errorf("action failed: %v", err) } diff --git a/go_backend/extension_manager_supplement_test.go b/go_backend/extension_manager_supplement_test.go index 7d91c1f5..517cdaa8 100644 --- a/go_backend/extension_manager_supplement_test.go +++ b/go_backend/extension_manager_supplement_test.go @@ -34,9 +34,13 @@ registerExtension({ }); ` pkgV1 := filepath.Join(dir, "manager-ext-v1.spotiflac-ext") - createTestExtensionPackage(t, pkgV1, "manager-ext", "1.0.0", js, map[string]string{"../unsafe.txt": "skip"}) + createTestExtensionPackage(t, pkgV1, "manager-ext", "1.0.0", js, nil) pkgV2 := filepath.Join(dir, "manager-ext-v2.spotiflac-ext") createTestExtensionPackage(t, pkgV2, "manager-ext", "1.1.0", js, nil) + brokenUpgrade := filepath.Join(dir, "manager-ext-broken.spotiflac-ext") + createTestExtensionPackage(t, brokenUpgrade, "manager-ext", "1.0.1", `registerExtension({`, nil) + unsafePkg := filepath.Join(dir, "unsafe.spotiflac-ext") + createTestExtensionPackage(t, unsafePkg, "manager-ext", "1.0.0", js, map[string]string{"../unsafe.txt": "blocked"}) if compareVersions("v1.2.0", "1.1.9") <= 0 || compareVersions("1.0.0", "1.0") != 0 || compareVersions("1.0.0", "1.0.1") >= 0 { t.Fatal("compareVersions mismatch") @@ -47,6 +51,9 @@ registerExtension({ if _, err := manager.LoadExtensionFromFile(filepath.Join(dir, "missing.spotiflac-ext")); err == nil { t.Fatal("expected invalid package error") } + if _, err := manager.LoadExtensionFromFile(unsafePkg); err == nil { + t.Fatal("expected unsafe archive path to reject the package") + } ext, err := manager.LoadExtensionFromFile(pkgV1) if err != nil { @@ -55,9 +62,6 @@ registerExtension({ if ext.ID != "manager-ext" || ext.Enabled || ext.SourceDir == "" { t.Fatalf("loaded extension = %#v", ext) } - if _, err := os.Stat(filepath.Join(ext.SourceDir, "unsafe.txt")); err == nil { - t.Fatal("unsafe archive path should not be extracted") - } if _, err := manager.LoadExtensionFromFile(pkgV1); err == nil { t.Fatal("expected duplicate version error") } @@ -93,6 +97,16 @@ registerExtension({ if err := manager.SetExtensionEnabled("manager-ext", false); err != nil { t.Fatalf("disable extension: %v", err) } + if _, err := manager.UpgradeExtension(brokenUpgrade); err == nil { + t.Fatal("expected invalid upgrade to be rejected") + } + stillInstalled, err := manager.GetExtension("manager-ext") + if err != nil || stillInstalled.Manifest.Version != "1.0.0" { + t.Fatalf("failed upgrade replaced the installed version: %#v/%v", stillInstalled, err) + } + if _, err := os.Stat(filepath.Join(stillInstalled.SourceDir, "index.js")); err != nil { + t.Fatalf("failed upgrade removed the working source: %v", err) + } if ext.VM != nil || ext.initialized { t.Fatalf("expected VM teardown, got %#v", ext) } diff --git a/go_backend/extension_manifest.go b/go_backend/extension_manifest.go index 0448a452..d2aebed0 100644 --- a/go_backend/extension_manifest.go +++ b/go_backend/extension_manifest.go @@ -4,9 +4,12 @@ import ( "encoding/json" "fmt" "net/url" + "regexp" "strings" ) +var extensionIDPattern = regexp.MustCompile(`^[a-z0-9][a-z0-9._-]{0,127}$`) + type ExtensionType string const ( @@ -185,6 +188,12 @@ func (m *ExtensionManifest) Validate() error { if strings.TrimSpace(m.Name) == "" { return &ManifestValidationError{Field: "name", Message: "name is required"} } + if !extensionIDPattern.MatchString(m.Name) { + return &ManifestValidationError{ + Field: "name", + Message: "name must be a lowercase extension ID containing only letters, numbers, '.', '_' or '-'", + } + } if strings.TrimSpace(m.Version) == "" { return &ManifestValidationError{Field: "version", Message: "version is required"} @@ -260,6 +269,9 @@ func (m *ExtensionManifest) Validate() error { } if m.SignedSession != nil { + if !m.Permissions.Storage { + return &ManifestValidationError{Field: "permissions.storage", Message: "signedSession requires storage permission"} + } if strings.TrimSpace(m.SignedSession.Namespace) == "" { return &ManifestValidationError{Field: "signedSession.namespace", Message: "namespace is required"} } @@ -278,10 +290,25 @@ func (m *ExtensionManifest) Validate() error { return &ManifestValidationError{Field: "signedSession.baseUrl", Message: "baseUrl host must be listed in permissions.network"} } } + if m.HasCapability("rawFfmpeg") && !m.Permissions.File { + return &ManifestValidationError{Field: "permissions.file", Message: "rawFfmpeg capability requires file permission"} + } return nil } +func (m *ExtensionManifest) HasCapability(name string) bool { + if m == nil || m.Capabilities == nil { + return false + } + value, ok := m.Capabilities[name] + if !ok { + return false + } + enabled, ok := value.(bool) + return ok && enabled +} + func (m *ExtensionManifest) HasType(t ExtensionType) bool { for _, et := range m.Types { if et == t { diff --git a/go_backend/extension_provider_wrapper.go b/go_backend/extension_provider_wrapper.go index 321bcbed..09b64e5b 100644 --- a/go_backend/extension_provider_wrapper.go +++ b/go_backend/extension_provider_wrapper.go @@ -106,6 +106,9 @@ func callExtensionScript[T any](p *extensionProviderWrapper, opts extCallOpts, p perf.recordJS(time.Since(jsStartedAt)) perf.recordPayload(result) if err != nil { + if IsRuntimeUnsafeError(err) { + quarantineRuntimeLocked(p.extension, p.vm) + } if opts.requestID != "" && isExtensionRequestCancelled(opts.requestID) { return zero, ErrExtensionRequestCancelled } @@ -402,6 +405,9 @@ func (p *extensionProviderWrapper) EnrichTrackForItemID(track *ExtTrackMetadata, perf.recordJS(time.Since(jsStartedAt)) perf.recordPayload(result) if err != nil { + if IsRuntimeUnsafeError(err) { + quarantineRuntimeLocked(p.extension, p.vm) + } if isDownloadCancelled(itemID) { return track, ErrDownloadCancelled } @@ -516,8 +522,9 @@ func (p *extensionProviderWrapper) Download(trackID, quality, outputPath, itemID }, nil } vmHealthy := false + cleanupSafe := true defer func() { - releaseIsolatedExtensionRuntime(p.extension, vm, runtime, vmHealthy) + releaseIsolatedExtensionRuntime(p.extension, vm, runtime, vmHealthy, cleanupSafe) }() if runtime != nil { runtime.setActiveDownloadItemID(itemID) @@ -565,6 +572,7 @@ func (p *extensionProviderWrapper) Download(trackID, quality, outputPath, itemID perf.recordJS(time.Since(jsStartedAt)) perf.recordPayload(result) vmHealthy = err == nil + cleanupSafe = !IsRuntimeUnsafeError(err) if err != nil { errMsg := err.Error() errType := "script_error" diff --git a/go_backend/extension_providers.go b/go_backend/extension_providers.go index 85720f91..b41fd94d 100644 --- a/go_backend/extension_providers.go +++ b/go_backend/extension_providers.go @@ -306,6 +306,12 @@ func (m *extensionManager) runPostProcessingCommon(input PostProcessInput, metad GoLog("%s Hook %s failed: %v\n", logTag, hook.ID, err) continue } + if result.Success { + if err := validatePostProcessResult(provider.extension, currentInput, result); err != nil { + GoLog("%s Hook %s returned an unsafe result: %v\n", logTag, hook.ID, err) + continue + } + } if result.Success && result.NewFilePath != "" { currentInput.Path = result.NewFilePath @@ -322,6 +328,38 @@ func (m *extensionManager) runPostProcessingCommon(input PostProcessInput, metad return &PostProcessResult{Success: true, NewFilePath: currentInput.Path, NewFileURI: currentInput.URI}, nil } +func validatePostProcessResult(ext *loadedExtension, input PostProcessInput, result *PostProcessResult) error { + if ext == nil || ext.Manifest == nil || result == nil { + return fmt.Errorf("invalid post-processing result") + } + if result.NewFileURI != "" && result.NewFileURI != input.URI { + return fmt.Errorf("an extension cannot replace the destination URI") + } + if result.NewFilePath == "" || filepath.Clean(result.NewFilePath) == filepath.Clean(input.Path) { + return nil + } + if !ext.Manifest.Permissions.File { + return fmt.Errorf("file permission is required to replace the processed file") + } + if !filepath.IsAbs(result.NewFilePath) { + return fmt.Errorf("replacement file path must be absolute") + } + if input.Path != "" && isPathWithinBase(filepath.Dir(input.Path), result.NewFilePath) { + return nil + } + if ext.DataDir != "" && isPathWithinBase(ext.DataDir, result.NewFilePath) { + return nil + } + allowedDownloadDirsMu.RLock() + defer allowedDownloadDirsMu.RUnlock() + for _, dir := range allowedDownloadDirs { + if isPathWithinBase(dir, result.NewFilePath) { + return nil + } + } + return fmt.Errorf("replacement file path is outside allowed directories") +} + func (m *extensionManager) RunPostProcessing(filePath string, metadata map[string]any) (*PostProcessResult, error) { result, err := m.runPostProcessingCommon(PostProcessInput{Path: filePath}, metadata, false) if err != nil { diff --git a/go_backend/extension_runtime.go b/go_backend/extension_runtime.go index 05c996d0..d3f11dc1 100644 --- a/go_backend/extension_runtime.go +++ b/go_backend/extension_runtime.go @@ -130,15 +130,10 @@ type extensionRuntime struct { storageMu sync.RWMutex storageCache map[string]any - storageLoaded bool - storageDirty bool storageClosed bool - storageTimer *time.Timer - credentialsMu sync.RWMutex - credentialsCache map[string]any - credentialsLoaded bool - storageFlushDelay time.Duration + credentialsMu sync.RWMutex + credentialsCache map[string]any // Set when a signed-session call inside the current script invocation // required verification. The provider wrapper consumes it after the @@ -189,13 +184,12 @@ func newExtensionRuntime(ext *loadedExtension) *extensionRuntime { jar, _ := newSimpleCookieJar() runtime := &extensionRuntime{ - extensionID: ext.ID, - manifest: ext.Manifest, - settings: make(map[string]any), - cookieJar: jar, - dataDir: ext.DataDir, - vm: ext.VM, - storageFlushDelay: defaultStorageFlushDelay, + extensionID: ext.ID, + manifest: ext.Manifest, + settings: make(map[string]any), + cookieJar: jar, + dataDir: ext.DataDir, + vm: ext.VM, } runtime.httpClient = newExtensionHTTPClient(ext, jar, extensionHTTPTimeout(ext, 30*time.Second), true) @@ -299,10 +293,10 @@ func (r *extensionRuntime) bindDownloadCancelContext(req *http.Request) *http.Re if requestID == "" { return req } - return req.WithContext(initExtensionRequestCancel(requestID)) + return req.WithContext(extensionRequestCancelContext(requestID)) } - return req.WithContext(initDownloadCancel(itemID)) + return req.WithContext(downloadCancelContext(itemID)) } // downloadStallTimeout is how long a download may go without receiving a single @@ -540,59 +534,67 @@ func (r *extensionRuntime) RegisterAPIs(vm *goja.Runtime) { httpObj.Set("clearCookies", r.httpClearCookies) vm.Set("http", httpObj) - storageObj := vm.NewObject() - storageObj.Set("get", r.storageGet) - storageObj.Set("set", r.storageSet) - storageObj.Set("remove", r.storageRemove) - vm.Set("storage", storageObj) + if r.manifest != nil && r.manifest.Permissions.Storage { + storageObj := vm.NewObject() + storageObj.Set("get", r.storageGet) + storageObj.Set("set", r.storageSet) + storageObj.Set("remove", r.storageRemove) + vm.Set("storage", storageObj) - credentialsObj := vm.NewObject() - credentialsObj.Set("store", r.credentialsStore) - credentialsObj.Set("get", r.credentialsGet) - credentialsObj.Set("remove", r.credentialsRemove) - credentialsObj.Set("has", r.credentialsHas) - vm.Set("credentials", credentialsObj) - - authObj := vm.NewObject() - authObj.Set("openAuthUrl", r.authOpenUrl) - authObj.Set("getAuthCode", r.authGetCode) - authObj.Set("setAuthCode", r.authSetCode) - authObj.Set("clearAuth", r.authClear) - authObj.Set("isAuthenticated", r.authIsAuthenticated) - authObj.Set("getTokens", r.authGetTokens) - authObj.Set("generatePKCE", r.authGeneratePKCE) - authObj.Set("getPKCE", r.authGetPKCE) - authObj.Set("startOAuthWithPKCE", r.authStartOAuthWithPKCE) - authObj.Set("exchangeCodeWithPKCE", r.authExchangeCodeWithPKCE) - vm.Set("auth", authObj) - - if r.manifest != nil && r.manifest.SignedSession != nil { - sessionObj := vm.NewObject() - sessionObj.Set("signedFetch", r.signedSessionFetch) - sessionObj.Set("completeGrant", r.signedSessionCompleteGrant) - sessionObj.Set("status", r.signedSessionStatus) - sessionObj.Set("clear", r.signedSessionClear) - vm.Set("session", sessionObj) + credentialsObj := vm.NewObject() + credentialsObj.Set("store", r.credentialsStore) + credentialsObj.Set("get", r.credentialsGet) + credentialsObj.Set("remove", r.credentialsRemove) + credentialsObj.Set("has", r.credentialsHas) + vm.Set("credentials", credentialsObj) } - fileObj := vm.NewObject() - fileObj.Set("download", r.fileDownload) - fileObj.Set("exists", r.fileExists) - fileObj.Set("delete", r.fileDelete) - fileObj.Set("read", r.fileRead) - fileObj.Set("readBytes", r.fileReadBytes) - fileObj.Set("write", r.fileWrite) - fileObj.Set("writeBytes", r.fileWriteBytes) - fileObj.Set("copy", r.fileCopy) - fileObj.Set("move", r.fileMove) - fileObj.Set("getSize", r.fileGetSize) - vm.Set("file", fileObj) + if r.manifest != nil && r.manifest.Permissions.Storage { + authObj := vm.NewObject() + authObj.Set("openAuthUrl", r.authOpenUrl) + authObj.Set("getAuthCode", r.authGetCode) + authObj.Set("setAuthCode", r.authSetCode) + authObj.Set("clearAuth", r.authClear) + authObj.Set("isAuthenticated", r.authIsAuthenticated) + authObj.Set("getTokens", r.authGetTokens) + authObj.Set("generatePKCE", r.authGeneratePKCE) + authObj.Set("getPKCE", r.authGetPKCE) + authObj.Set("startOAuthWithPKCE", r.authStartOAuthWithPKCE) + authObj.Set("exchangeCodeWithPKCE", r.authExchangeCodeWithPKCE) + vm.Set("auth", authObj) - ffmpegObj := vm.NewObject() - ffmpegObj.Set("execute", r.ffmpegExecute) - ffmpegObj.Set("getInfo", r.ffmpegGetInfo) - ffmpegObj.Set("convert", r.ffmpegConvert) - vm.Set("ffmpeg", ffmpegObj) + if r.manifest.SignedSession != nil { + sessionObj := vm.NewObject() + sessionObj.Set("signedFetch", r.signedSessionFetch) + sessionObj.Set("completeGrant", r.signedSessionCompleteGrant) + sessionObj.Set("status", r.signedSessionStatus) + sessionObj.Set("clear", r.signedSessionClear) + vm.Set("session", sessionObj) + } + } + + if r.manifest != nil && r.manifest.Permissions.File { + fileObj := vm.NewObject() + fileObj.Set("download", r.fileDownload) + fileObj.Set("exists", r.fileExists) + fileObj.Set("delete", r.fileDelete) + fileObj.Set("read", r.fileRead) + fileObj.Set("readBytes", r.fileReadBytes) + fileObj.Set("write", r.fileWrite) + fileObj.Set("writeBytes", r.fileWriteBytes) + fileObj.Set("copy", r.fileCopy) + fileObj.Set("move", r.fileMove) + fileObj.Set("getSize", r.fileGetSize) + vm.Set("file", fileObj) + + ffmpegObj := vm.NewObject() + if r.manifest.HasCapability("rawFfmpeg") { + ffmpegObj.Set("execute", r.ffmpegExecute) + } + ffmpegObj.Set("getInfo", r.ffmpegGetInfo) + ffmpegObj.Set("convert", r.ffmpegConvert) + vm.Set("ffmpeg", ffmpegObj) + } matchingObj := vm.NewObject() matchingObj.Set("compareStrings", r.matchingCompareStrings) diff --git a/go_backend/extension_runtime_ffmpeg.go b/go_backend/extension_runtime_ffmpeg.go index b47cdffd..21d225ef 100644 --- a/go_backend/extension_runtime_ffmpeg.go +++ b/go_backend/extension_runtime_ffmpeg.go @@ -51,11 +51,17 @@ func ClearFFmpegCommand(commandID string) { } func (r *extensionRuntime) ffmpegExecute(call goja.FunctionCall) goja.Value { + if r.manifest == nil || !r.manifest.Permissions.File || !r.manifest.HasCapability("rawFfmpeg") { + return r.jsError("raw FFmpeg execution permission denied") + } if len(call.Arguments) < 1 { return r.jsError("command is required") } - command := call.Arguments[0].String() + return r.executeFFmpegCommand(call.Arguments[0].String(), "", "") +} + +func (r *extensionRuntime) executeFFmpegCommand(command, inputPath, outputPath string) goja.Value { ffmpegCommandsMu.Lock() ffmpegCommandID++ @@ -63,6 +69,8 @@ func (r *extensionRuntime) ffmpegExecute(call goja.FunctionCall) goja.Value { ffmpegCommands[cmdID] = &FFmpegCommand{ ExtensionID: r.extensionID, Command: command, + InputPath: inputPath, + OutputPath: outputPath, Completed: false, } ffmpegCommandsMu.Unlock() @@ -102,11 +110,17 @@ func (r *extensionRuntime) ffmpegExecute(call goja.FunctionCall) goja.Value { } func (r *extensionRuntime) ffmpegGetInfo(call goja.FunctionCall) goja.Value { + if r.manifest == nil || !r.manifest.Permissions.File { + return r.jsError("file permission denied") + } if len(call.Arguments) < 1 { return r.jsError("file path is required") } - filePath := call.Arguments[0].String() + filePath, err := r.validatePath(call.Arguments[0].String()) + if err != nil { + return r.jsError("%s", err.Error()) + } quality, err := GetAudioQuality(filePath) if err != nil { @@ -123,12 +137,21 @@ func (r *extensionRuntime) ffmpegGetInfo(call goja.FunctionCall) goja.Value { } func (r *extensionRuntime) ffmpegConvert(call goja.FunctionCall) goja.Value { + if r.manifest == nil || !r.manifest.Permissions.File { + return r.jsError("file permission denied") + } if len(call.Arguments) < 2 { return r.jsError("input and output paths are required") } - inputPath := call.Arguments[0].String() - outputPath := call.Arguments[1].String() + inputPath, err := r.validatePath(call.Arguments[0].String()) + if err != nil { + return r.jsError("invalid input path: %v", err) + } + outputPath, err := r.validatePath(call.Arguments[1].String()) + if err != nil { + return r.jsError("invalid output path: %v", err) + } options := map[string]any{} if len(call.Arguments) > 2 && !goja.IsUndefined(call.Arguments[2]) && !goja.IsNull(call.Arguments[2]) { @@ -160,8 +183,5 @@ func (r *extensionRuntime) ffmpegConvert(call goja.FunctionCall) goja.Value { command := strings.Join(cmdParts, " ") - execCall := goja.FunctionCall{ - Arguments: []goja.Value{r.vm.ToValue(command)}, - } - return r.ffmpegExecute(execCall) + return r.executeFFmpegCommand(command, inputPath, outputPath) } diff --git a/go_backend/extension_runtime_storage.go b/go_backend/extension_runtime_storage.go index 70d00b00..390166e3 100644 --- a/go_backend/extension_runtime_storage.go +++ b/go_backend/extension_runtime_storage.go @@ -10,9 +10,7 @@ import ( "io" "os" "path/filepath" - "reflect" "sync" - "time" "github.com/dop251/goja" ) @@ -37,150 +35,81 @@ func writeExtensionFileLocked(path string, data []byte) error { return os.Rename(tmp, path) } -const ( - defaultStorageFlushDelay = 400 * time.Millisecond - storageFlushRetryDelay = 2 * time.Second -) - func (r *extensionRuntime) getStoragePath() string { return filepath.Join(r.dataDir, "storage.json") } -func cloneInterfaceMap(src map[string]any) map[string]any { - if len(src) == 0 { - return make(map[string]any) - } - dst := make(map[string]any, len(src)) - for k, v := range src { - dst[k] = v - } - return dst -} - -func (r *extensionRuntime) ensureStorageLoaded() error { - r.storageMu.RLock() - if r.storageLoaded { - r.storageMu.RUnlock() - return nil - } - r.storageMu.RUnlock() - - r.storageMu.Lock() - defer r.storageMu.Unlock() - if r.storageLoaded { - return nil - } - - storagePath := r.getStoragePath() - fileMu := extensionFileMu(storagePath) - fileMu.Lock() - data, err := os.ReadFile(storagePath) - fileMu.Unlock() +func readJSONMapFile(path string) (map[string]any, error) { + data, err := os.ReadFile(path) if err != nil { if os.IsNotExist(err) { - r.storageCache = make(map[string]any) - r.storageLoaded = true - return nil + return make(map[string]any), nil } + return nil, err + } + result := make(map[string]any) + if err := json.Unmarshal(data, &result); err != nil { + return nil, err + } + if result == nil { + result = make(map[string]any) + } + return result, nil +} + +func (r *extensionRuntime) refreshStorage() error { + path := r.getStoragePath() + fileMu := extensionFileMu(path) + fileMu.Lock() + snapshot, err := readJSONMapFile(path) + fileMu.Unlock() + if err != nil { return err } - - var storage map[string]any - if err := json.Unmarshal(data, &storage); err != nil { - return err - } - if storage == nil { - storage = make(map[string]any) - } - - r.storageCache = storage - r.storageLoaded = true + r.storageMu.Lock() + r.storageCache = snapshot + r.storageMu.Unlock() return nil } -func (r *extensionRuntime) queueStorageFlushLocked(delay time.Duration) { - if r.storageClosed { - return +func (r *extensionRuntime) mutateStorage(mutate func(map[string]any) bool) error { + r.storageMu.RLock() + closed := r.storageClosed + r.storageMu.RUnlock() + if closed { + return fmt.Errorf("storage is closed") } - if r.storageTimer != nil { - return - } - r.storageTimer = time.AfterFunc(delay, r.flushStorageDirtyAsync) -} -func (r *extensionRuntime) persistStorageSnapshot(storage map[string]any) error { - data, err := json.Marshal(storage) + path := r.getStoragePath() + fileMu := extensionFileMu(path) + fileMu.Lock() + snapshot, err := readJSONMapFile(path) + if err == nil && mutate(snapshot) { + var data []byte + data, err = json.Marshal(snapshot) + if err == nil { + err = writeExtensionFileLocked(path, data) + } + } + fileMu.Unlock() if err != nil { return err } - path := r.getStoragePath() - mu := extensionFileMu(path) - mu.Lock() - defer mu.Unlock() - - return writeExtensionFileLocked(path, data) -} - -func (r *extensionRuntime) flushStorageDirtyAsync() { - if err := r.flushStorageDirty(); err != nil { - GoLog("[Extension:%s] Storage flush error: %v\n", r.extensionID, err) - } -} - -func (r *extensionRuntime) flushStorageDirty() error { r.storageMu.Lock() - if r.storageClosed { - r.storageTimer = nil - r.storageMu.Unlock() - return nil - } - if !r.storageDirty { - r.storageTimer = nil - r.storageMu.Unlock() - return nil - } - snapshot := cloneInterfaceMap(r.storageCache) - r.storageDirty = false - r.storageTimer = nil + r.storageCache = snapshot r.storageMu.Unlock() - - if err := r.persistStorageSnapshot(snapshot); err != nil { - r.storageMu.Lock() - r.storageDirty = true - r.queueStorageFlushLocked(storageFlushRetryDelay) - r.storageMu.Unlock() - return err - } - return nil } func (r *extensionRuntime) flushStorageNow() error { - r.storageMu.Lock() - if r.storageTimer != nil { - r.storageTimer.Stop() - r.storageTimer = nil - } - if !r.storageLoaded || r.storageClosed { - r.storageMu.Unlock() - return nil - } - snapshot := cloneInterfaceMap(r.storageCache) - r.storageDirty = false - r.storageMu.Unlock() - - return r.persistStorageSnapshot(snapshot) + // Mutations are persisted synchronously under the process-wide file lock. + return nil } func (r *extensionRuntime) closeStorageFlusher() { r.storageMu.Lock() r.storageClosed = true - r.storageDirty = false - if r.storageTimer != nil { - r.storageTimer.Stop() - r.storageTimer = nil - } r.storageMu.Unlock() } @@ -191,7 +120,7 @@ func (r *extensionRuntime) storageGet(call goja.FunctionCall) goja.Value { key := call.Arguments[0].String() - if err := r.ensureStorageLoaded(); err != nil { + if err := r.refreshStorage(); err != nil { GoLog("[Extension:%s] Storage load error: %v\n", r.extensionID, err) return goja.Undefined() } @@ -217,27 +146,14 @@ func (r *extensionRuntime) storageSet(call goja.FunctionCall) goja.Value { key := call.Arguments[0].String() value := call.Arguments[1].Export() - if err := r.ensureStorageLoaded(); err != nil { - GoLog("[Extension:%s] Storage load error: %v\n", r.extensionID, err) + if err := r.mutateStorage(func(storage map[string]any) bool { + storage[key] = value + return true + }); err != nil { + GoLog("[Extension:%s] Storage save error: %v\n", r.extensionID, err) return r.vm.ToValue(false) } - r.storageMu.Lock() - if r.storageClosed { - r.storageMu.Unlock() - return r.vm.ToValue(false) - } - if existing, exists := r.storageCache[key]; exists { - if reflect.DeepEqual(existing, value) { - r.storageMu.Unlock() - return r.vm.ToValue(true) - } - } - r.storageCache[key] = value - r.storageDirty = true - r.queueStorageFlushLocked(r.storageFlushDelay) - r.storageMu.Unlock() - return r.vm.ToValue(true) } @@ -248,25 +164,17 @@ func (r *extensionRuntime) storageRemove(call goja.FunctionCall) goja.Value { key := call.Arguments[0].String() - if err := r.ensureStorageLoaded(); err != nil { - GoLog("[Extension:%s] Storage load error: %v\n", r.extensionID, err) + if err := r.mutateStorage(func(storage map[string]any) bool { + if _, exists := storage[key]; !exists { + return false + } + delete(storage, key) + return true + }); err != nil { + GoLog("[Extension:%s] Storage save error: %v\n", r.extensionID, err) return r.vm.ToValue(false) } - r.storageMu.Lock() - if r.storageClosed { - r.storageMu.Unlock() - return r.vm.ToValue(false) - } - if _, exists := r.storageCache[key]; !exists { - r.storageMu.Unlock() - return r.vm.ToValue(true) - } - delete(r.storageCache, key) - r.storageDirty = true - r.queueStorageFlushLocked(r.storageFlushDelay) - r.storageMu.Unlock() - return r.vm.ToValue(true) } @@ -315,80 +223,73 @@ func (r *extensionRuntime) getEncryptionKey() ([]byte, error) { return hash[:], nil } -func (r *extensionRuntime) ensureCredentialsLoaded() error { - r.credentialsMu.RLock() - if r.credentialsLoaded { - r.credentialsMu.RUnlock() - return nil - } - r.credentialsMu.RUnlock() - - r.credentialsMu.Lock() - defer r.credentialsMu.Unlock() - if r.credentialsLoaded { - return nil - } - - credPath := r.getCredentialsPath() - data, err := os.ReadFile(credPath) +func (r *extensionRuntime) readCredentialsFileLocked() (map[string]any, error) { + data, err := os.ReadFile(r.getCredentialsPath()) if err != nil { if os.IsNotExist(err) { - r.credentialsCache = make(map[string]any) - r.credentialsLoaded = true - return nil + return make(map[string]any), nil } - return err + return nil, err } - key, err := r.getEncryptionKey() if err != nil { - return fmt.Errorf("failed to get encryption key: %w", err) + return nil, fmt.Errorf("failed to get encryption key: %w", err) } decrypted, err := decryptAES(data, key) if err != nil { - return fmt.Errorf("failed to decrypt credentials: %w", err) + return nil, fmt.Errorf("failed to decrypt credentials: %w", err) } - - var creds map[string]any + creds := make(map[string]any) if err := json.Unmarshal(decrypted, &creds); err != nil { - return err + return nil, err } if creds == nil { creds = make(map[string]any) } + return creds, nil +} - r.credentialsCache = creds - r.credentialsLoaded = true +func (r *extensionRuntime) refreshCredentials() error { + path := r.getCredentialsPath() + fileMu := extensionFileMu(path) + fileMu.Lock() + snapshot, err := r.readCredentialsFileLocked() + fileMu.Unlock() + if err != nil { + return err + } + r.credentialsMu.Lock() + r.credentialsCache = snapshot + r.credentialsMu.Unlock() return nil } -func (r *extensionRuntime) saveCredentials(creds map[string]any) error { - data, err := json.Marshal(creds) +func (r *extensionRuntime) mutateCredentials(mutate func(map[string]any)) error { + path := r.getCredentialsPath() + fileMu := extensionFileMu(path) + fileMu.Lock() + snapshot, err := r.readCredentialsFileLocked() + if err == nil { + mutate(snapshot) + var data []byte + data, err = json.Marshal(snapshot) + if err == nil { + var key []byte + key, err = r.getEncryptionKey() + if err == nil { + data, err = encryptAES(data, key) + } + } + if err == nil { + err = writeExtensionFileLocked(path, data) + } + } + fileMu.Unlock() if err != nil { return err } - - key, err := r.getEncryptionKey() - if err != nil { - return fmt.Errorf("failed to get encryption key: %w", err) - } - encrypted, err := encryptAES(data, key) - if err != nil { - return fmt.Errorf("failed to encrypt credentials: %w", err) - } - - credPath := r.getCredentialsPath() - credMu := extensionFileMu(credPath) - credMu.Lock() - err = writeExtensionFileLocked(credPath, encrypted) - credMu.Unlock() - if err != nil { - return err - } - r.credentialsMu.Lock() - r.credentialsCache = cloneInterfaceMap(creds) - r.credentialsLoaded = true + r.credentialsCache = snapshot r.credentialsMu.Unlock() return nil } @@ -401,17 +302,9 @@ func (r *extensionRuntime) credentialsStore(call goja.FunctionCall) goja.Value { key := call.Arguments[0].String() value := call.Arguments[1].Export() - if err := r.ensureCredentialsLoaded(); err != nil { - GoLog("[Extension:%s] Credentials load error: %v\n", r.extensionID, err) - return r.jsError("%s", err.Error()) - } - - r.credentialsMu.RLock() - nextCreds := cloneInterfaceMap(r.credentialsCache) - r.credentialsMu.RUnlock() - nextCreds[key] = value - - if err := r.saveCredentials(nextCreds); err != nil { + if err := r.mutateCredentials(func(credentials map[string]any) { + credentials[key] = value + }); err != nil { GoLog("[Extension:%s] Credentials save error: %v\n", r.extensionID, err) return r.jsError("%s", err.Error()) } @@ -426,7 +319,7 @@ func (r *extensionRuntime) credentialsGet(call goja.FunctionCall) goja.Value { key := call.Arguments[0].String() - if err := r.ensureCredentialsLoaded(); err != nil { + if err := r.refreshCredentials(); err != nil { GoLog("[Extension:%s] Credentials load error: %v\n", r.extensionID, err) return goja.Undefined() } @@ -451,17 +344,9 @@ func (r *extensionRuntime) credentialsRemove(call goja.FunctionCall) goja.Value key := call.Arguments[0].String() - if err := r.ensureCredentialsLoaded(); err != nil { - GoLog("[Extension:%s] Credentials load error: %v\n", r.extensionID, err) - return r.vm.ToValue(false) - } - - r.credentialsMu.RLock() - nextCreds := cloneInterfaceMap(r.credentialsCache) - r.credentialsMu.RUnlock() - delete(nextCreds, key) - - if err := r.saveCredentials(nextCreds); err != nil { + if err := r.mutateCredentials(func(credentials map[string]any) { + delete(credentials, key) + }); err != nil { GoLog("[Extension:%s] Credentials save error: %v\n", r.extensionID, err) return r.vm.ToValue(false) } @@ -476,7 +361,7 @@ func (r *extensionRuntime) credentialsHas(call goja.FunctionCall) goja.Value { key := call.Arguments[0].String() - if err := r.ensureCredentialsLoaded(); err != nil { + if err := r.refreshCredentials(); err != nil { return r.vm.ToValue(false) } diff --git a/go_backend/extension_runtime_storage_test.go b/go_backend/extension_runtime_storage_test.go index df33d41d..dfa61164 100644 --- a/go_backend/extension_runtime_storage_test.go +++ b/go_backend/extension_runtime_storage_test.go @@ -24,6 +24,76 @@ func setStorageValue(t *testing.T, runtime *extensionRuntime, key string, value } } +func TestExtensionRuntimeStorageConcurrentRuntimesMergeWrites(t *testing.T) { + dataDir := t.TempDir() + ext := &loadedExtension{ID: "merge-test", Manifest: &ExtensionManifest{Name: "merge-test"}, DataDir: dataDir} + runtimeA := newExtensionRuntime(ext) + runtimeB := newExtensionRuntime(ext) + runtimeA.RegisterAPIs(goja.New()) + runtimeB.RegisterAPIs(goja.New()) + + start := make(chan struct{}) + done := make(chan bool, 2) + go func() { + <-start + result := runtimeA.storageSet(goja.FunctionCall{Arguments: []goja.Value{ + runtimeA.vm.ToValue("from_a"), runtimeA.vm.ToValue("a"), + }}) + done <- result.ToBoolean() + }() + go func() { + <-start + result := runtimeB.storageSet(goja.FunctionCall{Arguments: []goja.Value{ + runtimeB.vm.ToValue("from_b"), runtimeB.vm.ToValue("b"), + }}) + done <- result.ToBoolean() + }() + close(start) + if !<-done || !<-done { + t.Fatal("concurrent storage write failed") + } + + storage := readStorageMap(t, filepath.Join(dataDir, "storage.json")) + if storage["from_a"] != "a" || storage["from_b"] != "b" { + t.Fatalf("concurrent storage writes were not merged: %#v", storage) + } + + credStart := make(chan struct{}) + credDone := make(chan struct{}, 2) + for _, item := range []struct { + runtime *extensionRuntime + key string + }{ + {runtimeA, "token_a"}, + {runtimeB, "token_b"}, + } { + item := item + go func() { + <-credStart + result := item.runtime.credentialsStore(goja.FunctionCall{Arguments: []goja.Value{ + item.runtime.vm.ToValue(item.key), + item.runtime.vm.ToValue(item.key + "_value"), + }}) + if success, _ := result.Export().(map[string]any)["success"].(bool); !success { + t.Errorf("credentialsStore(%s) failed", item.key) + } + credDone <- struct{}{} + }() + } + close(credStart) + <-credDone + <-credDone + + reader := newExtensionRuntime(ext) + reader.RegisterAPIs(goja.New()) + for _, key := range []string{"token_a", "token_b"} { + got := reader.credentialsGet(goja.FunctionCall{Arguments: []goja.Value{reader.vm.ToValue(key)}}).String() + if got != key+"_value" { + t.Fatalf("credential %s = %q", key, got) + } + } +} + func readStorageMap(t *testing.T, storagePath string) map[string]any { t.Helper() data, err := os.ReadFile(storagePath) @@ -38,7 +108,7 @@ func readStorageMap(t *testing.T, storagePath string) map[string]any { return parsed } -func TestExtensionRuntimeStorage_DebouncedWriteCompactJSON(t *testing.T) { +func TestExtensionRuntimeStorage_AtomicWriteCompactJSON(t *testing.T) { ext := &loadedExtension{ ID: "storage-test", Manifest: &ExtensionManifest{ @@ -48,7 +118,6 @@ func TestExtensionRuntimeStorage_DebouncedWriteCompactJSON(t *testing.T) { } runtime := newExtensionRuntime(ext) - runtime.storageFlushDelay = 25 * time.Millisecond runtime.RegisterAPIs(goja.New()) setStorageValue(t, runtime, "k1", "v1") @@ -96,7 +165,6 @@ func TestUnloadExtension_FlushesPendingStorage(t *testing.T) { } runtime := newExtensionRuntime(ext) - runtime.storageFlushDelay = time.Hour runtime.RegisterAPIs(ext.VM) ext.runtime = runtime diff --git a/go_backend/extension_runtime_supplement_test.go b/go_backend/extension_runtime_supplement_test.go index d35192d2..61faf852 100644 --- a/go_backend/extension_runtime_supplement_test.go +++ b/go_backend/extension_runtime_supplement_test.go @@ -618,10 +618,9 @@ func TestExtensionStoreSettingsAndRuntimeStorage(t *testing.T) { vm := goja.New() runtime := &extensionRuntime{ - extensionID: "storage-ext", - dataDir: filepath.Join(dir, "runtime"), - vm: vm, - storageFlushDelay: time.Hour, + extensionID: "storage-ext", + dataDir: filepath.Join(dir, "runtime"), + vm: vm, } if err := os.MkdirAll(runtime.dataDir, 0755); err != nil { t.Fatal(err) diff --git a/go_backend/extension_test.go b/go_backend/extension_test.go index 46b897e8..4e5d0f41 100644 --- a/go_backend/extension_test.go +++ b/go_backend/extension_test.go @@ -5,6 +5,7 @@ import ( "errors" "net/http" "path/filepath" + "strconv" "testing" "time" @@ -76,6 +77,43 @@ func TestParseManifest_MissingName(t *testing.T) { } } +func TestParseManifestRejectsUnsafeExtensionIDs(t *testing.T) { + for _, name := range []string{"../escape", `..\\escape`, "/absolute", "UpperCase", "."} { + manifest := `{"name":` + strconv.Quote(name) + `,"version":"1.0.0","description":"test","type":["metadata_provider"]}` + if _, err := ParseManifest([]byte(manifest)); err == nil { + t.Fatalf("expected unsafe extension ID %q to be rejected", name) + } + } +} + +func TestManifestPrivilegedCapabilitiesRequirePermissions(t *testing.T) { + rawFFmpeg := &ExtensionManifest{ + Name: "raw-ffmpeg", + Version: "1.0.0", + Description: "test", + Types: []ExtensionType{ExtensionTypeDownloadProvider}, + Capabilities: map[string]any{"rawFfmpeg": true}, + } + if err := rawFFmpeg.Validate(); err == nil { + t.Fatal("expected rawFfmpeg without file permission to be rejected") + } + + signedSession := &ExtensionManifest{ + Name: "signed-session", + Version: "1.0.0", + Description: "test", + Types: []ExtensionType{ExtensionTypeDownloadProvider}, + Permissions: ExtensionPermissions{Network: []string{"api.example.com"}}, + SignedSession: &SignedSessionConfig{ + Namespace: "test", + BaseURL: "https://api.example.com", + }, + } + if err := signedSession.Validate(); err == nil { + t.Fatal("expected signedSession without storage permission to be rejected") + } +} + func TestParseManifest_MissingType(t *testing.T) { invalidManifest := `{ "name": "test-provider", @@ -329,6 +367,7 @@ func TestExtensionRuntime_BindDownloadCancelContext(t *testing.T) { runtime := newExtensionRuntime(ext) runtime.setActiveDownloadItemID("test-item") + initDownloadCancel("test-item") t.Cleanup(func() { clearDownloadCancel("test-item") runtime.clearActiveDownloadItemID() @@ -340,6 +379,12 @@ func TestExtensionRuntime_BindDownloadCancelContext(t *testing.T) { } req = runtime.bindDownloadCancelContext(req) + cancelMu.Lock() + refs := cancelMap["test-item"].refs + cancelMu.Unlock() + if refs != 1 { + t.Fatalf("binding a request leaked a cancellation reference: %d", refs) + } cancelDownload("test-item") select { @@ -365,6 +410,7 @@ func TestExtensionRuntime_BindDownloadCancelContextPreservesPreCancelledState(t runtime := newExtensionRuntime(ext) runtime.setActiveDownloadItemID("test-item") cancelDownload("test-item") + initDownloadCancel("test-item") t.Cleanup(func() { clearDownloadCancel("test-item") runtime.clearActiveDownloadItemID() @@ -415,12 +461,19 @@ func TestExtensionRuntime_BindExtensionRequestCancelContext(t *testing.T) { runtime.setActiveRequestID(requestID) defer runtime.clearActiveRequestID() + initExtensionRequestCancel(requestID) req, err := http.NewRequest(http.MethodGet, "https://example.com", nil) if err != nil { t.Fatalf("new request: %v", err) } req = runtime.bindDownloadCancelContext(req) + extensionRequestCancelMu.Lock() + refs := extensionRequestCancelMap[requestID].refs + extensionRequestCancelMu.Unlock() + if refs != 1 { + t.Fatalf("binding a request leaked a cancellation reference: %d", refs) + } cancelExtensionRequest(requestID) select { @@ -466,6 +519,61 @@ func TestExtensionRuntime_SSRFProtection(t *testing.T) { } } +func TestExtensionRuntimeAPIsRequireDeclaredPermissions(t *testing.T) { + withoutPermissions := &loadedExtension{ + ID: "no-permissions", + Manifest: &ExtensionManifest{Name: "no-permissions"}, + DataDir: t.TempDir(), + } + vm := goja.New() + newExtensionRuntime(withoutPermissions).RegisterAPIs(vm) + for _, api := range []string{"storage", "credentials", "auth", "session", "file", "ffmpeg"} { + if value := vm.Get(api); value != nil && !goja.IsUndefined(value) { + t.Fatalf("%s API was exposed without permission", api) + } + } + + withPermissions := &loadedExtension{ + ID: "with-permissions", + Manifest: &ExtensionManifest{ + Name: "with-permissions", + Permissions: ExtensionPermissions{Storage: true, File: true}, + Capabilities: map[string]any{"rawFfmpeg": true}, + }, + DataDir: t.TempDir(), + } + vm = goja.New() + newExtensionRuntime(withPermissions).RegisterAPIs(vm) + for _, api := range []string{"storage", "credentials", "auth", "file", "ffmpeg"} { + if value := vm.Get(api); value == nil || goja.IsUndefined(value) { + t.Fatalf("%s API was not exposed with permission", api) + } + } +} + +func TestValidatePostProcessResultRestrictsReplacementTargets(t *testing.T) { + workDir := t.TempDir() + input := PostProcessInput{Path: filepath.Join(workDir, "input.flac"), URI: "content://input"} + ext := &loadedExtension{ + ID: "post-process", + Manifest: &ExtensionManifest{Name: "post-process"}, + DataDir: t.TempDir(), + } + if err := validatePostProcessResult(ext, input, &PostProcessResult{NewFilePath: input.Path, NewFileURI: input.URI}); err != nil { + t.Fatalf("unchanged target should be accepted: %v", err) + } + if err := validatePostProcessResult(ext, input, &PostProcessResult{NewFilePath: filepath.Join(workDir, "output.flac")}); err == nil { + t.Fatal("replacement path should require file permission") + } + ext.Manifest.Permissions.File = true + if err := validatePostProcessResult(ext, input, &PostProcessResult{NewFilePath: filepath.Join(workDir, "output.flac")}); err != nil { + t.Fatalf("sibling replacement should be accepted with file permission: %v", err) + } + if err := validatePostProcessResult(ext, input, &PostProcessResult{NewFileURI: "content://other"}); err == nil { + t.Fatal("replacement URI should be rejected") + } +} + func TestIsPrivateIP(t *testing.T) { tests := []struct { host string diff --git a/go_backend/extension_timeout.go b/go_backend/extension_timeout.go index a1b6fbe9..943b1dcf 100644 --- a/go_backend/extension_timeout.go +++ b/go_backend/extension_timeout.go @@ -11,14 +11,22 @@ import ( ) type JSExecutionError struct { - Message string - IsTimeout bool + Message string + IsTimeout bool + RuntimeUnsafe bool + Cause error } func (e *JSExecutionError) Error() string { return e.Message } +func (e *JSExecutionError) Unwrap() error { + return e.Cause +} + +var jsInterruptGracePeriod = 5 * time.Second + func RunWithTimeoutContext(ctx context.Context, vm *goja.Runtime, script string, timeout time.Duration) (goja.Value, error) { if vm == nil { return nil, fmt.Errorf("extension runtime unavailable") @@ -87,27 +95,29 @@ func RunWithTimeoutContext(ctx context.Context, vm *goja.Runtime, script string, // caller will access the VM concurrently and crash with a nil // pointer dereference. select { - case res := <-resultCh: + case <-resultCh: if cancelled { return nil, ErrExtensionRequestCancelled } - if res.err != nil { - return nil, res.err - } return nil, &JSExecutionError{ Message: "execution timeout exceeded", IsTimeout: true, } - case <-time.After(60 * time.Second): + case <-time.After(jsInterruptGracePeriod): // Goroutine is truly stuck (e.g. HTTP read with no timeout). // Log a warning — the VM should NOT be reused after this. GoLog("[extensionRuntime] WARNING: JS goroutine did not exit within 60s after interrupt, VM may be unsafe\n") + message := "execution timeout exceeded (runtime quarantined)" + var cause error if cancelled { - return nil, ErrExtensionRequestCancelled + message = "extension request cancelled (runtime quarantined)" + cause = ErrExtensionRequestCancelled } return nil, &JSExecutionError{ - Message: "execution timeout exceeded (force)", - IsTimeout: true, + Message: message, + IsTimeout: !cancelled, + RuntimeUnsafe: true, + Cause: cause, } } } @@ -122,13 +132,18 @@ func RunWithTimeoutAndRecover(vm *goja.Runtime, script string, timeout time.Dura func RunWithTimeoutContextAndRecover(ctx context.Context, vm *goja.Runtime, script string, timeout time.Duration) (goja.Value, error) { result, err := RunWithTimeoutContext(ctx, vm, script, timeout) - if vm != nil { + if vm != nil && !IsRuntimeUnsafeError(err) { vm.ClearInterrupt() } return result, err } +func IsRuntimeUnsafeError(err error) bool { + jsErr, ok := err.(*JSExecutionError) + return ok && jsErr.RuntimeUnsafe +} + func IsTimeoutError(err error) bool { if jsErr, ok := err.(*JSExecutionError); ok { return jsErr.IsTimeout diff --git a/go_backend/log_progress_timeout_supplement_test.go b/go_backend/log_progress_timeout_supplement_test.go index ebb2246e..e2529e4b 100644 --- a/go_backend/log_progress_timeout_supplement_test.go +++ b/go_backend/log_progress_timeout_supplement_test.go @@ -2,6 +2,7 @@ package gobackend import ( "bytes" + "context" "encoding/json" "errors" "strings" @@ -111,3 +112,51 @@ func TestRunWithTimeoutBranches(t *testing.T) { t.Fatal("JSExecutionError Error mismatch") } } + +func TestRunWithTimeoutQuarantinesUnresponsiveRuntime(t *testing.T) { + previousGrace := jsInterruptGracePeriod + jsInterruptGracePeriod = 10 * time.Millisecond + defer func() { jsInterruptGracePeriod = previousGrace }() + + vm := goja.New() + release := make(chan struct{}) + if err := vm.Set("block", func() { <-release }); err != nil { + t.Fatal(err) + } + _, err := RunWithTimeoutAndRecover(vm, "block()", 10*time.Millisecond) + if !IsRuntimeUnsafeError(err) { + close(release) + t.Fatalf("expected unsafe runtime error, got %v", err) + } + close(release) +} + +func TestRunWithTimeoutQuarantinesUnresponsiveCancelledRuntime(t *testing.T) { + previousGrace := jsInterruptGracePeriod + jsInterruptGracePeriod = 10 * time.Millisecond + defer func() { jsInterruptGracePeriod = previousGrace }() + + vm := goja.New() + entered := make(chan struct{}) + release := make(chan struct{}) + if err := vm.Set("block", func() { + close(entered) + <-release + }); err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + result := make(chan error, 1) + go func() { + _, err := RunWithTimeoutContextAndRecover(ctx, vm, "block()", time.Second) + result <- err + }() + <-entered + cancel() + err := <-result + if !IsRuntimeUnsafeError(err) || !errors.Is(err, ErrExtensionRequestCancelled) { + close(release) + t.Fatalf("expected unsafe cancellation error, got %v", err) + } + close(release) +}