Commit 210734bb by CaIon

refactor(task): remove the custom-plugin layer switch

TaskPluginOverrideEnabled had no UI since the master switch landed, yet
when left off it marked every third-party plugin "disabled; platform
unavailable" regardless of its own toggle. Drop the option, env var,
registry flag, and the dead branch in ListTaskPlugins; the master switch
and per-plugin toggles are the only two levels now.
parent 92bc7ff7
...@@ -190,7 +190,6 @@ func initConstantEnv() { ...@@ -190,7 +190,6 @@ func initConstantEnv() {
constant.GetMediaTokenNotStream = GetEnvOrDefaultBool("GET_MEDIA_TOKEN_NOT_STREAM", false) constant.GetMediaTokenNotStream = GetEnvOrDefaultBool("GET_MEDIA_TOKEN_NOT_STREAM", false)
constant.UpdateTask = GetEnvOrDefaultBool("UPDATE_TASK", true) constant.UpdateTask = GetEnvOrDefaultBool("UPDATE_TASK", true)
constant.TaskPluginEnabled = GetEnvOrDefaultBool("TASK_PLUGIN_ENABLED", true) constant.TaskPluginEnabled = GetEnvOrDefaultBool("TASK_PLUGIN_ENABLED", true)
constant.TaskPluginOverrideEnabled = GetEnvOrDefaultBool("TASK_PLUGIN_OVERRIDE_ENABLED", true)
constant.AzureDefaultAPIVersion = GetEnvOrDefaultString("AZURE_DEFAULT_API_VERSION", "2025-04-01-preview") constant.AzureDefaultAPIVersion = GetEnvOrDefaultString("AZURE_DEFAULT_API_VERSION", "2025-04-01-preview")
constant.NotifyLimitCount = GetEnvOrDefault("NOTIFY_LIMIT_COUNT", 2) constant.NotifyLimitCount = GetEnvOrDefault("NOTIFY_LIMIT_COUNT", 2)
constant.NotificationLimitDurationMinute = GetEnvOrDefault("NOTIFICATION_LIMIT_DURATION_MINUTE", 10) constant.NotificationLimitDurationMinute = GetEnvOrDefault("NOTIFICATION_LIMIT_DURATION_MINUTE", 10)
......
...@@ -27,11 +27,6 @@ var legacyTaskActionAliases = map[string]string{ ...@@ -27,11 +27,6 @@ var legacyTaskActionAliases = map[string]string{
// When disabled, factory and override plugins both stop serving. // When disabled, factory and override plugins both stop serving.
var TaskPluginEnabled = true var TaskPluginEnabled = true
// TaskPluginOverrideEnabled controls whether the database override layer is
// active. When disabled, uploaded plugins are ignored and factory plugins are
// used instead; the factory layer is unaffected.
var TaskPluginOverrideEnabled = true
// NormalizeTaskAction maps persisted legacy action names to the canonical task // NormalizeTaskAction maps persisted legacy action names to the canonical task
// action vocabulary. Unknown platform-specific actions pass through unchanged. // action vocabulary. Unknown platform-specific actions pass through unchanged.
func NormalizeTaskAction(action string) string { func NormalizeTaskAction(action string) string {
......
...@@ -13,7 +13,6 @@ import ( ...@@ -13,7 +13,6 @@ import (
"time" "time"
"github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/logger" "github.com/QuantumNous/new-api/logger"
"github.com/QuantumNous/new-api/model" "github.com/QuantumNous/new-api/model"
"github.com/QuantumNous/new-api/pkg/jsplugin" "github.com/QuantumNous/new-api/pkg/jsplugin"
...@@ -181,9 +180,7 @@ func ListTaskPlugins(c *gin.Context) { ...@@ -181,9 +180,7 @@ func ListTaskPlugins(c *gin.Context) {
item.Active = row.Active item.Active = row.Active
item.SourceHash = row.SourceHash item.SourceHash = row.SourceHash
item.Remark = row.Remark item.Remark = row.Remark
if !constant.TaskPluginOverrideEnabled { if message := runtimeErrors[key]; message != "" {
item.RuntimeStatus = "disabled_fallback"
} else if message := runtimeErrors[key]; message != "" {
item.RuntimeStatus = "compile_failed" item.RuntimeStatus = "compile_failed"
item.RuntimeError = message item.RuntimeError = message
} else if runtimeMeta, ok := override[key]; ok { } else if runtimeMeta, ok := override[key]; ok {
......
...@@ -399,49 +399,6 @@ export function parseTaskResult() { return {}; } ...@@ -399,49 +399,6 @@ export function parseTaskResult() { return {}; }
t.Fatal("task plugin option not found") t.Fatal("task plugin option not found")
} }
func TestListTaskPluginsShowsDisabledFallbackWhenOverridesAreDisabled(t *testing.T) {
setupTaskPluginControllerTest(t)
factorySource, err := plugins.Source("kling")
require.NoError(t, err)
factoryVersion := klingFactoryVersion(t)
overrideSource := strings.Replace(factorySource, `version: "`+factoryVersion+`"`, `version: "`+factoryVersion+`-test-disabled-override"`, 1)
require.NotEqual(t, factorySource, overrideSource, "factory version marker must be found in kling source")
loaded, err := jsplugin.DefaultRegistry.Register(overrideSource, jsplugin.Options{})
require.NoError(t, err)
plugin := model.TaskPlugin{
Key: "kling", APIVersion: loaded.Meta.APIVersion, Version: loaded.Meta.Version,
Source: overrideSource, SourceHash: "test-hash", Enabled: true,
}
require.NoError(t, model.SaveTaskPlugin(&plugin))
originalEnabled := constant.TaskPluginOverrideEnabled
constant.TaskPluginOverrideEnabled = false
jsplugin.DefaultRegistry.SetOverrideEnabled(false)
t.Cleanup(func() {
constant.TaskPluginOverrideEnabled = originalEnabled
jsplugin.DefaultRegistry.SetOverrideEnabled(originalEnabled)
jsplugin.DefaultRegistry.Unregister("kling")
})
recorder := httptest.NewRecorder()
context, _ := gin.CreateTestContext(recorder)
context.Request = httptest.NewRequest(http.MethodGet, "/api/plugin/task", nil)
ListTaskPlugins(context)
var response struct {
Success bool `json:"success"`
Data []taskPluginListItem `json:"data"`
}
require.NoError(t, common.Unmarshal(recorder.Body.Bytes(), &response))
require.True(t, response.Success)
for _, item := range response.Data {
if item.Meta.Key == "kling" {
assert.Equal(t, "disabled_fallback", item.RuntimeStatus)
return
}
}
t.Fatal("kling plugin not found")
}
func TestDeleteActiveOverrideFallsBackToFactoryAndDeletesRecord(t *testing.T) { func TestDeleteActiveOverrideFallsBackToFactoryAndDeletesRecord(t *testing.T) {
setupTaskPluginControllerTest(t) setupTaskPluginControllerTest(t)
factorySource, err := plugins.Source("kling") factorySource, err := plugins.Source("kling")
...@@ -800,40 +757,6 @@ func TestSyncTaskPluginsCachesRejectedDesiredSourceWithoutLosingIncumbent(t *tes ...@@ -800,40 +757,6 @@ func TestSyncTaskPluginsCachesRejectedDesiredSourceWithoutLosingIncumbent(t *tes
assert.NotContains(t, jsplugin.DefaultRegistry.RoutingErrors(), pluginKey) assert.NotContains(t, jsplugin.DefaultRegistry.RoutingErrors(), pluginKey)
} }
func TestSyncTaskPluginsPreservesLastCompiledOverrideWhileOverridesAreDisabled(t *testing.T) {
setupTaskPluginControllerTest(t)
key := "sync-disabled-override"
cleanupTaskPluginControllerRuntime(t, key)
v1Source := taskPluginControllerTestSource(key, "1.0.0")
require.NoError(t, model.SaveTaskPlugin(&model.TaskPlugin{
Key: key, APIVersion: 1, Version: "1.0.0", Source: v1Source, SourceHash: "disabled-v1", Enabled: true,
}))
require.NoError(t, syncTaskPluginsOnce())
jsplugin.DefaultRegistry.SetOverrideEnabled(false)
t.Cleanup(func() { jsplugin.DefaultRegistry.SetOverrideEnabled(true) })
disabledGeneration := jsplugin.DefaultRegistry.Generation()
require.NoError(t, syncTaskPluginsOnce())
assert.Same(t, disabledGeneration, jsplugin.DefaultRegistry.Generation())
require.NoError(t, model.SaveTaskPlugin(&model.TaskPlugin{
Key: key, APIVersion: 1, Version: "2.0.0",
Source: "export const meta = {", SourceHash: "disabled-v2", Enabled: true,
}))
require.NoError(t, model.ActivateTaskPlugin(key, "2.0.0"))
require.NoError(t, syncTaskPluginsOnce())
assert.Equal(t, "1.0.0", jsplugin.DefaultRegistry.OverridePlugins()[key].Meta.Version)
taskPluginSyncState.Lock()
assert.Equal(t, "disabled-v1", taskPluginSyncState.hashes[key])
assert.NotEmpty(t, taskPluginSyncState.errors[key])
taskPluginSyncState.Unlock()
jsplugin.DefaultRegistry.SetOverrideEnabled(true)
active, ok := jsplugin.DefaultRegistry.Get(key)
require.True(t, ok)
assert.Equal(t, "1.0.0", active.Meta.Version)
}
const dryRunPluginSource = ` const dryRunPluginSource = `
export const meta = {apiVersion: 1, key: "dryrun-probe", name: "DryRun", version: "1.0.0", author: {name: "Test"}, models: ["doc-1"], fetchMode: "per_task"}; export const meta = {apiVersion: 1, key: "dryrun-probe", name: "DryRun", version: "1.0.0", author: {name: "Test"}, models: ["doc-1"], fetchMode: "per_task"};
export function buildSubmitRequest(payload) { export function buildSubmitRequest(payload) {
......
...@@ -56,8 +56,6 @@ func InitOptionMap() { ...@@ -56,8 +56,6 @@ func InitOptionMap() {
common.OptionMap["TaskEnabled"] = strconv.FormatBool(common.TaskEnabled) common.OptionMap["TaskEnabled"] = strconv.FormatBool(common.TaskEnabled)
common.OptionMap["TaskPluginEnabled"] = strconv.FormatBool(constant.TaskPluginEnabled) common.OptionMap["TaskPluginEnabled"] = strconv.FormatBool(constant.TaskPluginEnabled)
jsplugin.DefaultRegistry.SetEnabled(constant.TaskPluginEnabled) jsplugin.DefaultRegistry.SetEnabled(constant.TaskPluginEnabled)
common.OptionMap["TaskPluginOverrideEnabled"] = strconv.FormatBool(constant.TaskPluginOverrideEnabled)
jsplugin.DefaultRegistry.SetOverrideEnabled(constant.TaskPluginOverrideEnabled)
common.OptionMap[setting.TaskPluginMarketplaceSourcesKey] = setting.TaskPluginMarketplaceSources2JsonString() common.OptionMap[setting.TaskPluginMarketplaceSourcesKey] = setting.TaskPluginMarketplaceSources2JsonString()
common.OptionMap[setting.TaskPluginDisabledFactoryKeysKey] = "[]" common.OptionMap[setting.TaskPluginDisabledFactoryKeysKey] = "[]"
jsplugin.DefaultRegistry.SetDisabledFactoryKeys(nil) jsplugin.DefaultRegistry.SetDisabledFactoryKeys(nil)
...@@ -368,9 +366,6 @@ func updateOptionMap(key string, value string) (err error) { ...@@ -368,9 +366,6 @@ func updateOptionMap(key string, value string) (err error) {
case "TaskPluginEnabled": case "TaskPluginEnabled":
constant.TaskPluginEnabled = boolValue constant.TaskPluginEnabled = boolValue
jsplugin.DefaultRegistry.SetEnabled(boolValue) jsplugin.DefaultRegistry.SetEnabled(boolValue)
case "TaskPluginOverrideEnabled":
constant.TaskPluginOverrideEnabled = boolValue
jsplugin.DefaultRegistry.SetOverrideEnabled(boolValue)
case "DataExportEnabled": case "DataExportEnabled":
common.DataExportEnabled = boolValue common.DataExportEnabled = boolValue
case "DefaultCollapseSidebar": case "DefaultCollapseSidebar":
......
...@@ -46,22 +46,6 @@ export function parseTaskResult() { return {}; } ...@@ -46,22 +46,6 @@ export function parseTaskResult() { return {}; }
assert.True(t, ok) assert.True(t, ok)
} }
func TestTaskPluginOverrideEnabledOptionUpdatesRuntimeSwitch(t *testing.T) {
originalEnabled := constant.TaskPluginOverrideEnabled
originalMap := common.OptionMap
common.OptionMap = map[string]string{}
t.Cleanup(func() {
constant.TaskPluginOverrideEnabled = originalEnabled
jsplugin.DefaultRegistry.SetOverrideEnabled(originalEnabled)
common.OptionMap = originalMap
})
require.NoError(t, updateOptionMap("TaskPluginOverrideEnabled", "false"))
assert.False(t, constant.TaskPluginOverrideEnabled)
assert.Equal(t, "false", common.OptionMap["TaskPluginOverrideEnabled"])
}
func TestTaskPluginDisabledFactoryKeysOptionUpdatesRegistry(t *testing.T) { func TestTaskPluginDisabledFactoryKeysOptionUpdatesRegistry(t *testing.T) {
originalMap := common.OptionMap originalMap := common.OptionMap
common.OptionMap = map[string]string{} common.OptionMap = map[string]string{}
......
...@@ -169,7 +169,6 @@ type Registry struct { ...@@ -169,7 +169,6 @@ type Registry struct {
activeOverride map[string]*LoadedPlugin activeOverride map[string]*LoadedPlugin
disabledFactory map[string]struct{} disabledFactory map[string]struct{}
masterEnabled atomic.Bool masterEnabled atomic.Bool
overrideEnabled atomic.Bool
generation atomic.Pointer[RoutingGeneration] generation atomic.Pointer[RoutingGeneration]
preparer RoutingGenerationPreparer preparer RoutingGenerationPreparer
routingErrors map[string]string routingErrors map[string]string
...@@ -185,8 +184,7 @@ func NewRegistry() *Registry { ...@@ -185,8 +184,7 @@ func NewRegistry() *Registry {
routingErrors: make(map[string]string), routingErrors: make(map[string]string),
} }
registry.masterEnabled.Store(true) registry.masterEnabled.Store(true)
registry.overrideEnabled.Store(true) generation, _ := buildRoutingGeneration(registry.factory, registry.override, 0)
generation, _ := buildRoutingGeneration(registry.factory, registry.override, true, 0)
registry.generation.Store(generation) registry.generation.Store(generation)
registry.lastRebuild = RoutingRebuildOutcome{ registry.lastRebuild = RoutingRebuildOutcome{
Status: "success", Status: "success",
...@@ -221,8 +219,7 @@ func (r *Registry) register(source string, options Options, factory bool) (*Load ...@@ -221,8 +219,7 @@ func (r *Registry) register(source string, options Options, factory bool) (*Load
} else { } else {
overridePlugins[plugin.Meta.Key] = plugin overridePlugins[plugin.Meta.Key] = plugin
} }
enabled := r.overrideEnabled.Load() generation, routingErrors, err := r.prepareGeneration(filterDisabledFactory(factoryPlugins, r.disabledFactory), overridePlugins, false, nil)
generation, routingErrors, err := r.prepareGeneration(filterDisabledFactory(factoryPlugins, r.disabledFactory), overridePlugins, enabled, false, nil)
if err != nil { if err != nil {
r.recordRebuildFailure(err) r.recordRebuildFailure(err)
return nil, err return nil, err
...@@ -233,7 +230,7 @@ func (r *Registry) register(source string, options Options, factory bool) (*Load ...@@ -233,7 +230,7 @@ func (r *Registry) register(source string, options Options, factory bool) (*Load
} }
r.factory = factoryPlugins r.factory = factoryPlugins
r.override = overridePlugins r.override = overridePlugins
r.publishGeneration(generation, routingErrors, r.resolveActiveOverrides(generation, overridePlugins, enabled)) r.publishGeneration(generation, routingErrors, r.resolveActiveOverrides(generation, overridePlugins))
return plugin, nil return plugin, nil
} }
...@@ -460,37 +457,17 @@ func (r *Registry) SetEnabled(enabled bool) { ...@@ -460,37 +457,17 @@ func (r *Registry) SetEnabled(enabled bool) {
} }
previous := r.masterEnabled.Load() previous := r.masterEnabled.Load()
r.masterEnabled.Store(enabled) r.masterEnabled.Store(enabled)
overrideEnabled := r.overrideEnabled.Load()
var retainCurrent map[string]struct{}
if enabled && overrideEnabled {
retainCurrent = pluginMapKeys(r.override)
}
generation, routingErrors, err := r.prepareGeneration(filterDisabledFactory(r.factory, r.disabledFactory), r.override, overrideEnabled, true, retainCurrent)
if err != nil {
r.masterEnabled.Store(previous)
r.recordRebuildFailure(err)
return
}
r.publishGeneration(generation, routingErrors, r.resolveActiveOverrides(generation, r.override, overrideEnabled))
}
func (r *Registry) SetOverrideEnabled(enabled bool) {
r.mu.Lock()
defer r.mu.Unlock()
if r.overrideEnabled.Load() == enabled {
return
}
var retainCurrent map[string]struct{} var retainCurrent map[string]struct{}
if enabled { if enabled {
retainCurrent = pluginMapKeys(r.override) retainCurrent = pluginMapKeys(r.override)
} }
generation, routingErrors, err := r.prepareGeneration(filterDisabledFactory(r.factory, r.disabledFactory), r.override, enabled, true, retainCurrent) generation, routingErrors, err := r.prepareGeneration(filterDisabledFactory(r.factory, r.disabledFactory), r.override, true, retainCurrent)
if err != nil { if err != nil {
r.masterEnabled.Store(previous)
r.recordRebuildFailure(err) r.recordRebuildFailure(err)
return return
} }
r.overrideEnabled.Store(enabled) r.publishGeneration(generation, routingErrors, r.resolveActiveOverrides(generation, r.override))
r.publishGeneration(generation, routingErrors, r.resolveActiveOverrides(generation, r.override, enabled))
} }
func (r *Registry) SetDisabledFactoryKeys(keys []string) { func (r *Registry) SetDisabledFactoryKeys(keys []string) {
...@@ -518,18 +495,14 @@ func (r *Registry) SetDisabledFactoryKeys(keys []string) { ...@@ -518,18 +495,14 @@ func (r *Registry) SetDisabledFactoryKeys(keys []string) {
} }
} }
enabled := r.overrideEnabled.Load() retainCurrent := pluginMapKeys(r.override)
var retainCurrent map[string]struct{} generation, routingErrors, err := r.prepareGeneration(filterDisabledFactory(r.factory, next), r.override, true, retainCurrent)
if enabled {
retainCurrent = pluginMapKeys(r.override)
}
generation, routingErrors, err := r.prepareGeneration(filterDisabledFactory(r.factory, next), r.override, enabled, true, retainCurrent)
if err != nil { if err != nil {
r.recordRebuildFailure(err) r.recordRebuildFailure(err)
return return
} }
r.disabledFactory = next r.disabledFactory = next
r.publishGeneration(generation, routingErrors, r.resolveActiveOverrides(generation, r.override, enabled)) r.publishGeneration(generation, routingErrors, r.resolveActiveOverrides(generation, r.override))
} }
func (r *Registry) Unregister(key string) error { func (r *Registry) Unregister(key string) error {
...@@ -540,18 +513,14 @@ func (r *Registry) Unregister(key string) error { ...@@ -540,18 +513,14 @@ func (r *Registry) Unregister(key string) error {
} }
overridePlugins := clonePluginMap(r.override) overridePlugins := clonePluginMap(r.override)
delete(overridePlugins, key) delete(overridePlugins, key)
enabled := r.overrideEnabled.Load() retainCurrent := pluginMapKeys(overridePlugins)
var retainCurrent map[string]struct{} generation, routingErrors, err := r.prepareGeneration(filterDisabledFactory(r.factory, r.disabledFactory), overridePlugins, true, retainCurrent)
if enabled {
retainCurrent = pluginMapKeys(overridePlugins)
}
generation, routingErrors, err := r.prepareGeneration(filterDisabledFactory(r.factory, r.disabledFactory), overridePlugins, enabled, true, retainCurrent)
if err != nil { if err != nil {
r.recordRebuildFailure(err) r.recordRebuildFailure(err)
return err return err
} }
r.override = overridePlugins r.override = overridePlugins
r.publishGeneration(generation, routingErrors, r.resolveActiveOverrides(generation, overridePlugins, enabled)) r.publishGeneration(generation, routingErrors, r.resolveActiveOverrides(generation, overridePlugins))
return nil return nil
} }
...@@ -573,18 +542,14 @@ func (r *Registry) ReplaceOverrides(plugins []*LoadedPlugin) error { ...@@ -573,18 +542,14 @@ func (r *Registry) ReplaceOverrides(plugins []*LoadedPlugin) error {
if samePluginMap(r.override, overridePlugins) { if samePluginMap(r.override, overridePlugins) {
return nil return nil
} }
enabled := r.overrideEnabled.Load() retainCurrent := pluginMapKeys(overridePlugins)
var retainCurrent map[string]struct{} generation, routingErrors, err := r.prepareGeneration(filterDisabledFactory(r.factory, r.disabledFactory), overridePlugins, true, retainCurrent)
if enabled {
retainCurrent = pluginMapKeys(overridePlugins)
}
generation, routingErrors, err := r.prepareGeneration(filterDisabledFactory(r.factory, r.disabledFactory), overridePlugins, enabled, true, retainCurrent)
if err != nil { if err != nil {
r.recordRebuildFailure(err) r.recordRebuildFailure(err)
return err return err
} }
r.override = overridePlugins r.override = overridePlugins
r.publishGeneration(generation, routingErrors, r.resolveActiveOverrides(generation, overridePlugins, enabled)) r.publishGeneration(generation, routingErrors, r.resolveActiveOverrides(generation, overridePlugins))
return nil return nil
} }
...@@ -610,18 +575,14 @@ func (r *Registry) SetGenerationPreparer(preparer RoutingGenerationPreparer) err ...@@ -610,18 +575,14 @@ func (r *Registry) SetGenerationPreparer(preparer RoutingGenerationPreparer) err
previous := r.preparer previous := r.preparer
r.preparer = preparer r.preparer = preparer
enabled := r.overrideEnabled.Load() retainCurrent := pluginMapKeys(r.override)
var retainCurrent map[string]struct{} generation, routingErrors, err := r.prepareGeneration(filterDisabledFactory(r.factory, r.disabledFactory), r.override, true, retainCurrent)
if enabled {
retainCurrent = pluginMapKeys(r.override)
}
generation, routingErrors, err := r.prepareGeneration(filterDisabledFactory(r.factory, r.disabledFactory), r.override, enabled, true, retainCurrent)
if err != nil { if err != nil {
r.preparer = previous r.preparer = previous
r.recordRebuildFailure(err) r.recordRebuildFailure(err)
return err return err
} }
r.publishGeneration(generation, routingErrors, r.resolveActiveOverrides(generation, r.override, enabled)) r.publishGeneration(generation, routingErrors, r.resolveActiveOverrides(generation, r.override))
return nil return nil
} }
...@@ -659,7 +620,7 @@ func (r *Registry) RoutingStatus() RoutingStatus { ...@@ -659,7 +620,7 @@ func (r *Registry) RoutingStatus() RoutingStatus {
func (r *Registry) prepareGeneration( func (r *Registry) prepareGeneration(
factory, override map[string]*LoadedPlugin, factory, override map[string]*LoadedPlugin,
enabled, tolerateConflicts bool, tolerateConflicts bool,
retainCurrent map[string]struct{}, retainCurrent map[string]struct{},
) (*RoutingGeneration, map[string]string, error) { ) (*RoutingGeneration, map[string]string, error) {
if !r.masterEnabled.Load() { if !r.masterEnabled.Load() {
...@@ -677,21 +638,14 @@ func (r *Registry) prepareGeneration( ...@@ -677,21 +638,14 @@ func (r *Registry) prepareGeneration(
err error err error
) )
if tolerateConflicts { if tolerateConflicts {
generation, routingErrors, err = buildRoutingGenerationAdmitting(factory, override, enabled, number, current, retainCurrent) generation, routingErrors, err = buildRoutingGenerationAdmitting(factory, override, number, current, retainCurrent)
} else { } else {
generation, err = buildRoutingGeneration(factory, override, enabled, number) generation, err = buildRoutingGeneration(factory, override, number)
routingErrors = make(map[string]string) routingErrors = make(map[string]string)
} }
if err != nil { if err != nil {
return nil, nil, err return nil, nil, err
} }
if !tolerateConflicts {
// Both runtime switch positions must remain publishable so toggling the
// override layer never exposes an invalid generation.
if _, err = buildRoutingGeneration(factory, override, !enabled, number); err != nil {
return nil, nil, err
}
}
if r.preparer != nil { if r.preparer != nil {
prepared, prepareErr := r.preparer(generation, current) prepared, prepareErr := r.preparer(generation, current)
if prepareErr != nil { if prepareErr != nil {
...@@ -921,12 +875,8 @@ func pluginMapKeys(plugins map[string]*LoadedPlugin) map[string]struct{} { ...@@ -921,12 +875,8 @@ func pluginMapKeys(plugins map[string]*LoadedPlugin) map[string]struct{} {
func (r *Registry) resolveActiveOverrides( func (r *Registry) resolveActiveOverrides(
generation *RoutingGeneration, generation *RoutingGeneration,
override map[string]*LoadedPlugin, override map[string]*LoadedPlugin,
enabled bool,
) map[string]*LoadedPlugin { ) map[string]*LoadedPlugin {
active := make(map[string]*LoadedPlugin) active := make(map[string]*LoadedPlugin)
if !enabled {
return active
}
for _, plugin := range generation.plugins { for _, plugin := range generation.plugins {
desired, hasOverride := override[plugin.Meta.Key] desired, hasOverride := override[plugin.Meta.Key]
if !hasOverride { if !hasOverride {
......
...@@ -55,10 +55,6 @@ func TestRegistrySetDisabledFactoryKeysHidesFactoryPlugin(t *testing.T) { ...@@ -55,10 +55,6 @@ func TestRegistrySetDisabledFactoryKeysHidesFactoryPlugin(t *testing.T) {
assert.Equal(t, "factory-off", snapshot.Factory[0].Key) assert.Equal(t, "factory-off", snapshot.Factory[0].Key)
assert.Equal(t, []string{"factory-off"}, snapshot.DisabledFactory) assert.Equal(t, []string{"factory-off"}, snapshot.DisabledFactory)
registry.SetOverrideEnabled(false)
_, ok = registry.Get("factory-off")
assert.False(t, ok)
registry.SetDisabledFactoryKeys(nil) registry.SetDisabledFactoryKeys(nil)
plugin, ok := registry.Get("factory-off") plugin, ok := registry.Get("factory-off")
require.True(t, ok) require.True(t, ok)
......
...@@ -121,15 +121,9 @@ func TestRegistrySetEnabledNoOpDoesNotBumpGeneration(t *testing.T) { ...@@ -121,15 +121,9 @@ func TestRegistrySetEnabledNoOpDoesNotBumpGeneration(t *testing.T) {
func TestRegistryMasterSwitchIsOrthogonalToLayerFlags(t *testing.T) { func TestRegistryMasterSwitchIsOrthogonalToLayerFlags(t *testing.T) {
registry := NewRegistry() registry := NewRegistry()
require.NoError(t, registerTestPlugin(registry, "1.0.0-factory", true)) require.NoError(t, registerTestPlugin(registry, "1.0.0-factory", true))
require.NoError(t, registerTestPlugin(registry, "1.0.0-override", false))
registry.SetOverrideEnabled(false)
plugin, ok := registry.Get("test")
require.True(t, ok)
assert.Equal(t, "1.0.0-factory", plugin.Meta.Version)
registry.SetDisabledFactoryKeys([]string{"test"}) registry.SetDisabledFactoryKeys([]string{"test"})
_, ok = registry.Get("test") _, ok := registry.Get("test")
assert.False(t, ok) assert.False(t, ok)
registry.SetEnabled(false) registry.SetEnabled(false)
...@@ -141,7 +135,7 @@ func TestRegistryMasterSwitchIsOrthogonalToLayerFlags(t *testing.T) { ...@@ -141,7 +135,7 @@ func TestRegistryMasterSwitchIsOrthogonalToLayerFlags(t *testing.T) {
assert.False(t, ok) assert.False(t, ok)
registry.SetDisabledFactoryKeys(nil) registry.SetDisabledFactoryKeys(nil)
plugin, ok = registry.Get("test") plugin, ok := registry.Get("test")
require.True(t, ok) require.True(t, ok)
assert.Equal(t, "1.0.0-factory", plugin.Meta.Version) assert.Equal(t, "1.0.0-factory", plugin.Meta.Version)
} }
...@@ -33,22 +33,6 @@ func TestRegistryUnregisterFallsBackToFactory(t *testing.T) { ...@@ -33,22 +33,6 @@ func TestRegistryUnregisterFallsBackToFactory(t *testing.T) {
assert.Equal(t, "1.0.0-factory", plugin.Meta.Version) assert.Equal(t, "1.0.0-factory", plugin.Meta.Version)
} }
func TestRegistryDisabledOverrideFallsBackToFactoryAndCanBeRestored(t *testing.T) {
registry := NewRegistry()
require.NoError(t, registerTestPlugin(registry, "1.0.0-factory", true))
require.NoError(t, registerTestPlugin(registry, "1.0.0-override", false))
registry.SetOverrideEnabled(false)
plugin, ok := registry.Get("test")
require.True(t, ok)
assert.Equal(t, "1.0.0-factory", plugin.Meta.Version)
registry.SetOverrideEnabled(true)
plugin, ok = registry.Get("test")
require.True(t, ok)
assert.Equal(t, "1.0.0-override", plugin.Meta.Version)
}
func TestRegistrySnapshotSeparatesLayersWithoutExposingEntries(t *testing.T) { func TestRegistrySnapshotSeparatesLayersWithoutExposingEntries(t *testing.T) {
registry := NewRegistry() registry := NewRegistry()
require.NoError(t, registerTestPlugin(registry, "1.0.0-factory", true)) require.NoError(t, registerTestPlugin(registry, "1.0.0-factory", true))
......
...@@ -531,7 +531,7 @@ func (g *RoutingGeneration) RebuildWithPlugins(plugins []*LoadedPlugin) (*Routin ...@@ -531,7 +531,7 @@ func (g *RoutingGeneration) RebuildWithPlugins(plugins []*LoadedPlugin) (*Routin
} }
byKey[plugin.Meta.Key] = plugin byKey[plugin.Meta.Key] = plugin
} }
rebuilt, err := buildRoutingGeneration(byKey, nil, false, g.Number) rebuilt, err := buildRoutingGeneration(byKey, nil, g.Number)
if err != nil { if err != nil {
return nil, err return nil, err
} }
...@@ -764,19 +764,18 @@ func ResolveRouteAction(route Route, resolvedAction string) string { ...@@ -764,19 +764,18 @@ func ResolveRouteAction(route Route, resolvedAction string) string {
return route.Action return route.Action
} }
func buildRoutingGeneration(factory, override map[string]*LoadedPlugin, overrideEnabled bool, number uint64) (*RoutingGeneration, error) { func buildRoutingGeneration(factory, override map[string]*LoadedPlugin, number uint64) (*RoutingGeneration, error) {
effective := effectivePlugins(factory, override, overrideEnabled) effective := effectivePlugins(factory, override)
return buildRoutingGenerationFromPlugins(effective, number) return buildRoutingGenerationFromPlugins(effective, number)
} }
func buildRoutingGenerationAdmitting( func buildRoutingGenerationAdmitting(
factory, override map[string]*LoadedPlugin, factory, override map[string]*LoadedPlugin,
overrideEnabled bool,
number uint64, number uint64,
current *RoutingGeneration, current *RoutingGeneration,
retainCurrent map[string]struct{}, retainCurrent map[string]struct{},
) (*RoutingGeneration, map[string]string, error) { ) (*RoutingGeneration, map[string]string, error) {
candidates := effectivePlugins(factory, override, overrideEnabled) candidates := effectivePlugins(factory, override)
accepted := make(map[string]*LoadedPlugin, len(candidates)) accepted := make(map[string]*LoadedPlugin, len(candidates))
currentByKey := make(map[string]*LoadedPlugin) currentByKey := make(map[string]*LoadedPlugin)
if current != nil { if current != nil {
...@@ -843,12 +842,10 @@ func buildRoutingGenerationAdmitting( ...@@ -843,12 +842,10 @@ func buildRoutingGenerationAdmitting(
return generation, routingErrors, nil return generation, routingErrors, nil
} }
func effectivePlugins(factory, override map[string]*LoadedPlugin, overrideEnabled bool) map[string]*LoadedPlugin { func effectivePlugins(factory, override map[string]*LoadedPlugin) map[string]*LoadedPlugin {
effective := make(map[string]*LoadedPlugin, len(factory)+len(override)) effective := make(map[string]*LoadedPlugin, len(factory)+len(override))
maps.Copy(effective, factory) maps.Copy(effective, factory)
if overrideEnabled { maps.Copy(effective, override)
maps.Copy(effective, override)
}
return effective return effective
} }
......
...@@ -388,9 +388,6 @@ func TestRegistryNoOpMutationsKeepCurrentGeneration(t *testing.T) { ...@@ -388,9 +388,6 @@ func TestRegistryNoOpMutationsKeepCurrentGeneration(t *testing.T) {
require.NoError(t, registry.ReplaceOverrides([]*LoadedPlugin{plugin})) require.NoError(t, registry.ReplaceOverrides([]*LoadedPlugin{plugin}))
assert.Same(t, current, registry.Generation()) assert.Same(t, current, registry.Generation())
registry.SetOverrideEnabled(true)
assert.Same(t, current, registry.Generation())
require.NoError(t, registry.Unregister("missing")) require.NoError(t, registry.Unregister("missing"))
assert.Same(t, current, registry.Generation()) assert.Same(t, current, registry.Generation())
} }
...@@ -686,23 +683,6 @@ func TestRejectedNewOverrideRetainsFactoryIncumbent(t *testing.T) { ...@@ -686,23 +683,6 @@ func TestRejectedNewOverrideRetainsFactoryIncumbent(t *testing.T) {
assert.Contains(t, registry.RoutingErrors()["factory-fallback"], "channelType 122 conflicts") assert.Contains(t, registry.RoutingErrors()["factory-fallback"], "channelType 122 conflicts")
} }
func TestDisablingOverridesPublishesFactoryInsteadOfRetainingOverride(t *testing.T) {
registry := NewRegistry()
factorySource := routingTestPluginSource("switchable", 93, `["factory"]`, "", "")
factory, err := registry.RegisterFactory(factorySource, Options{})
require.NoError(t, err)
override := mustCompileRoutingPlugin(t, "switchable", 94, `["override"]`, "", "")
require.NoError(t, registry.ReplaceOverrides([]*LoadedPlugin{override}))
registry.SetOverrideEnabled(false)
active, ok := registry.Get("switchable")
require.True(t, ok)
assert.Same(t, factory, active)
assert.Same(t, override, registry.OverridePlugins()["switchable"])
assert.Empty(t, registry.ActiveOverridePlugins())
}
func TestGenericChannelTypesDoNotCreateLegacyIdentityConflicts(t *testing.T) { func TestGenericChannelTypesDoNotCreateLegacyIdentityConflicts(t *testing.T) {
for _, channelType := range []int{0, constant.ChannelTypeTaskPlugin} { for _, channelType := range []int{0, constant.ChannelTypeTaskPlugin} {
t.Run(fmt.Sprintf("channel_%d", channelType), func(t *testing.T) { t.Run(fmt.Sprintf("channel_%d", channelType), func(t *testing.T) {
......
Markdown is supported
0% or
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or sign in to comment