diff --git a/app/router/balancing.go b/app/router/balancing.go index 5f8cb1c21b0b..5b7ef60a6c6c 100644 --- a/app/router/balancing.go +++ b/app/router/balancing.go @@ -136,7 +136,7 @@ func (b *Balancer) SelectOutbounds() ([]string, error) { // GetPrincipleTarget implements routing.BalancerPrincipleTarget func (r *Router) GetPrincipleTarget(tag string) ([]string, error) { - if b, ok := r.balancers[tag]; ok { + if b, ok := (*r.balancers.Load())[tag]; ok { if s, ok := b.strategy.(BalancingPrincipleTarget); ok { candidates, err := b.SelectOutbounds() if err != nil { @@ -151,7 +151,7 @@ func (r *Router) GetPrincipleTarget(tag string) ([]string, error) { // SetOverrideTarget implements routing.BalancerOverrider func (r *Router) SetOverrideTarget(tag, target string) error { - if b, ok := r.balancers[tag]; ok { + if b, ok := (*r.balancers.Load())[tag]; ok { b.override.Put(target) return nil } @@ -160,7 +160,7 @@ func (r *Router) SetOverrideTarget(tag, target string) error { // GetOverrideTarget implements routing.BalancerOverrider func (r *Router) GetOverrideTarget(tag string) (string, error) { - if b, ok := r.balancers[tag]; ok { + if b, ok := (*r.balancers.Load())[tag]; ok { return b.override.Get(), nil } return "", errors.New("cannot find tag") diff --git a/app/router/balancing_override.go b/app/router/balancing_override.go index 96c0aea16997..e3a34b75d0bf 100644 --- a/app/router/balancing_override.go +++ b/app/router/balancing_override.go @@ -2,25 +2,8 @@ package router import ( sync "sync" - - "github.com/xtls/xray-core/common/errors" ) -func (r *Router) OverrideBalancer(balancer string, target string) error { - var b *Balancer - for tag, bl := range r.balancers { - if tag == balancer { - b = bl - break - } - } - if b == nil { - return errors.New("balancer '", balancer, "' not found") - } - b.override.Put(target) - return nil -} - type overrideSettings struct { target string } diff --git a/app/router/router.go b/app/router/router.go index 97a1d0d57bd8..a0b58e5b612e 100644 --- a/app/router/router.go +++ b/app/router/router.go @@ -2,7 +2,10 @@ package router import ( "context" + "maps" + "slices" "sync" + "sync/atomic" "github.com/xtls/xray-core/common" "github.com/xtls/xray-core/common/errors" @@ -17,8 +20,8 @@ import ( // Router is an implementation of routing.Router. type Router struct { domainStrategy Config_DomainStrategy - rules []*Rule - balancers map[string]*Balancer + rules atomic.Pointer[[]*Rule] + balancers atomic.Pointer[map[string]*Balancer] dns dns.Client ctx context.Context @@ -43,21 +46,25 @@ func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm out r.ohm = ohm r.dispatcher = dispatcher - r.balancers = make(map[string]*Balancer, len(config.BalancingRule)) + balancers := make(map[string]*Balancer, len(config.BalancingRule)) for _, rule := range config.BalancingRule { + if _, found := balancers[rule.Tag]; found { + return errors.New("duplicate balancer tag") + } balancer, err := rule.Build(ohm, dispatcher) if err != nil { return err } balancer.InjectContext(ctx) - r.balancers[rule.Tag] = balancer + balancers[rule.Tag] = balancer } - r.rules = make([]*Rule, 0, len(config.Rule)) + // Webhooks that never went live need no cleanup on failure: constructing + // a notifier has no side effects, dropping it is enough. + rules := make([]*Rule, 0, len(config.Rule)) for _, rule := range config.Rule { cond, err := rule.BuildCondition() if err != nil { - r.closeWebhooks() return err } rr := &Rule{ @@ -68,26 +75,23 @@ func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm out if wh := rule.GetWebhook(); wh != nil { notifier, err := NewWebhookNotifier(wh) if err != nil { - r.closeWebhooks() return err } rr.Webhook = notifier } btag := rule.GetBalancingTag() if len(btag) > 0 { - brule, found := r.balancers[btag] + brule, found := balancers[btag] if !found { - if rr.Webhook != nil { - rr.Webhook.Close() - } - r.closeWebhooks() return errors.New("balancer ", btag, " not found") } rr.Balancer = brule } - r.rules = append(r.rules, rr) + rules = append(rules, rr) } + r.balancers.Store(&balancers) + r.rules.Store(&rules) return nil } @@ -124,18 +128,18 @@ func (r *Router) ReloadRules(config *Config, shouldAppend bool) error { r.mu.Lock() defer r.mu.Unlock() - if !shouldAppend { - for _, rule := range r.rules { - if rule.Webhook != nil { - rule.Webhook.Close() - } - } - r.balancers = make(map[string]*Balancer, len(config.BalancingRule)) - r.rules = make([]*Rule, 0, len(config.Rule)) + oldRules := *r.rules.Load() + oldBalancers := *r.balancers.Load() + + newRules := make([]*Rule, 0, len(oldRules)+len(config.Rule)) + newBalancers := make(map[string]*Balancer, len(oldBalancers)+len(config.BalancingRule)) + if shouldAppend { + newRules = append(newRules, oldRules...) + maps.Copy(newBalancers, oldBalancers) } + for _, rule := range config.BalancingRule { - _, found := r.balancers[rule.Tag] - if found { + if _, found := newBalancers[rule.Tag]; found { return errors.New("duplicate balancer tag") } balancer, err := rule.Build(r.ohm, r.dispatcher) @@ -143,27 +147,15 @@ func (r *Router) ReloadRules(config *Config, shouldAppend bool) error { return err } balancer.InjectContext(r.ctx) - r.balancers[rule.Tag] = balancer - } - - startIdx := len(r.rules) - closeNewWebhooks := func() { - for i := startIdx; i < len(r.rules); i++ { - if r.rules[i].Webhook != nil { - r.rules[i].Webhook.Close() - } - } - r.rules = r.rules[:startIdx] + newBalancers[rule.Tag] = balancer } for _, rule := range config.Rule { - if r.RuleExists(rule.GetRuleTag()) { - closeNewWebhooks() - return errors.New("duplicate ruleTag ", rule.GetRuleTag()) + if tag := rule.GetRuleTag(); tag != "" && slices.ContainsFunc(newRules, func(nr *Rule) bool { return nr.RuleTag == tag }) { + return errors.New("duplicate ruleTag ", tag) } cond, err := rule.BuildCondition() if err != nil { - closeNewWebhooks() return err } rr := &Rule{ @@ -174,66 +166,57 @@ func (r *Router) ReloadRules(config *Config, shouldAppend bool) error { if wh := rule.GetWebhook(); wh != nil { notifier, err := NewWebhookNotifier(wh) if err != nil { - closeNewWebhooks() return err } rr.Webhook = notifier } - btag := rule.GetBalancingTag() - if len(btag) > 0 { - brule, found := r.balancers[btag] + if btag := rule.GetBalancingTag(); len(btag) > 0 { + brule, found := newBalancers[btag] if !found { - if rr.Webhook != nil { - rr.Webhook.Close() - } - closeNewWebhooks() return errors.New("balancer ", btag, " not found") } rr.Balancer = brule } - r.rules = append(r.rules, rr) + newRules = append(newRules, rr) } - return nil -} - -func (r *Router) RuleExists(tag string) bool { - if tag != "" { - for _, rule := range r.rules { - if rule.RuleTag == tag { - return true - } - } + r.balancers.Store(&newBalancers) + r.rules.Store(&newRules) + if !shouldAppend { + closeWebhooks(oldRules) } - return false + return nil } // RemoveRule implements routing.Router. func (r *Router) RemoveRule(tag string) error { + if tag == "" { + return errors.New("empty tag name!") + } + r.mu.Lock() defer r.mu.Unlock() - newRules := []*Rule{} - if tag != "" { - for _, rule := range r.rules { - if rule.RuleTag != tag { - newRules = append(newRules, rule) - } else if rule.Webhook != nil { - rule.Webhook.Close() - } + oldRules := *r.rules.Load() + newRules := make([]*Rule, 0, len(oldRules)) + var removed []*Rule + for _, rule := range oldRules { + if rule.RuleTag != tag { + newRules = append(newRules, rule) + } else { + removed = append(removed, rule) } - r.rules = newRules - return nil } - return errors.New("empty tag name!") + r.rules.Store(&newRules) + closeWebhooks(removed) + return nil } // ListRule implements routing.Router func (r *Router) ListRule() []routing.Route { - r.mu.Lock() - defer r.mu.Unlock() - ruleList := make([]routing.Route, 0) - for _, rule := range r.rules { + rules := *r.rules.Load() + ruleList := make([]routing.Route, 0, len(rules)) + for _, rule := range rules { ruleList = append(ruleList, &Route{ outboundTag: rule.Tag, ruleTag: rule.RuleTag, @@ -252,7 +235,9 @@ func (r *Router) pickRouteInternal(ctx routing.Context) (*Rule, routing.Context, ctx = routing_dns.ContextWithDNSClient(ctx, r.dns) } - for _, rule := range r.rules { + rules := *r.rules.Load() + + for _, rule := range rules { if rule.Apply(ctx) { return rule, ctx, nil } @@ -265,7 +250,7 @@ func (r *Router) pickRouteInternal(ctx routing.Context) (*Rule, routing.Context, ctx = routing_dns.ContextWithDNSClient(ctx, r.dns) // Try applying rules again if we have IPs. - for _, rule := range r.rules { + for _, rule := range rules { if rule.Apply(ctx) { return rule, ctx, nil } @@ -279,9 +264,9 @@ func (r *Router) Start() error { return nil } -// closeWebhooks closes all webhook notifiers in the current rule set. -func (r *Router) closeWebhooks() { - for _, rule := range r.rules { +// closeWebhooks closes all webhook notifiers in the given rule set. +func closeWebhooks(rules []*Rule) { + for _, rule := range rules { if rule.Webhook != nil { rule.Webhook.Close() } @@ -292,7 +277,7 @@ func (r *Router) closeWebhooks() { func (r *Router) Close() error { r.mu.Lock() defer r.mu.Unlock() - r.closeWebhooks() + closeWebhooks(*r.rules.Load()) return nil } diff --git a/app/router/webhook.go b/app/router/webhook.go index 0c7a3e4d6680..cac82ee5f92f 100644 --- a/app/router/webhook.go +++ b/app/router/webhook.go @@ -8,6 +8,7 @@ import ( "net" "net/http" "sync" + "sync/atomic" "time" "github.com/xtls/xray-core/common/errors" @@ -40,6 +41,7 @@ type WebhookNotifier struct { deduplication uint32 client *http.Client seen sync.Map + lastSweep atomic.Int64 done chan struct{} wg sync.WaitGroup closeOnce sync.Once @@ -77,11 +79,6 @@ func NewWebhookNotifier(cfg *WebhookConfig) (*WebhookNotifier, error) { } } - if h.deduplication > 0 { - h.wg.Add(1) - go h.cleanupLoop() - } - return h, nil } @@ -201,6 +198,7 @@ func (h *WebhookNotifier) isDuplicate(email string) bool { } ttl := time.Duration(h.deduplication) * time.Second now := time.Now() + h.maybeSweep(now, ttl) if v, loaded := h.seen.LoadOrStore(email, now); loaded { if now.Sub(v.(time.Time)) < ttl { return true @@ -210,27 +208,23 @@ func (h *WebhookNotifier) isDuplicate(email string) bool { return false } -func (h *WebhookNotifier) cleanupLoop() { - defer h.wg.Done() - ttl := time.Duration(h.deduplication) * time.Second - ticker := time.NewTicker(ttl) - defer ticker.Stop() - for { - select { - case <-h.done: - return - case <-ticker.C: - now := time.Now() - h.seen.Range(func(key, value any) bool { - if now.Sub(value.(time.Time)) >= ttl { - h.seen.Delete(key) - } - return true - }) - } +func (h *WebhookNotifier) maybeSweep(now time.Time, ttl time.Duration) { + last := h.lastSweep.Load() + if now.UnixNano()-last < int64(ttl) { + return + } + if !h.lastSweep.CompareAndSwap(last, now.UnixNano()) { + return // another goroutine did the sweep } + h.seen.Range(func(key, value any) bool { + if now.Sub(value.(time.Time)) >= ttl { + h.seen.Delete(key) + } + return true + }) } +// Only need to call if the Notifier is really used, otherwise GC can clean it func (h *WebhookNotifier) Close() error { h.closeOnce.Do(func() { close(h.done)