Compare commits

..
Author SHA1 Message Date
Fangliding a591870b44 Pass ctx correctly 2026-08-08 13:07:09 +08:00
5 changed files with 159 additions and 81 deletions
+2 -2
View File
@@ -172,13 +172,13 @@ func (h *HealthPing) doCheck(ctx context.Context, tags []string, duration time.D
for _, tag := range tags { for _, tag := range tags {
handler := tag handler := tag
client := newPingClient( client := newPingClient(
h.ctx, ctx,
h.dispatcher, h.dispatcher,
h.Settings.Destination, h.Settings.Destination,
h.Settings.Timeout, h.Settings.Timeout,
handler, handler,
) )
for i := 0; i < rounds; i++ { for range rounds {
delay := time.Duration(0) delay := time.Duration(0)
if duration > 0 { if duration > 0 {
delay = time.Duration(dice.RollInt63n(int64(duration))) delay = time.Duration(dice.RollInt63n(int64(duration)))
+3 -3
View File
@@ -136,7 +136,7 @@ func (b *Balancer) SelectOutbounds() ([]string, error) {
// GetPrincipleTarget implements routing.BalancerPrincipleTarget // GetPrincipleTarget implements routing.BalancerPrincipleTarget
func (r *Router) GetPrincipleTarget(tag string) ([]string, error) { func (r *Router) GetPrincipleTarget(tag string) ([]string, error) {
if b, ok := (*r.balancers.Load())[tag]; ok { if b, ok := r.balancers[tag]; ok {
if s, ok := b.strategy.(BalancingPrincipleTarget); ok { if s, ok := b.strategy.(BalancingPrincipleTarget); ok {
candidates, err := b.SelectOutbounds() candidates, err := b.SelectOutbounds()
if err != nil { if err != nil {
@@ -151,7 +151,7 @@ func (r *Router) GetPrincipleTarget(tag string) ([]string, error) {
// SetOverrideTarget implements routing.BalancerOverrider // SetOverrideTarget implements routing.BalancerOverrider
func (r *Router) SetOverrideTarget(tag, target string) error { func (r *Router) SetOverrideTarget(tag, target string) error {
if b, ok := (*r.balancers.Load())[tag]; ok { if b, ok := r.balancers[tag]; ok {
b.override.Put(target) b.override.Put(target)
return nil return nil
} }
@@ -160,7 +160,7 @@ func (r *Router) SetOverrideTarget(tag, target string) error {
// GetOverrideTarget implements routing.BalancerOverrider // GetOverrideTarget implements routing.BalancerOverrider
func (r *Router) GetOverrideTarget(tag string) (string, error) { func (r *Router) GetOverrideTarget(tag string) (string, error) {
if b, ok := (*r.balancers.Load())[tag]; ok { if b, ok := r.balancers[tag]; ok {
return b.override.Get(), nil return b.override.Get(), nil
} }
return "", errors.New("cannot find tag") return "", errors.New("cannot find tag")
+17
View File
@@ -2,8 +2,25 @@ package router
import ( import (
sync "sync" 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 { type overrideSettings struct {
target string target string
} }
+111 -56
View File
@@ -2,9 +2,7 @@ package router
import ( import (
"context" "context"
"maps"
"sync" "sync"
"sync/atomic"
"github.com/xtls/xray-core/common" "github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/errors"
@@ -19,8 +17,8 @@ import (
// Router is an implementation of routing.Router. // Router is an implementation of routing.Router.
type Router struct { type Router struct {
domainStrategy Config_DomainStrategy domainStrategy Config_DomainStrategy
rules atomic.Pointer[[]*Rule] rules []*Rule
balancers atomic.Pointer[map[string]*Balancer] balancers map[string]*Balancer
dns dns.Client dns dns.Client
ctx context.Context ctx context.Context
@@ -45,9 +43,52 @@ func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm out
r.ohm = ohm r.ohm = ohm
r.dispatcher = dispatcher r.dispatcher = dispatcher
r.rules.Store(new([]*Rule)) r.balancers = make(map[string]*Balancer, len(config.BalancingRule))
r.balancers.Store(&map[string]*Balancer{}) for _, rule := range config.BalancingRule {
return r.ReloadRules(config, false) balancer, err := rule.Build(ohm, dispatcher)
if err != nil {
return err
}
balancer.InjectContext(ctx)
r.balancers[rule.Tag] = balancer
}
r.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{
Condition: cond,
Tag: rule.GetTag(),
RuleTag: rule.GetRuleTag(),
}
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]
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)
}
return nil
} }
// PickRoute implements routing.Router. // PickRoute implements routing.Router.
@@ -83,22 +124,18 @@ func (r *Router) ReloadRules(config *Config, shouldAppend bool) error {
r.mu.Lock() r.mu.Lock()
defer r.mu.Unlock() defer r.mu.Unlock()
oldRules := *r.rules.Load() if !shouldAppend {
oldBalancers := *r.balancers.Load() for _, rule := range r.rules {
if rule.Webhook != nil {
var newRules []*Rule rule.Webhook.Close()
newBalancers := make(map[string]*Balancer)
existTags := make(map[string]bool, len(oldRules)+len(config.Rule))
if shouldAppend {
newRules = append(newRules, oldRules...)
maps.Copy(newBalancers, oldBalancers)
for _, rule := range oldRules {
existTags[rule.RuleTag] = true
} }
} }
r.balancers = make(map[string]*Balancer, len(config.BalancingRule))
r.rules = make([]*Rule, 0, len(config.Rule))
}
for _, rule := range config.BalancingRule { for _, rule := range config.BalancingRule {
if _, found := newBalancers[rule.Tag]; found { _, found := r.balancers[rule.Tag]
if found {
return errors.New("duplicate balancer tag") return errors.New("duplicate balancer tag")
} }
balancer, err := rule.Build(r.ohm, r.dispatcher) balancer, err := rule.Build(r.ohm, r.dispatcher)
@@ -106,12 +143,27 @@ func (r *Router) ReloadRules(config *Config, shouldAppend bool) error {
return err return err
} }
balancer.InjectContext(r.ctx) balancer.InjectContext(r.ctx)
newBalancers[rule.Tag] = balancer 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]
} }
for _, rule := range config.Rule { for _, rule := range config.Rule {
if r.RuleExists(rule.GetRuleTag()) {
closeNewWebhooks()
return errors.New("duplicate ruleTag ", rule.GetRuleTag())
}
cond, err := rule.BuildCondition() cond, err := rule.BuildCondition()
if err != nil { if err != nil {
closeNewWebhooks()
return err return err
} }
rr := &Rule{ rr := &Rule{
@@ -119,64 +171,69 @@ func (r *Router) ReloadRules(config *Config, shouldAppend bool) error {
Tag: rule.GetTag(), Tag: rule.GetTag(),
RuleTag: rule.GetRuleTag(), RuleTag: rule.GetRuleTag(),
} }
if rr.RuleTag != "" && existTags[rr.RuleTag] {
return errors.New("duplicate ruleTag ", rr.RuleTag)
}
existTags[rr.RuleTag] = true
if wh := rule.GetWebhook(); wh != nil { if wh := rule.GetWebhook(); wh != nil {
notifier, err := NewWebhookNotifier(wh) notifier, err := NewWebhookNotifier(wh)
if err != nil { if err != nil {
closeNewWebhooks()
return err return err
} }
rr.Webhook = notifier rr.Webhook = notifier
} }
if btag := rule.GetBalancingTag(); len(btag) > 0 { btag := rule.GetBalancingTag()
brule, found := newBalancers[btag] if len(btag) > 0 {
brule, found := r.balancers[btag]
if !found { if !found {
if rr.Webhook != nil {
rr.Webhook.Close()
}
closeNewWebhooks()
return errors.New("balancer ", btag, " not found") return errors.New("balancer ", btag, " not found")
} }
rr.Balancer = brule rr.Balancer = brule
} }
newRules = append(newRules, rr) r.rules = append(r.rules, rr)
} }
r.balancers.Store(&newBalancers)
r.rules.Store(&newRules)
if !shouldAppend {
closeWebhooks(oldRules)
}
return nil return nil
} }
func (r *Router) RuleExists(tag string) bool {
if tag != "" {
for _, rule := range r.rules {
if rule.RuleTag == tag {
return true
}
}
}
return false
}
// RemoveRule implements routing.Router. // RemoveRule implements routing.Router.
func (r *Router) RemoveRule(tag string) error { func (r *Router) RemoveRule(tag string) error {
if tag == "" {
return errors.New("empty tag name!")
}
r.mu.Lock() r.mu.Lock()
defer r.mu.Unlock() defer r.mu.Unlock()
oldRules := *r.rules.Load() newRules := []*Rule{}
newRules := make([]*Rule, 0, len(oldRules)) if tag != "" {
var removed []*Rule for _, rule := range r.rules {
for _, rule := range oldRules {
if rule.RuleTag != tag { if rule.RuleTag != tag {
newRules = append(newRules, rule) newRules = append(newRules, rule)
} else { } else if rule.Webhook != nil {
removed = append(removed, rule) rule.Webhook.Close()
} }
} }
r.rules.Store(&newRules) r.rules = newRules
closeWebhooks(removed)
return nil return nil
} }
return errors.New("empty tag name!")
}
// ListRule implements routing.Router // ListRule implements routing.Router
func (r *Router) ListRule() []routing.Route { func (r *Router) ListRule() []routing.Route {
rules := *r.rules.Load() r.mu.Lock()
ruleList := make([]routing.Route, 0, len(rules)) defer r.mu.Unlock()
for _, rule := range rules { ruleList := make([]routing.Route, 0)
for _, rule := range r.rules {
ruleList = append(ruleList, &Route{ ruleList = append(ruleList, &Route{
outboundTag: rule.Tag, outboundTag: rule.Tag,
ruleTag: rule.RuleTag, ruleTag: rule.RuleTag,
@@ -195,9 +252,7 @@ func (r *Router) pickRouteInternal(ctx routing.Context) (*Rule, routing.Context,
ctx = routing_dns.ContextWithDNSClient(ctx, r.dns) ctx = routing_dns.ContextWithDNSClient(ctx, r.dns)
} }
rules := *r.rules.Load() for _, rule := range r.rules {
for _, rule := range rules {
if rule.Apply(ctx) { if rule.Apply(ctx) {
return rule, ctx, nil return rule, ctx, nil
} }
@@ -210,7 +265,7 @@ func (r *Router) pickRouteInternal(ctx routing.Context) (*Rule, routing.Context,
ctx = routing_dns.ContextWithDNSClient(ctx, r.dns) ctx = routing_dns.ContextWithDNSClient(ctx, r.dns)
// Try applying rules again if we have IPs. // Try applying rules again if we have IPs.
for _, rule := range rules { for _, rule := range r.rules {
if rule.Apply(ctx) { if rule.Apply(ctx) {
return rule, ctx, nil return rule, ctx, nil
} }
@@ -224,9 +279,9 @@ func (r *Router) Start() error {
return nil return nil
} }
// closeWebhooks closes all webhook notifiers in the given rule set. // closeWebhooks closes all webhook notifiers in the current rule set.
func closeWebhooks(rules []*Rule) { func (r *Router) closeWebhooks() {
for _, rule := range rules { for _, rule := range r.rules {
if rule.Webhook != nil { if rule.Webhook != nil {
rule.Webhook.Close() rule.Webhook.Close()
} }
@@ -237,7 +292,7 @@ func closeWebhooks(rules []*Rule) {
func (r *Router) Close() error { func (r *Router) Close() error {
r.mu.Lock() r.mu.Lock()
defer r.mu.Unlock() defer r.mu.Unlock()
closeWebhooks(*r.rules.Load()) r.closeWebhooks()
return nil return nil
} }
+17 -11
View File
@@ -8,7 +8,6 @@ import (
"net" "net"
"net/http" "net/http"
"sync" "sync"
"sync/atomic"
"time" "time"
"github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/errors"
@@ -41,7 +40,6 @@ type WebhookNotifier struct {
deduplication uint32 deduplication uint32
client *http.Client client *http.Client
seen sync.Map seen sync.Map
lastSweep atomic.Int64
done chan struct{} done chan struct{}
wg sync.WaitGroup wg sync.WaitGroup
closeOnce sync.Once closeOnce sync.Once
@@ -79,6 +77,11 @@ func NewWebhookNotifier(cfg *WebhookConfig) (*WebhookNotifier, error) {
} }
} }
if h.deduplication > 0 {
h.wg.Add(1)
go h.cleanupLoop()
}
return h, nil return h, nil
} }
@@ -198,7 +201,6 @@ func (h *WebhookNotifier) isDuplicate(email string) bool {
} }
ttl := time.Duration(h.deduplication) * time.Second ttl := time.Duration(h.deduplication) * time.Second
now := time.Now() now := time.Now()
h.maybeSweep(now, ttl)
if v, loaded := h.seen.LoadOrStore(email, now); loaded { if v, loaded := h.seen.LoadOrStore(email, now); loaded {
if now.Sub(v.(time.Time)) < ttl { if now.Sub(v.(time.Time)) < ttl {
return true return true
@@ -208,14 +210,17 @@ func (h *WebhookNotifier) isDuplicate(email string) bool {
return false return false
} }
func (h *WebhookNotifier) maybeSweep(now time.Time, ttl time.Duration) { func (h *WebhookNotifier) cleanupLoop() {
last := h.lastSweep.Load() defer h.wg.Done()
if now.UnixNano()-last < int64(ttl) { ttl := time.Duration(h.deduplication) * time.Second
ticker := time.NewTicker(ttl)
defer ticker.Stop()
for {
select {
case <-h.done:
return return
} case <-ticker.C:
if !h.lastSweep.CompareAndSwap(last, now.UnixNano()) { now := time.Now()
return // another goroutine did the sweep
}
h.seen.Range(func(key, value any) bool { h.seen.Range(func(key, value any) bool {
if now.Sub(value.(time.Time)) >= ttl { if now.Sub(value.(time.Time)) >= ttl {
h.seen.Delete(key) h.seen.Delete(key)
@@ -223,8 +228,9 @@ func (h *WebhookNotifier) maybeSweep(now time.Time, ttl time.Duration) {
return true return true
}) })
} }
}
}
// Only need to call if the Notifier is really used, otherwise GC can clean it
func (h *WebhookNotifier) Close() error { func (h *WebhookNotifier) Close() error {
h.closeOnce.Do(func() { h.closeOnce.Do(func() {
close(h.done) close(h.done)