mirror of
https://github.com/shtorm-7/sing-box-extended.git
synced 2026-09-15 21:00:27 +00:00
daemon: Fix start or reload race
This commit is contained in:
+13
-4
@@ -1,6 +1,7 @@
|
||||
package adapter
|
||||
|
||||
import (
|
||||
"context"
|
||||
"reflect"
|
||||
"strings"
|
||||
"time"
|
||||
@@ -74,11 +75,15 @@ func getServiceName(service any) string {
|
||||
return strings.ToLower(t.Name())
|
||||
}
|
||||
|
||||
func Start(logger log.ContextLogger, stage StartStage, services ...Lifecycle) error {
|
||||
func Start(ctx context.Context, logger log.ContextLogger, stage StartStage, services ...Lifecycle) error {
|
||||
for _, service := range services {
|
||||
err := ctx.Err()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
name := getServiceName(service)
|
||||
done := LogElapsed(logger, stage, " ", name)
|
||||
err := service.Start(stage)
|
||||
err = service.Start(stage)
|
||||
done()
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -87,10 +92,14 @@ func Start(logger log.ContextLogger, stage StartStage, services ...Lifecycle) er
|
||||
return nil
|
||||
}
|
||||
|
||||
func StartNamed(logger log.ContextLogger, stage StartStage, services []LifecycleService) error {
|
||||
func StartNamed(ctx context.Context, logger log.ContextLogger, stage StartStage, services []LifecycleService) error {
|
||||
for _, service := range services {
|
||||
err := ctx.Err()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
done := LogElapsed(logger, stage, " ", service.Name())
|
||||
err := service.Start(stage)
|
||||
err = service.Start(stage)
|
||||
done()
|
||||
if err != nil {
|
||||
return E.Cause(err, stage.String(), " ", service.Name())
|
||||
|
||||
@@ -43,6 +43,7 @@ import (
|
||||
var _ adapter.SimpleLifecycle = (*Box)(nil)
|
||||
|
||||
type Box struct {
|
||||
ctx context.Context
|
||||
createdAt time.Time
|
||||
debugOptions option.DebugOptions
|
||||
debugHTTPServer *http.Server
|
||||
@@ -462,6 +463,7 @@ func New(options Options) (*Box, error) {
|
||||
internalServices = append(internalServices, adapter.NewLifecycleService(ntpService, "ntp service"))
|
||||
}
|
||||
return &Box{
|
||||
ctx: ctx,
|
||||
network: networkManager,
|
||||
endpoint: endpointManager,
|
||||
inbound: inboundManager,
|
||||
@@ -533,23 +535,23 @@ func (s *Box) preStart() error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
err = adapter.StartNamed(s.logger, adapter.StartStateInitialize, s.internalService) // cache-file clash-api v2ray-api
|
||||
err = adapter.StartNamed(s.ctx, s.logger, adapter.StartStateInitialize, s.internalService) // cache-file clash-api v2ray-api
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
err = adapter.Start(s.logger, adapter.StartStateInitialize, s.network, s.dnsTransport, s.dnsRouter, s.connection, s.router, s.outbound, s.inbound, s.endpoint, s.service, s.certificateProvider)
|
||||
err = adapter.Start(s.ctx, s.logger, adapter.StartStateInitialize, s.network, s.dnsTransport, s.dnsRouter, s.connection, s.router, s.outbound, s.inbound, s.endpoint, s.service, s.certificateProvider)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
err = adapter.Start(s.logger, adapter.StartStateStart, s.outbound, s.dnsTransport, s.network, s.connection)
|
||||
err = adapter.Start(s.ctx, s.logger, adapter.StartStateStart, s.outbound, s.dnsTransport, s.network, s.connection)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
err = adapter.StartNamed(s.logger, adapter.StartStateStart, []adapter.LifecycleService{s.httpClientService})
|
||||
err = adapter.StartNamed(s.ctx, s.logger, adapter.StartStateStart, []adapter.LifecycleService{s.httpClientService})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
err = adapter.Start(s.logger, adapter.StartStateStart, s.router, s.dnsRouter)
|
||||
err = adapter.Start(s.ctx, s.logger, adapter.StartStateStart, s.router, s.dnsRouter)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -561,35 +563,35 @@ func (s *Box) start() error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
err = adapter.StartNamed(s.logger, adapter.StartStateStart, s.internalService)
|
||||
err = adapter.StartNamed(s.ctx, s.logger, adapter.StartStateStart, s.internalService)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
err = adapter.Start(s.logger, adapter.StartStateStart, s.endpoint)
|
||||
err = adapter.Start(s.ctx, s.logger, adapter.StartStateStart, s.endpoint)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
err = adapter.Start(s.logger, adapter.StartStateStart, s.certificateProvider)
|
||||
err = adapter.Start(s.ctx, s.logger, adapter.StartStateStart, s.certificateProvider)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
err = adapter.Start(s.logger, adapter.StartStateStart, s.inbound, s.service)
|
||||
err = adapter.Start(s.ctx, s.logger, adapter.StartStateStart, s.inbound, s.service)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
err = adapter.Start(s.logger, adapter.StartStatePostStart, s.outbound, s.network, s.dnsTransport, s.dnsRouter, s.connection, s.router, s.endpoint, s.certificateProvider, s.inbound, s.service)
|
||||
err = adapter.Start(s.ctx, s.logger, adapter.StartStatePostStart, s.outbound, s.network, s.dnsTransport, s.dnsRouter, s.connection, s.router, s.endpoint, s.certificateProvider, s.inbound, s.service)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
err = adapter.StartNamed(s.logger, adapter.StartStatePostStart, s.internalService)
|
||||
err = adapter.StartNamed(s.ctx, s.logger, adapter.StartStatePostStart, s.internalService)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
err = adapter.Start(s.logger, adapter.StartStateStarted, s.network, s.dnsTransport, s.dnsRouter, s.connection, s.router, s.outbound, s.endpoint, s.certificateProvider, s.inbound, s.service)
|
||||
err = adapter.Start(s.ctx, s.logger, adapter.StartStateStarted, s.network, s.dnsTransport, s.dnsRouter, s.connection, s.router, s.outbound, s.endpoint, s.certificateProvider, s.inbound, s.service)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
err = adapter.StartNamed(s.logger, adapter.StartStateStarted, s.internalService)
|
||||
err = adapter.StartNamed(s.ctx, s.logger, adapter.StartStateStarted, s.internalService)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
+56
-34
@@ -57,7 +57,10 @@ type StartedService struct {
|
||||
// userID int
|
||||
// groupID int
|
||||
// systemProxyEnabled bool
|
||||
lifecycleAccess sync.Mutex
|
||||
serviceAccess sync.RWMutex
|
||||
closed bool
|
||||
startInterrupted bool
|
||||
serviceStatus *ServiceStatus
|
||||
serviceStatusSubscriber *observable.Subscriber[*ServiceStatus]
|
||||
serviceStatusObserver *observable.Observer[*ServiceStatus]
|
||||
@@ -149,12 +152,19 @@ func (s *StartedService) updateStatus(newStatus ServiceStatus_Type) {
|
||||
s.serviceStatus = statusObject
|
||||
}
|
||||
|
||||
func (s *StartedService) updateStatusError(err error) error {
|
||||
func (s *StartedService) updateStatusError(err error) {
|
||||
statusObject := &ServiceStatus{Status: ServiceStatus_FATAL, ErrorMessage: err.Error()}
|
||||
s.serviceStatusSubscriber.Emit(statusObject)
|
||||
s.serviceStatus = statusObject
|
||||
}
|
||||
|
||||
func (s *StartedService) interruptStart() {
|
||||
s.serviceAccess.Lock()
|
||||
if s.serviceStatus.Status == ServiceStatus_STARTING && s.instance != nil {
|
||||
s.startInterrupted = true
|
||||
s.instance.cancel()
|
||||
}
|
||||
s.serviceAccess.Unlock()
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *StartedService) waitForStarted(ctx context.Context) error {
|
||||
@@ -239,12 +249,13 @@ func (s *StartedService) followInstance(ctx context.Context, run func(ctx contex
|
||||
}
|
||||
|
||||
func (s *StartedService) StartOrReloadService(ctx context.Context, profileContent string, options *OverrideOptions) error {
|
||||
s.interruptStart()
|
||||
s.lifecycleAccess.Lock()
|
||||
defer s.lifecycleAccess.Unlock()
|
||||
s.serviceAccess.Lock()
|
||||
switch s.serviceStatus.Status {
|
||||
case ServiceStatus_IDLE, ServiceStatus_STARTED, ServiceStatus_STARTING, ServiceStatus_FATAL:
|
||||
default:
|
||||
if s.closed {
|
||||
s.serviceAccess.Unlock()
|
||||
return os.ErrInvalid
|
||||
return os.ErrClosed
|
||||
}
|
||||
oldInstance := s.instance
|
||||
if oldInstance != nil {
|
||||
@@ -255,28 +266,41 @@ func (s *StartedService) StartOrReloadService(ctx context.Context, profileConten
|
||||
runtimeDebug.FreeOSMemory()
|
||||
s.serviceAccess.Lock()
|
||||
}
|
||||
s.startInterrupted = false
|
||||
s.updateStatus(ServiceStatus_STARTING)
|
||||
s.resetLogs()
|
||||
s.serviceAccess.Unlock()
|
||||
instance, err := s.newInstance(ctx, profileContent, options)
|
||||
if err != nil {
|
||||
return s.updateStatusError(err)
|
||||
s.serviceAccess.Lock()
|
||||
s.updateStatusError(err)
|
||||
s.serviceAccess.Unlock()
|
||||
return err
|
||||
}
|
||||
s.instance = instance
|
||||
instance.urlTestHistoryStorage.AddUpdateHook(s.urlTestSubscriber)
|
||||
if instance.clashServer != nil {
|
||||
instance.clashServer.AddModeUpdateHook(s.clashModeSubscriber)
|
||||
}
|
||||
s.serviceAccess.Lock()
|
||||
s.instance = instance
|
||||
s.serviceAccess.Unlock()
|
||||
err = instance.Start()
|
||||
s.serviceAccess.Lock()
|
||||
if s.serviceStatus.Status != ServiceStatus_STARTING {
|
||||
if s.startInterrupted {
|
||||
s.startInterrupted = false
|
||||
s.instance = nil
|
||||
s.serviceAccess.Unlock()
|
||||
_ = instance.Close()
|
||||
runtimeDebug.FreeOSMemory()
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
s.instance = nil
|
||||
s.updateStatusError(err)
|
||||
s.serviceAccess.Unlock()
|
||||
_ = instance.Close()
|
||||
return s.updateStatusError(err)
|
||||
runtimeDebug.FreeOSMemory()
|
||||
return err
|
||||
}
|
||||
s.startedAt = time.Now()
|
||||
s.updateStatus(ServiceStatus_STARTED)
|
||||
@@ -286,6 +310,9 @@ func (s *StartedService) StartOrReloadService(ctx context.Context, profileConten
|
||||
}
|
||||
|
||||
func (s *StartedService) Close() {
|
||||
s.serviceAccess.Lock()
|
||||
s.closed = true
|
||||
s.serviceAccess.Unlock()
|
||||
s.serviceStatusSubscriber.Close()
|
||||
s.logSubscriber.Close()
|
||||
s.urlTestSubscriber.Close()
|
||||
@@ -294,19 +321,22 @@ func (s *StartedService) Close() {
|
||||
}
|
||||
|
||||
func (s *StartedService) CloseService() error {
|
||||
s.interruptStart()
|
||||
s.lifecycleAccess.Lock()
|
||||
defer s.lifecycleAccess.Unlock()
|
||||
s.serviceAccess.Lock()
|
||||
switch s.serviceStatus.Status {
|
||||
case ServiceStatus_STARTING, ServiceStatus_STARTED:
|
||||
default:
|
||||
instance := s.instance
|
||||
if instance == nil && s.serviceStatus.Status != ServiceStatus_STARTING && s.serviceStatus.Status != ServiceStatus_STARTED {
|
||||
s.serviceAccess.Unlock()
|
||||
return nil
|
||||
}
|
||||
s.updateStatus(ServiceStatus_STOPPING)
|
||||
instance := s.instance
|
||||
s.instance = nil
|
||||
s.updateStatus(ServiceStatus_STOPPING)
|
||||
s.serviceAccess.Unlock()
|
||||
if instance != nil {
|
||||
_ = instance.Close()
|
||||
}
|
||||
s.serviceAccess.Lock()
|
||||
s.startedAt = time.Time{}
|
||||
s.updateStatus(ServiceStatus_IDLE)
|
||||
s.serviceAccess.Unlock()
|
||||
@@ -317,6 +347,7 @@ func (s *StartedService) CloseService() error {
|
||||
func (s *StartedService) SetError(err error) {
|
||||
s.serviceAccess.Lock()
|
||||
s.updateStatusError(err)
|
||||
s.serviceAccess.Unlock()
|
||||
s.WriteMessage(log.LevelError, err.Error())
|
||||
}
|
||||
|
||||
@@ -417,15 +448,12 @@ func (s *StartedService) SubscribeLog(empty *emptypb.Empty, server grpc.ServerSt
|
||||
|
||||
func (s *StartedService) GetDefaultLogLevel(ctx context.Context, empty *emptypb.Empty) (*DefaultLogLevel, error) {
|
||||
s.serviceAccess.RLock()
|
||||
switch s.serviceStatus.Status {
|
||||
case ServiceStatus_STARTING, ServiceStatus_STARTED:
|
||||
default:
|
||||
s.serviceAccess.RUnlock()
|
||||
boxService := s.instance
|
||||
s.serviceAccess.RUnlock()
|
||||
if boxService == nil {
|
||||
return nil, os.ErrInvalid
|
||||
}
|
||||
logLevel := s.instance.logFactory.Level()
|
||||
s.serviceAccess.RUnlock()
|
||||
return &DefaultLogLevel{Level: LogLevel(logLevel)}, nil
|
||||
return &DefaultLogLevel{Level: LogLevel(boxService.logFactory.Level())}, nil
|
||||
}
|
||||
|
||||
func (s *StartedService) ClearLogs(ctx context.Context, empty *emptypb.Empty) (*emptypb.Empty, error) {
|
||||
@@ -746,14 +774,11 @@ func (s *StartedService) URLTest(ctx context.Context, request *URLTestRequest) (
|
||||
|
||||
func (s *StartedService) SelectOutbound(ctx context.Context, request *SelectOutboundRequest) (*emptypb.Empty, error) {
|
||||
s.serviceAccess.RLock()
|
||||
switch s.serviceStatus.Status {
|
||||
case ServiceStatus_STARTING, ServiceStatus_STARTED:
|
||||
default:
|
||||
s.serviceAccess.RUnlock()
|
||||
return nil, os.ErrInvalid
|
||||
}
|
||||
boxService := s.instance
|
||||
s.serviceAccess.RUnlock()
|
||||
if boxService == nil {
|
||||
return nil, os.ErrInvalid
|
||||
}
|
||||
outboundGroup, isLoaded := boxService.outboundManager.Outbound(request.GroupTag)
|
||||
if !isLoaded {
|
||||
return nil, status.Error(codes.NotFound, "selector not found: "+request.GroupTag)
|
||||
@@ -1080,14 +1105,11 @@ func buildConnectionProto(metadata *trafficcontrol.TrackerMetadata) *Connection
|
||||
|
||||
func (s *StartedService) CloseConnection(ctx context.Context, request *CloseConnectionRequest) (*emptypb.Empty, error) {
|
||||
s.serviceAccess.RLock()
|
||||
switch s.serviceStatus.Status {
|
||||
case ServiceStatus_STARTING, ServiceStatus_STARTED:
|
||||
default:
|
||||
s.serviceAccess.RUnlock()
|
||||
return nil, os.ErrInvalid
|
||||
}
|
||||
boxService := s.instance
|
||||
s.serviceAccess.RUnlock()
|
||||
if boxService == nil {
|
||||
return nil, os.ErrInvalid
|
||||
}
|
||||
if boxService.trafficManager == nil {
|
||||
return nil, status.Error(codes.Unimplemented, "connection tracking not available")
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user