daemon: Fix start or reload race

This commit is contained in:
世界
2026-08-30 17:41:46 +08:00
parent 2e238abe01
commit cfddedb3da
3 changed files with 84 additions and 51 deletions
+13 -4
View File
@@ -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())
+15 -13
View File
@@ -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
View File
@@ -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")
}