Files
flvx/go-gost/x/socket/service.go
T

634 lines
16 KiB
Go

package socket
import (
"errors"
"fmt"
"reflect"
"strings"
"time"
"github.com/go-gost/core/service"
"github.com/go-gost/x/config"
parser "github.com/go-gost/x/config/parsing/service"
kill "github.com/go-gost/x/internal/util/port"
"github.com/go-gost/x/registry"
xservice "github.com/go-gost/x/service"
)
type serviceReplacement struct {
config config.ServiceConfig
oldConfig *config.ServiceConfig
}
func createServices(req createServicesRequest) error {
if len(req.Data) == 0 {
return errors.New("services list cannot be empty")
}
// 第一阶段:验证所有服务配置
var parsedServices []struct {
config config.ServiceConfig
service service.Service
}
for _, serviceConfig := range req.Data {
name := strings.TrimSpace(serviceConfig.Name)
if name == "" {
return errors.New("service name is required")
}
serviceConfig.Name = name
if registry.ServiceRegistry().IsRegistered(name) {
return errors.New("service " + name + " already exists")
}
svc, err := parser.ParseService(&serviceConfig)
if err != nil {
return errors.New("create service " + name + " failed: " + err.Error())
}
parsedServices = append(parsedServices, struct {
config config.ServiceConfig
service service.Service
}{serviceConfig, svc})
}
// 第二阶段:注册所有服务
var registeredServices []string
for _, ps := range parsedServices {
if err := registry.ServiceRegistry().Register(ps.config.Name, ps.service); err != nil {
// 如果注册失败,回滚已注册的服务
for _, regName := range registeredServices {
if registry.ServiceRegistry().Get(regName) != nil {
registry.ServiceRegistry().Unregister(regName)
}
}
return errors.New("service " + ps.config.Name + " already exists")
}
registeredServices = append(registeredServices, ps.config.Name)
}
// 第三阶段:启动所有服务
for _, ps := range parsedServices {
if svc := registry.ServiceRegistry().Get(ps.config.Name); svc != nil {
go svc.Serve()
}
}
// 第四阶段:更新配置
return config.OnUpdate(func(c *config.Config) error {
for _, ps := range parsedServices {
c.Services = append(c.Services, &ps.config)
}
return nil
})
}
func updateServices(req updateServicesRequest) error {
if len(req.Data) == 0 {
return errors.New("services list cannot be empty")
}
// 第一阶段:验证所有服务名称有效性
for i := range req.Data {
name := strings.TrimSpace(req.Data[i].Name)
if name == "" {
return errors.New("service name is required")
}
req.Data[i].Name = name
}
// 第二阶段:逐个更新服务(Upsert模式:存在则更新,不存在则创建)。
// 配置变更命令由 WebSocket reporter 串行调度,但这里仍保留完整回滚,
// 避免新配置解析或监听失败后旧服务永久消失。
originalConfig := config.Global()
changedServices := make([]serviceReplacement, 0, len(req.Data))
for i := range req.Data {
serviceConfig := &req.Data[i]
name := serviceConfig.Name
if registry.ServiceRegistry().Get(name) != nil && serviceConfigUnchanged(name, *serviceConfig) {
continue
}
var oldConfig *config.ServiceConfig
if originalConfig != nil {
for _, current := range originalConfig.Services {
if current != nil && strings.TrimSpace(current.Name) == name {
oldConfig = current
break
}
}
}
// 1. 关闭并移除旧服务(如果存在)。同名监听必须先释放端口,
// 才能创建新 listener。
if registry.ServiceRegistry().Get(name) != nil {
registry.ServiceRegistry().Unregister(name)
}
// 2. 解析新服务配置。
svc, err := parser.ParseService(serviceConfig)
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())
}
// 3. 注册并启动新服务。
if err := registry.ServiceRegistry().Register(name, svc); err != nil {
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")
}
go svc.Serve()
changedServices = append(changedServices, serviceReplacement{
config: *serviceConfig,
oldConfig: oldConfig,
})
}
if len(changedServices) == 0 {
return nil
}
// 第三阶段:更新配置
if err := config.OnUpdate(func(c *config.Config) error {
for i := range changedServices {
// 创建副本以确保指针安全
cfgCopy := changedServices[i].config
found := false
for j := range c.Services {
if c.Services[j].Name == cfgCopy.Name {
c.Services[j] = &cfgCopy
found = true
break
}
}
if !found {
c.Services = append(c.Services, &cfgCopy)
}
}
return 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 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 {
cfg := config.Global()
if cfg == nil {
return false
}
next.Status = nil
for _, current := range cfg.Services {
if current == nil || strings.TrimSpace(current.Name) != name {
continue
}
currentCopy := *current
currentCopy.Status = nil
return reflect.DeepEqual(currentCopy, next)
}
return false
}
func deleteServices(req deleteServicesRequest) error {
if len(req.Services) == 0 {
return errors.New("services list cannot be empty")
}
// 第一阶段:验证所有服务是否存在
var servicesToDelete []struct {
name string
service service.Service
}
var namesToRemove []string
for _, serviceName := range req.Services {
name := strings.TrimSpace(serviceName)
if name == "" {
return errors.New("service name is required")
}
namesToRemove = append(namesToRemove, name)
svc := registry.ServiceRegistry().Get(name)
if svc != nil {
servicesToDelete = append(servicesToDelete, struct {
name string
service service.Service
}{name, svc})
}
}
// 第二阶段:删除所有服务
for _, std := range servicesToDelete {
registry.ServiceRegistry().Unregister(std.name)
}
// 确保所有请求删除的服务都从注册表中移除(即使之前未找到实例)
for _, name := range namesToRemove {
if registry.ServiceRegistry().IsRegistered(name) {
registry.ServiceRegistry().Unregister(name)
}
}
// 第三阶段:更新配置
err := config.OnUpdate(func(c *config.Config) error {
services := c.Services
c.Services = nil
for _, s := range services {
shouldDelete := false
for _, name := range namesToRemove {
if s.Name == name {
shouldDelete = true
break
}
}
if !shouldDelete {
c.Services = append(c.Services, s)
}
}
return nil
})
xservice.GetGlobalTrafficManager().RemoveServices(namesToRemove...)
return err
}
func pauseServices(req pauseServicesRequest) error {
if len(req.Services) == 0 {
return errors.New("services list cannot be empty")
}
// 第一阶段:验证所有服务是否存在,并筛选需要暂停的服务
var servicesToPause []struct {
name string
service service.Service
}
//var skippedServices []string
cfg := config.Global()
for _, serviceName := range req.Services {
name := strings.TrimSpace(serviceName)
if name == "" {
return errors.New("service name is required")
}
svc := registry.ServiceRegistry().Get(name)
if svc == nil {
return errors.New(fmt.Sprintf("service %s not found", name))
}
//// 检查服务是否已经暂停
//var serviceConfig *config.ServiceConfig
//for _, s := range cfg.Services {
// if s.Name == name {
// serviceConfig = s
// break
// }
//}
//
//// 如果服务已经暂停,跳过
//if serviceConfig != nil && serviceConfig.Metadata != nil {
// if pausedVal, exists := serviceConfig.Metadata["paused"]; exists && pausedVal == true {
// skippedServices = append(skippedServices, name)
// continue
// }
//}
servicesToPause = append(servicesToPause, struct {
name string
service service.Service
}{name, svc})
}
// 第二阶段:事务性暂停所有服务
var pausedServices []struct {
name string
service service.Service
serviceConfig *config.ServiceConfig
}
// 获取服务配置
serviceConfigs := make(map[string]*config.ServiceConfig)
for _, s := range cfg.Services {
serviceConfigs[s.Name] = s
}
// 逐个暂停服务,如果失败则回滚
for _, stp := range servicesToPause {
serviceConfig := serviceConfigs[stp.name]
if serviceConfig == nil {
// 找不到配置,回滚已暂停的服务
rollbackPausedServices(pausedServices)
return errors.New(fmt.Sprintf("service %s configuration not found", stp.name))
}
// 暂停服务
stp.service.Close()
// 强制断开端口的所有连接
if serviceConfig.Addr != "" {
_ = kill.ForceClosePortConnections(serviceConfig.Addr)
}
// 记录已暂停的服务
pausedServices = append(pausedServices, struct {
name string
service service.Service
serviceConfig *config.ServiceConfig
}{stp.name, stp.service, serviceConfig})
}
// 第三阶段:更新配置,标记暂停状态
err := config.OnUpdate(func(c *config.Config) error {
for _, stp := range servicesToPause {
for i := range c.Services {
if c.Services[i].Name == stp.name {
if c.Services[i].Metadata == nil {
c.Services[i].Metadata = make(map[string]any)
}
c.Services[i].Metadata["paused"] = true
break
}
}
}
return nil
})
if err != nil {
// 配置更新失败,需要回滚所有暂停的服务
rollbackPausedServices(pausedServices)
return errors.New(fmt.Sprintf("Failed to update config, rolling back paused services: %v", err))
}
return nil
}
func resumeServices(req resumeServicesRequest) error {
if len(req.Services) == 0 {
return errors.New("services list cannot be empty")
}
// 第一阶段:验证所有服务是否存在,并筛选需要恢复的服务
var servicesToResume []struct {
name string
service service.Service
serviceConfig *config.ServiceConfig
}
var skippedServices []string
cfg := config.Global()
for _, serviceName := range req.Services {
name := strings.TrimSpace(serviceName)
if name == "" {
return errors.New("service name is required")
}
// 检查服务是否存在
svc := registry.ServiceRegistry().Get(name)
if svc == nil {
return errors.New(fmt.Sprintf("service %s not found", name))
}
// 查找配置中的服务
var serviceConfig *config.ServiceConfig
for _, s := range cfg.Services {
if s.Name == name {
serviceConfig = s
break
}
}
if serviceConfig == nil {
return errors.New(fmt.Sprintf("service %s configuration not found", name))
}
// 检查是否处于暂停状态
paused := false
if serviceConfig.Metadata != nil {
if pausedVal, exists := serviceConfig.Metadata["paused"]; exists && pausedVal == true {
paused = true
}
}
// 如果服务没有暂停(即正在运行),跳过
if !paused {
skippedServices = append(skippedServices, name)
continue
}
servicesToResume = append(servicesToResume, struct {
name string
service service.Service
serviceConfig *config.ServiceConfig
}{name, svc, serviceConfig})
}
// 第二阶段:事务性恢复所有服务
var resumedServices []struct {
name string
service service.Service
serviceConfig *config.ServiceConfig
}
// 逐个恢复服务,如果失败则回滚
for _, str := range servicesToResume {
// 先关闭现有服务
str.service.Close()
registry.ServiceRegistry().Unregister(str.name)
// 强制断开端口的所有连接
if str.serviceConfig.Addr != "" {
_ = kill.ForceClosePortConnections(str.serviceConfig.Addr)
}
// 等待端口释放
time.Sleep(500 * time.Millisecond)
// 重新解析并启动服务
svc, err := parser.ParseService(str.serviceConfig)
if err != nil {
// 恢复失败,回滚已恢复的服务
rollbackResumedServices(resumedServices)
return errors.New(fmt.Sprintf("resume service %s failed: %s", str.name, err.Error()))
}
if err := registry.ServiceRegistry().Register(str.name, svc); err != nil {
svc.Close()
// 恢复失败,回滚已恢复的服务
rollbackResumedServices(resumedServices)
return errors.New(fmt.Sprintf("service %s already exists", str.name))
}
go svc.Serve()
// 记录已成功恢复的服务
resumedServices = append(resumedServices, str)
}
// 第三阶段:更新配置,移除暂停状态
err := config.OnUpdate(func(c *config.Config) error {
for _, str := range servicesToResume {
for i := range c.Services {
if c.Services[i].Name == str.name {
if c.Services[i].Metadata != nil {
delete(c.Services[i].Metadata, "paused")
// 如果 metadata 为空,设置为 nil
if len(c.Services[i].Metadata) == 0 {
c.Services[i].Metadata = nil
}
}
break
}
}
}
return nil
})
if err != nil {
// 配置更新失败,回滚所有已恢复的服务
rollbackResumedServices(resumedServices)
return errors.New(fmt.Sprintf("Failed to update config, rolling back resumed services: %v", err))
}
return nil
}
func rollbackPausedServices(pausedServices []struct {
name string
service service.Service
serviceConfig *config.ServiceConfig
}) {
for _, pss := range pausedServices {
// 重新解析并启动服务
svc, err := parser.ParseService(pss.serviceConfig)
if err != nil {
continue // 回滚失败,记录日志但继续处理其他服务
}
if err := registry.ServiceRegistry().Register(pss.name, svc); err != nil {
svc.Close()
continue // 回滚失败,记录日志但继续处理其他服务
}
go svc.Serve()
// 移除暂停状态标记
config.OnUpdate(func(c *config.Config) error {
for i := range c.Services {
if c.Services[i].Name == pss.name {
if c.Services[i].Metadata != nil {
delete(c.Services[i].Metadata, "paused")
if len(c.Services[i].Metadata) == 0 {
c.Services[i].Metadata = nil
}
}
break
}
}
return nil
})
}
}
func rollbackResumedServices(resumedServices []struct {
name string
service service.Service
serviceConfig *config.ServiceConfig
}) {
for _, rss := range resumedServices {
// 关闭已恢复的服务
if svc := registry.ServiceRegistry().Get(rss.name); svc != nil {
svc.Close()
}
// 重新标记为暂停状态
config.OnUpdate(func(c *config.Config) error {
for i := range c.Services {
if c.Services[i].Name == rss.name {
if c.Services[i].Metadata == nil {
c.Services[i].Metadata = make(map[string]any)
}
c.Services[i].Metadata["paused"] = true
break
}
}
return nil
})
}
}
type resumeServicesRequest struct {
Services []string `json:"services"`
}
type pauseServicesRequest struct {
Services []string `json:"services"`
}
type deleteServicesRequest struct {
Services []string `json:"services"`
}
type updateServicesRequest struct {
Data []config.ServiceConfig `json:"data"`
}
type createServicesRequest struct {
Data []config.ServiceConfig `json:"data"`
}