fix(extensions): harden package and runtime lifecycle

This commit is contained in:
zarzet
2026-07-15 21:31:16 +07:00
parent ed9511b17f
commit bab954c6d9
15 changed files with 815 additions and 439 deletions
+24
View File
@@ -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
+3
View File
@@ -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
}
+232 -116
View File
@@ -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)
}
@@ -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)
}
+27
View File
@@ -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 {
+9 -1
View File
@@ -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"
+38
View File
@@ -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 {
+67 -65
View File
@@ -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)
+28 -8
View File
@@ -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)
}
+112 -227
View File
@@ -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)
}
+71 -3
View File
@@ -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
@@ -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)
+108
View File
@@ -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
+26 -11
View File
@@ -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
@@ -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)
}