From cfddedb3dab8c3e326d8b11f92608cc3ba21d514 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Fri, 21 Aug 2026 10:44:46 +0800 Subject: [PATCH] daemon: Fix start or reload race --- adapter/lifecycle.go | 17 ++++++-- box.go | 28 ++++++------ daemon/started_service.go | 90 ++++++++++++++++++++++++--------------- 3 files changed, 84 insertions(+), 51 deletions(-) diff --git a/adapter/lifecycle.go b/adapter/lifecycle.go index 2a6a8ed7..8fd8830b 100644 --- a/adapter/lifecycle.go +++ b/adapter/lifecycle.go @@ -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()) diff --git a/box.go b/box.go index eded0b8b..f06a1176 100644 --- a/box.go +++ b/box.go @@ -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 } diff --git a/daemon/started_service.go b/daemon/started_service.go index 59e315c1..06335249 100644 --- a/daemon/started_service.go +++ b/daemon/started_service.go @@ -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") }