Compare commits

..
Author SHA1 Message Date
Fangliding 002433cc46 Fix more 2026-08-25 22:57:39 +08:00
Fangliding 34ebaed38f de dup 2026-08-25 04:00:22 +08:00
Fangliding 5a2afdefdd Dead code 2026-08-25 03:57:19 +08:00
Fangliding 076954fbe8 Refactor router to fix api data race 2026-08-24 15:23:59 +08:00
6 changed files with 113 additions and 249 deletions
+3 -3
View File
@@ -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")
-17
View File
@@ -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
}
+59 -114
View File
@@ -2,7 +2,9 @@ package router
import (
"context"
"maps"
"sync"
"sync/atomic"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/errors"
@@ -17,8 +19,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,52 +45,9 @@ 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))
for _, rule := range config.BalancingRule {
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
r.rules.Store(new([]*Rule))
r.balancers.Store(&map[string]*Balancer{})
return r.ReloadRules(config, false)
}
// PickRoute implements routing.Router.
@@ -124,18 +83,22 @@ 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()
}
oldRules := *r.rules.Load()
oldBalancers := *r.balancers.Load()
var newRules []*Rule
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 {
_, 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 +106,12 @@ 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())
}
cond, err := rule.BuildCondition()
if err != nil {
closeNewWebhooks()
return err
}
rr := &Rule{
@@ -171,69 +119,64 @@ func (r *Router) ReloadRules(config *Config, shouldAppend bool) error {
Tag: rule.GetTag(),
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 {
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)
}
r.balancers.Store(&newBalancers)
r.rules.Store(&newRules)
if !shouldAppend {
closeWebhooks(oldRules)
}
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.
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 +195,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 +210,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 +224,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 +237,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
}
+17 -23
View File
@@ -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)
+34 -25
View File
@@ -3,8 +3,11 @@ package bittorrent
import (
"encoding/binary"
"errors"
"math"
"time"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
)
type SniffHeader struct{}
@@ -36,44 +39,50 @@ func SniffUTP(b []byte) (*SniffHeader, error) {
return nil, common.ErrNoClue
}
// type 4 (ST_SYN), version 1
if b[0] != 0x41 {
buffer := buf.FromBytes(b)
var typeAndVersion uint8
if binary.Read(buffer, binary.BigEndian, &typeAndVersion) != nil {
return nil, common.ErrNoClue
} else if b[0]>>4&0xF > 4 || b[0]&0xF != 1 {
return nil, errNotBittorrent
}
// timestamp_difference is always 0 in new connections
if binary.BigEndian.Uint32(b[8:12]) != 0 {
var extension uint8
if binary.Read(buffer, binary.BigEndian, &extension) != nil {
return nil, common.ErrNoClue
} else if extension != 0 && extension != 1 {
return nil, errNotBittorrent
}
// Walk the extension chain. Selective ack (1) and extension bits (2)
extension, offset := b[1], 20
for extension != 0 {
if len(b) < offset+2 {
if extension != 1 {
return nil, errNotBittorrent
}
length := int(b[offset+1])
switch extension {
case 1: // selective ack
if length < 4 || length%4 != 0 {
return nil, errNotBittorrent
}
case 2: // extension bits: fixed 8 bytes, sent in ST_SYN by µTorrent
if length != 8 {
return nil, errNotBittorrent
}
default:
return nil, errNotBittorrent
if binary.Read(buffer, binary.BigEndian, &extension) != nil {
return nil, common.ErrNoClue
}
if len(b) < offset+2+length {
return nil, errNotBittorrent
var length uint8
if err := binary.Read(buffer, binary.BigEndian, &length); err != nil {
return nil, common.ErrNoClue
}
if common.Error2(buffer.ReadBytes(int32(length))) != nil {
return nil, common.ErrNoClue
}
extension = b[offset]
offset += 2 + length
}
// extensions should consume all ST_SYN payload
if len(b) != offset {
if common.Error2(buffer.ReadBytes(2)) != nil {
return nil, common.ErrNoClue
}
var timestamp uint32
if err := binary.Read(buffer, binary.BigEndian, &timestamp); err != nil {
return nil, common.ErrNoClue
}
if math.Abs(float64(time.Now().UnixMicro()-int64(timestamp))) > float64(24*time.Hour) {
return nil, errNotBittorrent
}
@@ -1,67 +0,0 @@
package bittorrent
import (
"encoding/binary"
"testing"
"github.com/xtls/xray-core/common"
)
// utpPacket builds the fixed 20-byte header defined by BEP 29.
func utpPacket(packetType, extension byte, tsDiff uint32, payload ...byte) []byte {
b := make([]byte, 20)
b[0] = packetType<<4 | 1
b[1] = extension
binary.BigEndian.PutUint16(b[2:4], 0x4a3f) // connection_id, random
binary.BigEndian.PutUint32(b[4:8], 0x8c3a91d2) // timestamp_microseconds, sender's clock
binary.BigEndian.PutUint32(b[8:12], tsDiff)
binary.BigEndian.PutUint32(b[12:16], 0x00100000) // wnd_size
binary.BigEndian.PutUint16(b[16:18], 0x71ee) // seq_nr, random in libutp/libtorrent
binary.BigEndian.PutUint16(b[18:20], 0x0000) // ack_nr
return append(b, payload...)
}
func TestSniffUTP(t *testing.T) {
selectiveAck := []byte{0, 4, 0xff, 0x00, 0xff, 0x00}
wrongVersion := utpPacket(4, 0, 0)
wrongVersion[0] = 4<<4 | 2
cases := []struct {
name string
payload []byte
err error
}{
{"syn", utpPacket(4, 0, 0), nil},
{"syn with selective ack", append(utpPacket(4, 1, 0), selectiveAck...), nil},
{"syn with extension bits", append(utpPacket(4, 2, 0), 0, 8, 1, 2, 3, 4, 5, 6, 7, 8), nil},
{"extension bits with wrong length", append(utpPacket(4, 2, 0), 0, 4, 1, 2, 3, 4), errNotBittorrent},
{"syn with nonzero timestamp_difference", utpPacket(4, 0, 0x1234), errNotBittorrent},
{"syn with trailing payload", utpPacket(4, 0, 0, 'x'), errNotBittorrent},
// txid 0x4100, no EDNS0: the worst case colliding with the uTP header
{"dns query", []byte{
0x41, 0x00, 0x01, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
0x01, 'a', 0x02, 'c', 'o', 0x00, 0x00, 0x01, 0x00, 0x01,
}, errNotBittorrent},
{"established connection packets", utpPacket(0, 0, 0x5678, 'x', 'y', 'z'), errNotBittorrent},
{"state", utpPacket(2, 0, 0x5678), errNotBittorrent},
{"fin", utpPacket(1, 0, 0x5678), errNotBittorrent},
{"wrong version", wrongVersion, errNotBittorrent},
{"unknown packet type", utpPacket(5, 0, 0), errNotBittorrent},
{"unknown extension", utpPacket(4, 2, 0), errNotBittorrent},
{"extension chain past the datagram", utpPacket(4, 1, 0, 0, 8, 0xff), errNotBittorrent},
{"selective ack not in multiples of 4", append(utpPacket(4, 1, 0), 0, 3, 0xff, 0x00, 0xff), errNotBittorrent},
{"shorter than the header", utpPacket(4, 0, 0)[:19], common.ErrNoClue},
}
for _, c := range cases {
t.Run(c.name, func(t *testing.T) {
h, err := SniffUTP(c.payload)
if err != c.err {
t.Fatalf("expected error %v, got %v", c.err, err)
}
if err == nil && h == nil {
t.Fatal("expected a sniff header, got nil")
}
})
}
}