diff --git a/go-gost/x/socket/chain.go b/go-gost/x/socket/chain.go index 5d6149e..74b6325 100644 --- a/go-gost/x/socket/chain.go +++ b/go-gost/x/socket/chain.go @@ -43,8 +43,8 @@ func updateChain(req updateChainRequest) error { name := strings.TrimSpace(req.Chain) - if !registry.ChainRegistry().IsRegistered(name) { - return errors.New("chain " + name + " not found") + if registry.ChainRegistry().IsRegistered(name) { + registry.ChainRegistry().Unregister(name) } req.Data.Name = name @@ -54,19 +54,22 @@ func updateChain(req updateChainRequest) error { return errors.New("create chain " + name + " failed: " + err.Error()) } - registry.ChainRegistry().Unregister(name) - if err := registry.ChainRegistry().Register(name, v); err != nil { return errors.New("chain " + name + " already exists") } config.OnUpdate(func(c *config.Config) error { + found := false for i := range c.Chains { if c.Chains[i].Name == name { c.Chains[i] = &req.Data + found = true break } } + if !found { + c.Chains = append(c.Chains, &req.Data) + } return nil }) @@ -77,10 +80,9 @@ func deleteChain(req deleteChainRequest) error { name := strings.TrimSpace(req.Chain) - if !registry.ChainRegistry().IsRegistered(name) { - return errors.New("chain " + name + " not found") + if registry.ChainRegistry().IsRegistered(name) { + registry.ChainRegistry().Unregister(name) } - registry.ChainRegistry().Unregister(name) config.OnUpdate(func(c *config.Config) error { chains := c.Chains diff --git a/go-gost/x/socket/limiter.go b/go-gost/x/socket/limiter.go index a41db66..61ed9e5 100644 --- a/go-gost/x/socket/limiter.go +++ b/go-gost/x/socket/limiter.go @@ -37,27 +37,30 @@ func updateLimiter(req updateLimiterRequest) error { name := strings.TrimSpace(req.Limiter) - if !registry.TrafficLimiterRegistry().IsRegistered(name) { - return errors.New("limiter " + name + " not found") + if registry.TrafficLimiterRegistry().IsRegistered(name) { + registry.TrafficLimiterRegistry().Unregister(name) } req.Data.Name = name v := parser.ParseTrafficLimiter(&req.Data) - registry.TrafficLimiterRegistry().Unregister(name) - if err := registry.TrafficLimiterRegistry().Register(name, v); err != nil { return errors.New("limiter " + name + " already exists") } config.OnUpdate(func(c *config.Config) error { + found := false for i := range c.Limiters { if c.Limiters[i].Name == name { c.Limiters[i] = &req.Data + found = true break } } + if !found { + c.Limiters = append(c.Limiters, &req.Data) + } return nil }) @@ -68,10 +71,9 @@ func deleteLimiter(req deleteLimiterRequest) error { name := strings.TrimSpace(req.Limiter) - if !registry.TrafficLimiterRegistry().IsRegistered(name) { - return errors.New("limiter " + name + " not found") + if registry.TrafficLimiterRegistry().IsRegistered(name) { + registry.TrafficLimiterRegistry().Unregister(name) } - registry.TrafficLimiterRegistry().Unregister(name) config.OnUpdate(func(c *config.Config) error { limiteres := c.Limiters