mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 07:36:38 +08:00
fix: stabilize federated agent runtime updates (#550)
This commit is contained in:
@@ -3,6 +3,7 @@ package config
|
|||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"io"
|
"io"
|
||||||
|
"reflect"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -30,9 +31,87 @@ func Global() *Config {
|
|||||||
globalMux.RLock()
|
globalMux.RLock()
|
||||||
defer globalMux.RUnlock()
|
defer globalMux.RUnlock()
|
||||||
|
|
||||||
cfg := &Config{}
|
return cloneConfig(global)
|
||||||
*cfg = *global
|
}
|
||||||
return cfg
|
|
||||||
|
// cloneConfig returns a detached snapshot of the runtime config. Runtime
|
||||||
|
// commands mutate slices, pointers and metadata maps, so a shallow struct copy
|
||||||
|
// is not sufficient once readers and persistence run concurrently.
|
||||||
|
func cloneConfig(c *Config) *Config {
|
||||||
|
if c == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
v := cloneConfigValue(reflect.ValueOf(c))
|
||||||
|
if !v.IsValid() || v.IsNil() {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return v.Interface().(*Config)
|
||||||
|
}
|
||||||
|
|
||||||
|
func cloneConfigValue(v reflect.Value) reflect.Value {
|
||||||
|
if !v.IsValid() {
|
||||||
|
return reflect.Value{}
|
||||||
|
}
|
||||||
|
|
||||||
|
switch v.Kind() {
|
||||||
|
case reflect.Interface:
|
||||||
|
if v.IsNil() {
|
||||||
|
return reflect.Zero(v.Type())
|
||||||
|
}
|
||||||
|
cloned := cloneConfigValue(v.Elem())
|
||||||
|
out := reflect.New(v.Type()).Elem()
|
||||||
|
if cloned.IsValid() && cloned.Type().AssignableTo(v.Type()) {
|
||||||
|
out.Set(cloned)
|
||||||
|
} else if cloned.IsValid() && cloned.Type().Implements(v.Type()) {
|
||||||
|
out.Set(cloned)
|
||||||
|
} else if cloned.IsValid() {
|
||||||
|
out.Set(cloned)
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
case reflect.Pointer:
|
||||||
|
if v.IsNil() {
|
||||||
|
return reflect.Zero(v.Type())
|
||||||
|
}
|
||||||
|
out := reflect.New(v.Type().Elem())
|
||||||
|
out.Elem().Set(cloneConfigValue(v.Elem()))
|
||||||
|
return out
|
||||||
|
case reflect.Map:
|
||||||
|
if v.IsNil() {
|
||||||
|
return reflect.Zero(v.Type())
|
||||||
|
}
|
||||||
|
out := reflect.MakeMapWithSize(v.Type(), v.Len())
|
||||||
|
iter := v.MapRange()
|
||||||
|
for iter.Next() {
|
||||||
|
out.SetMapIndex(cloneConfigValue(iter.Key()), cloneConfigValue(iter.Value()))
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
case reflect.Slice:
|
||||||
|
if v.IsNil() {
|
||||||
|
return reflect.Zero(v.Type())
|
||||||
|
}
|
||||||
|
out := reflect.MakeSlice(v.Type(), v.Len(), v.Len())
|
||||||
|
for i := 0; i < v.Len(); i++ {
|
||||||
|
out.Index(i).Set(cloneConfigValue(v.Index(i)))
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
case reflect.Array:
|
||||||
|
out := reflect.New(v.Type()).Elem()
|
||||||
|
for i := 0; i < v.Len(); i++ {
|
||||||
|
out.Index(i).Set(cloneConfigValue(v.Index(i)))
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
case reflect.Struct:
|
||||||
|
out := reflect.New(v.Type()).Elem()
|
||||||
|
out.Set(v)
|
||||||
|
for i := 0; i < v.NumField(); i++ {
|
||||||
|
if out.Field(i).CanSet() && v.Field(i).CanInterface() {
|
||||||
|
out.Field(i).Set(cloneConfigValue(v.Field(i)))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
default:
|
||||||
|
return v
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func Set(c *Config) {
|
func Set(c *Config) {
|
||||||
|
|||||||
@@ -0,0 +1,121 @@
|
|||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"path/filepath"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestGlobalReturnsDetachedSnapshot(t *testing.T) {
|
||||||
|
original := Global()
|
||||||
|
t.Cleanup(func() { Set(original) })
|
||||||
|
|
||||||
|
Set(&Config{Services: []*ServiceConfig{{
|
||||||
|
Name: "snapshot-service",
|
||||||
|
Metadata: map[string]any{"paused": false},
|
||||||
|
Handler: &HandlerConfig{
|
||||||
|
Type: "relay",
|
||||||
|
Metadata: map[string]any{"retries": 2},
|
||||||
|
},
|
||||||
|
}}})
|
||||||
|
|
||||||
|
snapshot := Global()
|
||||||
|
snapshot.Services[0].Name = "changed"
|
||||||
|
snapshot.Services[0].Metadata["paused"] = true
|
||||||
|
snapshot.Services[0].Handler.Metadata["retries"] = 9
|
||||||
|
|
||||||
|
current := Global()
|
||||||
|
if current.Services[0].Name != "snapshot-service" {
|
||||||
|
t.Fatalf("snapshot mutated global service name: %q", current.Services[0].Name)
|
||||||
|
}
|
||||||
|
if paused, _ := current.Services[0].Metadata["paused"].(bool); paused {
|
||||||
|
t.Fatalf("snapshot mutated global service metadata")
|
||||||
|
}
|
||||||
|
if retries, _ := current.Services[0].Handler.Metadata["retries"].(int); retries != 2 {
|
||||||
|
t.Fatalf("snapshot mutated nested handler metadata: %v", current.Services[0].Handler.Metadata["retries"])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConcurrentGlobalSnapshotAndUpdate(t *testing.T) {
|
||||||
|
original := Global()
|
||||||
|
t.Cleanup(func() { Set(original) })
|
||||||
|
|
||||||
|
Set(&Config{Services: []*ServiceConfig{{
|
||||||
|
Name: "concurrent-service",
|
||||||
|
Metadata: map[string]any{"generation": 0},
|
||||||
|
}}})
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for worker := 0; worker < 4; worker++ {
|
||||||
|
wg.Add(1)
|
||||||
|
go func(worker int) {
|
||||||
|
defer wg.Done()
|
||||||
|
for i := 0; i < 500; i++ {
|
||||||
|
if err := OnUpdate(func(c *Config) error {
|
||||||
|
c.Services[0] = &ServiceConfig{
|
||||||
|
Name: "concurrent-service",
|
||||||
|
Metadata: map[string]any{"generation": worker*500 + i},
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}); err != nil {
|
||||||
|
t.Errorf("OnUpdate: %v", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}(worker)
|
||||||
|
}
|
||||||
|
|
||||||
|
for i := 0; i < 2000; i++ {
|
||||||
|
if _, err := json.Marshal(Global()); err != nil {
|
||||||
|
t.Fatalf("marshal snapshot: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConcurrentPersistProducesValidConfig(t *testing.T) {
|
||||||
|
original := Global()
|
||||||
|
originalPath := PersistPath()
|
||||||
|
persistMu.Lock()
|
||||||
|
originalEnabled := persistEnable
|
||||||
|
persistMu.Unlock()
|
||||||
|
t.Cleanup(func() {
|
||||||
|
Set(original)
|
||||||
|
SetPersistPath(originalPath)
|
||||||
|
persistMu.Lock()
|
||||||
|
persistEnable = originalEnabled
|
||||||
|
persistMu.Unlock()
|
||||||
|
})
|
||||||
|
|
||||||
|
path := filepath.Join(t.TempDir(), "gost.json")
|
||||||
|
SetPersistPath(path)
|
||||||
|
EnablePersist()
|
||||||
|
Set(&Config{Services: []*ServiceConfig{{Name: "persist-service"}}})
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for worker := 0; worker < 4; worker++ {
|
||||||
|
wg.Add(1)
|
||||||
|
go func(worker int) {
|
||||||
|
defer wg.Done()
|
||||||
|
for i := 0; i < 50; i++ {
|
||||||
|
if err := OnUpdate(func(c *Config) error {
|
||||||
|
c.Services[0].Metadata = map[string]any{"generation": worker*50 + i}
|
||||||
|
return nil
|
||||||
|
}); err != nil {
|
||||||
|
t.Errorf("persist update: %v", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}(worker)
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
|
||||||
|
var persisted Config
|
||||||
|
if err := persisted.ReadFile(path); err != nil {
|
||||||
|
t.Fatalf("read persisted config: %v", err)
|
||||||
|
}
|
||||||
|
if len(persisted.Services) != 1 || persisted.Services[0] == nil || persisted.Services[0].Name != "persist-service" {
|
||||||
|
t.Fatalf("unexpected persisted services: %#v", persisted.Services)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -42,9 +42,9 @@ func EnablePersist() {
|
|||||||
// persist writes the current global config to the configured file atomically.
|
// persist writes the current global config to the configured file atomically.
|
||||||
func persist() error {
|
func persist() error {
|
||||||
persistMu.Lock()
|
persistMu.Lock()
|
||||||
|
defer persistMu.Unlock()
|
||||||
path := persistPath
|
path := persistPath
|
||||||
enabled := persistEnable
|
enabled := persistEnable
|
||||||
persistMu.Unlock()
|
|
||||||
|
|
||||||
if !enabled || path == "" {
|
if !enabled || path == "" {
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
@@ -0,0 +1,117 @@
|
|||||||
|
package socket
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
coreservice "github.com/go-gost/core/service"
|
||||||
|
"github.com/go-gost/x/config"
|
||||||
|
"github.com/go-gost/x/registry"
|
||||||
|
)
|
||||||
|
|
||||||
|
type blockingCommandService struct {
|
||||||
|
started chan struct{}
|
||||||
|
release chan struct{}
|
||||||
|
startedOnce sync.Once
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *blockingCommandService) Serve() error { return nil }
|
||||||
|
func (s *blockingCommandService) Addr() net.Addr { return nil }
|
||||||
|
func (s *blockingCommandService) Close() error {
|
||||||
|
s.startedOnce.Do(func() { close(s.started) })
|
||||||
|
<-s.release
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMutationCommandsAreSerialized(t *testing.T) {
|
||||||
|
mutations := []string{
|
||||||
|
"AddService", "UpdateService", "DeleteService", "PauseService", "ResumeService",
|
||||||
|
"AddChains", "UpdateChains", "DeleteChains",
|
||||||
|
"AddLimiters", "UpdateLimiters", "DeleteLimiters",
|
||||||
|
"AddCLimiters", "UpdateCLimiters", "DeleteCLimiters",
|
||||||
|
"SetProtocol", "UpgradeAgent", "RollbackAgent", "reload",
|
||||||
|
}
|
||||||
|
for _, command := range mutations {
|
||||||
|
if !isMutationCommand(command) {
|
||||||
|
t.Fatalf("expected %s to use the serialized mutation queue", command)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReadOnlyCommandsRemainBoundedAsync(t *testing.T) {
|
||||||
|
for _, command := range []string{"TcpPing", "UdpPing", "ServiceMonitorCheck"} {
|
||||||
|
if isMutationCommand(command) {
|
||||||
|
t.Fatalf("expected %s to remain a read-only command", command)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCommandResponseTypePreservesRequestContract(t *testing.T) {
|
||||||
|
if got := commandResponseType("UpdateService"); got != "UpdateServiceResponse" {
|
||||||
|
t.Fatalf("unexpected response type: %s", got)
|
||||||
|
}
|
||||||
|
if got := commandResponseType(""); got != "UnknownCommandResponse" {
|
||||||
|
t.Fatalf("unexpected empty command response type: %s", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMutationQueueExecutesCommandsInArrivalOrder(t *testing.T) {
|
||||||
|
originalConfig := config.Global()
|
||||||
|
t.Cleanup(func() { config.Set(originalConfig) })
|
||||||
|
|
||||||
|
firstName := "mutation_queue_first_tdd"
|
||||||
|
secondName := "mutation_queue_second_tdd"
|
||||||
|
first := &blockingCommandService{started: make(chan struct{}), release: make(chan struct{})}
|
||||||
|
second := &blockingCommandService{started: make(chan struct{}), release: make(chan struct{})}
|
||||||
|
for _, name := range []string{firstName, secondName} {
|
||||||
|
registry.ServiceRegistry().Unregister(name)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() {
|
||||||
|
select {
|
||||||
|
case <-first.release:
|
||||||
|
default:
|
||||||
|
close(first.release)
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case <-second.release:
|
||||||
|
default:
|
||||||
|
close(second.release)
|
||||||
|
}
|
||||||
|
registry.ServiceRegistry().Unregister(firstName)
|
||||||
|
registry.ServiceRegistry().Unregister(secondName)
|
||||||
|
})
|
||||||
|
if err := registry.ServiceRegistry().Register(firstName, coreservice.Service(first)); err != nil {
|
||||||
|
t.Fatalf("register first service: %v", err)
|
||||||
|
}
|
||||||
|
if err := registry.ServiceRegistry().Register(secondName, coreservice.Service(second)); err != nil {
|
||||||
|
t.Fatalf("register second service: %v", err)
|
||||||
|
}
|
||||||
|
config.Set(&config.Config{Services: []*config.ServiceConfig{{Name: firstName}, {Name: secondName}}})
|
||||||
|
|
||||||
|
reporter := NewWebSocketReporter("", "mutation-queue-test-secret")
|
||||||
|
go reporter.runMutationCommands()
|
||||||
|
t.Cleanup(reporter.Stop)
|
||||||
|
reporter.dispatchCommand(CommandMessage{Type: "DeleteService", Data: map[string]any{"services": []string{firstName}}})
|
||||||
|
reporter.dispatchCommand(CommandMessage{Type: "DeleteService", Data: map[string]any{"services": []string{secondName}}})
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-first.started:
|
||||||
|
case <-time.After(time.Second):
|
||||||
|
t.Fatal("first mutation did not start")
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case <-second.started:
|
||||||
|
t.Fatal("second mutation started before first mutation completed")
|
||||||
|
case <-time.After(100 * time.Millisecond):
|
||||||
|
}
|
||||||
|
|
||||||
|
close(first.release)
|
||||||
|
select {
|
||||||
|
case <-second.started:
|
||||||
|
case <-time.After(time.Second):
|
||||||
|
t.Fatal("second mutation did not start after first mutation completed")
|
||||||
|
}
|
||||||
|
close(second.release)
|
||||||
|
}
|
||||||
+82
-18
@@ -15,6 +15,11 @@ import (
|
|||||||
xservice "github.com/go-gost/x/service"
|
xservice "github.com/go-gost/x/service"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
type serviceReplacement struct {
|
||||||
|
config config.ServiceConfig
|
||||||
|
oldConfig *config.ServiceConfig
|
||||||
|
}
|
||||||
|
|
||||||
func createServices(req createServicesRequest) error {
|
func createServices(req createServicesRequest) error {
|
||||||
|
|
||||||
if len(req.Data) == 0 {
|
if len(req.Data) == 0 {
|
||||||
@@ -95,11 +100,11 @@ func updateServices(req updateServicesRequest) error {
|
|||||||
req.Data[i].Name = name
|
req.Data[i].Name = name
|
||||||
}
|
}
|
||||||
|
|
||||||
// 第二阶段:逐个更新服务(Upsert模式:存在则更新,不存在则创建)
|
// 第二阶段:逐个更新服务(Upsert模式:存在则更新,不存在则创建)。
|
||||||
changedServices := make([]struct {
|
// 配置变更命令由 WebSocket reporter 串行调度,但这里仍保留完整回滚,
|
||||||
config config.ServiceConfig
|
// 避免新配置解析或监听失败后旧服务永久消失。
|
||||||
service service.Service
|
originalConfig := config.Global()
|
||||||
}, 0, len(req.Data))
|
changedServices := make([]serviceReplacement, 0, len(req.Data))
|
||||||
for i := range req.Data {
|
for i := range req.Data {
|
||||||
serviceConfig := &req.Data[i]
|
serviceConfig := &req.Data[i]
|
||||||
name := serviceConfig.Name
|
name := serviceConfig.Name
|
||||||
@@ -107,33 +112,48 @@ func updateServices(req updateServicesRequest) error {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
// 1. 获取旧服务
|
var oldConfig *config.ServiceConfig
|
||||||
old := registry.ServiceRegistry().Get(name)
|
if originalConfig != nil {
|
||||||
|
for _, current := range originalConfig.Services {
|
||||||
|
if current != nil && strings.TrimSpace(current.Name) == name {
|
||||||
|
oldConfig = current
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// 2. 关闭旧服务 (如果存在)
|
// 1. 关闭并移除旧服务(如果存在)。同名监听必须先释放端口,
|
||||||
if old != nil {
|
// 才能创建新 listener。
|
||||||
// 3. 从注册表移除旧服务;registry 会负责关闭旧服务。
|
if registry.ServiceRegistry().Get(name) != nil {
|
||||||
registry.ServiceRegistry().Unregister(name)
|
registry.ServiceRegistry().Unregister(name)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 4. 解析新服务配置
|
// 2. 解析新服务配置。
|
||||||
svc, err := parser.ParseService(serviceConfig)
|
svc, err := parser.ParseService(serviceConfig)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
rollbackErr := restoreServiceRuntime(name, oldConfig)
|
||||||
|
rollbackErr = errors.Join(rollbackErr, rollbackServiceReplacements(changedServices))
|
||||||
|
if rollbackErr != nil {
|
||||||
|
return fmt.Errorf("create service %s failed: %v; restore previous service failed: %w", name, err, rollbackErr)
|
||||||
|
}
|
||||||
return errors.New("create service " + name + " failed: " + err.Error())
|
return errors.New("create service " + name + " failed: " + err.Error())
|
||||||
}
|
}
|
||||||
changedServices = append(changedServices, struct {
|
|
||||||
config config.ServiceConfig
|
|
||||||
service service.Service
|
|
||||||
}{*serviceConfig, svc})
|
|
||||||
|
|
||||||
// 5. 注册新服务
|
// 3. 注册并启动新服务。
|
||||||
if err := registry.ServiceRegistry().Register(name, svc); err != nil {
|
if err := registry.ServiceRegistry().Register(name, svc); err != nil {
|
||||||
svc.Close()
|
svc.Close()
|
||||||
|
rollbackErr := restoreServiceRuntime(name, oldConfig)
|
||||||
|
rollbackErr = errors.Join(rollbackErr, rollbackServiceReplacements(changedServices))
|
||||||
|
if rollbackErr != nil {
|
||||||
|
return fmt.Errorf("service %s already exists; restore previous service failed: %w", name, rollbackErr)
|
||||||
|
}
|
||||||
return errors.New("service " + name + " already exists")
|
return errors.New("service " + name + " already exists")
|
||||||
}
|
}
|
||||||
|
|
||||||
// 6. 启动新服务
|
|
||||||
go svc.Serve()
|
go svc.Serve()
|
||||||
|
changedServices = append(changedServices, serviceReplacement{
|
||||||
|
config: *serviceConfig,
|
||||||
|
oldConfig: oldConfig,
|
||||||
|
})
|
||||||
}
|
}
|
||||||
if len(changedServices) == 0 {
|
if len(changedServices) == 0 {
|
||||||
return nil
|
return nil
|
||||||
@@ -158,12 +178,56 @@ func updateServices(req updateServicesRequest) error {
|
|||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
|
config.Set(originalConfig)
|
||||||
|
if rollbackErr := rollbackServiceReplacements(changedServices); rollbackErr != nil {
|
||||||
|
return fmt.Errorf("%w; restore previous services failed: %v", err, rollbackErr)
|
||||||
|
}
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func restoreServiceRuntime(name string, serviceConfig *config.ServiceConfig) error {
|
||||||
|
name = strings.TrimSpace(name)
|
||||||
|
if name == "" || serviceConfig == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if registry.ServiceRegistry().Get(name) != nil {
|
||||||
|
registry.ServiceRegistry().Unregister(name)
|
||||||
|
}
|
||||||
|
|
||||||
|
cfgCopy := *serviceConfig
|
||||||
|
cfgCopy.Name = name
|
||||||
|
svc, err := parser.ParseService(&cfgCopy)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := registry.ServiceRegistry().Register(name, svc); err != nil {
|
||||||
|
svc.Close()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
go svc.Serve()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func rollbackServiceReplacements(replacements []serviceReplacement) error {
|
||||||
|
var rollbackErr error
|
||||||
|
for i := len(replacements) - 1; i >= 0; i-- {
|
||||||
|
name := strings.TrimSpace(replacements[i].config.Name)
|
||||||
|
if name == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if registry.ServiceRegistry().Get(name) != nil {
|
||||||
|
registry.ServiceRegistry().Unregister(name)
|
||||||
|
}
|
||||||
|
if err := restoreServiceRuntime(name, replacements[i].oldConfig); err != nil {
|
||||||
|
rollbackErr = errors.Join(rollbackErr, fmt.Errorf("restore service %s: %w", name, err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return rollbackErr
|
||||||
|
}
|
||||||
|
|
||||||
func serviceConfigUnchanged(name string, next config.ServiceConfig) bool {
|
func serviceConfigUnchanged(name string, next config.ServiceConfig) bool {
|
||||||
cfg := config.Global()
|
cfg := config.Global()
|
||||||
if cfg == nil {
|
if cfg == nil {
|
||||||
|
|||||||
@@ -7,6 +7,8 @@ import (
|
|||||||
corelogger "github.com/go-gost/core/logger"
|
corelogger "github.com/go-gost/core/logger"
|
||||||
"github.com/go-gost/core/service"
|
"github.com/go-gost/core/service"
|
||||||
"github.com/go-gost/x/config"
|
"github.com/go-gost/x/config"
|
||||||
|
_ "github.com/go-gost/x/handler/auto"
|
||||||
|
_ "github.com/go-gost/x/listener/tcp"
|
||||||
xlogger "github.com/go-gost/x/logger"
|
xlogger "github.com/go-gost/x/logger"
|
||||||
"github.com/go-gost/x/registry"
|
"github.com/go-gost/x/registry"
|
||||||
)
|
)
|
||||||
@@ -15,6 +17,43 @@ type recordingService struct {
|
|||||||
closed int
|
closed int
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestUpdateServicesParseFailureRestoresPreviousRuntime(t *testing.T) {
|
||||||
|
corelogger.SetDefault(xlogger.Nop())
|
||||||
|
|
||||||
|
name := "restore_after_failed_update_tdd"
|
||||||
|
existing := &recordingService{}
|
||||||
|
registry.ServiceRegistry().Unregister(name)
|
||||||
|
t.Cleanup(func() { registry.ServiceRegistry().Unregister(name) })
|
||||||
|
if err := registry.ServiceRegistry().Register(name, service.Service(existing)); err != nil {
|
||||||
|
t.Fatalf("register existing service: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
originalConfig := config.Global()
|
||||||
|
t.Cleanup(func() { config.Set(originalConfig) })
|
||||||
|
serviceConfig := config.ServiceConfig{Name: name, Addr: "127.0.0.1:0"}
|
||||||
|
config.Set(&config.Config{Services: []*config.ServiceConfig{&serviceConfig}})
|
||||||
|
|
||||||
|
invalid := serviceConfig
|
||||||
|
invalid.Listener = &config.ListenerConfig{Type: "listener-does-not-exist"}
|
||||||
|
if err := updateServices(updateServicesRequest{Data: []config.ServiceConfig{invalid}}); err == nil {
|
||||||
|
t.Fatalf("expected invalid service update to fail")
|
||||||
|
}
|
||||||
|
if existing.closed != 1 {
|
||||||
|
t.Fatalf("expected old runtime to be closed once, got %d", existing.closed)
|
||||||
|
}
|
||||||
|
if registry.ServiceRegistry().Get(name) == nil {
|
||||||
|
t.Fatalf("expected previous runtime to be restored")
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg := config.Global()
|
||||||
|
if len(cfg.Services) != 1 || cfg.Services[0] == nil || cfg.Services[0].Name != name {
|
||||||
|
t.Fatalf("expected previous config to remain, got %#v", cfg.Services)
|
||||||
|
}
|
||||||
|
if cfg.Services[0].Listener != nil {
|
||||||
|
t.Fatalf("expected previous listener config to remain unchanged")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (s *recordingService) Serve() error { return nil }
|
func (s *recordingService) Serve() error { return nil }
|
||||||
func (s *recordingService) Addr() net.Addr { return nil }
|
func (s *recordingService) Addr() net.Addr { return nil }
|
||||||
func (s *recordingService) Close() error {
|
func (s *recordingService) Close() error {
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ import (
|
|||||||
"os"
|
"os"
|
||||||
"os/exec"
|
"os/exec"
|
||||||
"runtime"
|
"runtime"
|
||||||
|
"runtime/debug"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
@@ -152,6 +153,8 @@ const (
|
|||||||
maxBackoff = 2 * time.Minute // 重连最大退避
|
maxBackoff = 2 * time.Minute // 重连最大退避
|
||||||
defaultMetricReportInterval = 5 * time.Second
|
defaultMetricReportInterval = 5 * time.Second
|
||||||
maxConcurrentTCPPings = 8
|
maxConcurrentTCPPings = 8
|
||||||
|
maxConcurrentReadCommands = 16
|
||||||
|
maxQueuedMutationCommands = 256
|
||||||
)
|
)
|
||||||
|
|
||||||
type WebSocketReporter struct {
|
type WebSocketReporter struct {
|
||||||
@@ -174,6 +177,8 @@ type WebSocketReporter struct {
|
|||||||
connMutex sync.Mutex // 连接状态锁
|
connMutex sync.Mutex // 连接状态锁
|
||||||
aesCrypto *crypto.AESCrypto // AES加密器
|
aesCrypto *crypto.AESCrypto // AES加密器
|
||||||
tcpPingSem chan struct{} // 限制诊断探测并发,避免离线目标耗尽连接
|
tcpPingSem chan struct{} // 限制诊断探测并发,避免离线目标耗尽连接
|
||||||
|
readCommandSem chan struct{} // 限制只读命令并发,避免诊断请求耗尽资源
|
||||||
|
mutationQueue chan CommandMessage
|
||||||
}
|
}
|
||||||
|
|
||||||
var wsDial = func(dialer *websocket.Dialer, rawURL string) (*websocket.Conn, *http.Response, error) {
|
var wsDial = func(dialer *websocket.Dialer, rawURL string) (*websocket.Conn, *http.Response, error) {
|
||||||
@@ -204,6 +209,8 @@ func NewWebSocketReporter(serverURL string, secret string) *WebSocketReporter {
|
|||||||
connecting: false,
|
connecting: false,
|
||||||
aesCrypto: aesCrypto,
|
aesCrypto: aesCrypto,
|
||||||
tcpPingSem: make(chan struct{}, maxConcurrentTCPPings),
|
tcpPingSem: make(chan struct{}, maxConcurrentTCPPings),
|
||||||
|
readCommandSem: make(chan struct{}, maxConcurrentReadCommands),
|
||||||
|
mutationQueue: make(chan CommandMessage, maxQueuedMutationCommands),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -231,6 +238,7 @@ func (w *WebSocketReporter) releaseTCPPingSlot() {
|
|||||||
|
|
||||||
// Start 启动WebSocket报告器
|
// Start 启动WebSocket报告器
|
||||||
func (w *WebSocketReporter) Start() {
|
func (w *WebSocketReporter) Start() {
|
||||||
|
go w.runMutationCommands()
|
||||||
go w.run()
|
go w.run()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -777,8 +785,7 @@ func (w *WebSocketReporter) handleReceivedMessage(messageType int, message []byt
|
|||||||
}
|
}
|
||||||
|
|
||||||
if cmdMsg.Type != "call" {
|
if cmdMsg.Type != "call" {
|
||||||
// 所有命令统一异步执行,避免阻塞消息接收循环
|
w.dispatchCommand(cmdMsg)
|
||||||
go w.routeCommand(cmdMsg)
|
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
// 处理普通消息
|
// 处理普通消息
|
||||||
@@ -789,8 +796,7 @@ func (w *WebSocketReporter) handleReceivedMessage(messageType int, message []byt
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
if cmdMsg.Type != "call" {
|
if cmdMsg.Type != "call" {
|
||||||
// 所有命令统一异步执行,避免阻塞消息接收循环
|
w.dispatchCommand(cmdMsg)
|
||||||
go w.routeCommand(cmdMsg)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -799,6 +805,86 @@ func (w *WebSocketReporter) handleReceivedMessage(messageType int, message []byt
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// dispatchCommand keeps all runtime mutations ordered while allowing bounded
|
||||||
|
// concurrency for read-only diagnostics. Mutations share process-wide
|
||||||
|
// registries and configuration, so running them concurrently can corrupt the
|
||||||
|
// persisted config or interleave service lifecycle operations.
|
||||||
|
func (w *WebSocketReporter) dispatchCommand(cmd CommandMessage) {
|
||||||
|
if isMutationCommand(cmd.Type) {
|
||||||
|
select {
|
||||||
|
case w.mutationQueue <- cmd:
|
||||||
|
case <-w.ctx.Done():
|
||||||
|
w.sendCommandFailure(cmd, "Agent is shutting down")
|
||||||
|
default:
|
||||||
|
w.sendCommandFailure(cmd, "运行时配置命令队列已满,请稍后重试")
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case w.readCommandSem <- struct{}{}:
|
||||||
|
go func() {
|
||||||
|
defer func() { <-w.readCommandSem }()
|
||||||
|
w.routeCommandSafely(cmd)
|
||||||
|
}()
|
||||||
|
case <-w.ctx.Done():
|
||||||
|
w.sendCommandFailure(cmd, "Agent is shutting down")
|
||||||
|
default:
|
||||||
|
w.sendCommandFailure(cmd, "只读命令并发过多,请稍后重试")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *WebSocketReporter) runMutationCommands() {
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-w.ctx.Done():
|
||||||
|
return
|
||||||
|
case cmd := <-w.mutationQueue:
|
||||||
|
w.routeCommandSafely(cmd)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *WebSocketReporter) routeCommandSafely(cmd CommandMessage) {
|
||||||
|
defer func() {
|
||||||
|
if recovered := recover(); recovered != nil {
|
||||||
|
fmt.Printf("❌ 命令处理 panic: type=%s panic=%v\n%s", cmd.Type, recovered, debug.Stack())
|
||||||
|
w.sendCommandFailure(cmd, fmt.Sprintf("命令处理异常: %v", recovered))
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
w.routeCommand(cmd)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *WebSocketReporter) sendCommandFailure(cmd CommandMessage, message string) {
|
||||||
|
w.sendResponse(CommandResponse{
|
||||||
|
Type: commandResponseType(cmd.Type),
|
||||||
|
Success: false,
|
||||||
|
Message: message,
|
||||||
|
RequestId: cmd.RequestId,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func commandResponseType(commandType string) string {
|
||||||
|
commandType = strings.TrimSpace(commandType)
|
||||||
|
if commandType == "" {
|
||||||
|
return "UnknownCommandResponse"
|
||||||
|
}
|
||||||
|
return commandType + "Response"
|
||||||
|
}
|
||||||
|
|
||||||
|
func isMutationCommand(commandType string) bool {
|
||||||
|
switch strings.ToLower(strings.TrimSpace(commandType)) {
|
||||||
|
case "addservice", "updateservice", "deleteservice", "pauseservice", "resumeservice",
|
||||||
|
"addchains", "updatechains", "deletechains",
|
||||||
|
"addlimiters", "updatelimiters", "deletelimiters",
|
||||||
|
"addclimiters", "updateclimiters", "deleteclimiters",
|
||||||
|
"setprotocol", "upgradeagent", "rollbackagent", "reload":
|
||||||
|
return true
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// routeCommand 路由命令到对应的处理函数
|
// routeCommand 路由命令到对应的处理函数
|
||||||
func (w *WebSocketReporter) routeCommand(cmd CommandMessage) {
|
func (w *WebSocketReporter) routeCommand(cmd CommandMessage) {
|
||||||
jsonBytes, errs := json.Marshal(cmd)
|
jsonBytes, errs := json.Marshal(cmd)
|
||||||
|
|||||||
@@ -405,6 +405,7 @@ After=network.target
|
|||||||
[Service]
|
[Service]
|
||||||
WorkingDirectory=$INSTALL_DIR
|
WorkingDirectory=$INSTALL_DIR
|
||||||
ExecStart=$INSTALL_DIR/flux_agent
|
ExecStart=$INSTALL_DIR/flux_agent
|
||||||
|
Environment=GODEBUG=disablethp=1
|
||||||
Restart=on-failure
|
Restart=on-failure
|
||||||
StandardOutput=null
|
StandardOutput=null
|
||||||
StandardError=null
|
StandardError=null
|
||||||
|
|||||||
@@ -277,6 +277,8 @@ EOF
|
|||||||
local expected=$'{\n "addr": "panel\\"addr",\n "secret": "sec\\\\ret\\"1"\n}'
|
local expected=$'{\n "addr": "panel\\"addr",\n "secret": "sec\\\\ret\\"1"\n}'
|
||||||
|
|
||||||
assert_equals "$expected" "$actual" "install_flux_agent should JSON-escape config values"
|
assert_equals "$expected" "$actual" "install_flux_agent should JSON-escape config values"
|
||||||
|
grep -Fqx 'Environment=GODEBUG=disablethp=1' "$FLUX_AGENT_SYSTEMD_SERVICE_FILE" || \
|
||||||
|
fail "systemd service should disable transparent huge pages for the Go heap"
|
||||||
)
|
)
|
||||||
|
|
||||||
test_install_script_bootstraps_bash_for_alpine() (
|
test_install_script_bootstraps_bash_for_alpine() (
|
||||||
|
|||||||
Reference in New Issue
Block a user