Add balancer: & refactor

This commit is contained in:
Fangliding
2026-09-15 01:33:26 +08:00
parent 52a412d9e2
commit bf3230ec51
3 changed files with 120 additions and 73 deletions
+107 -73
View File
@@ -2,33 +2,43 @@ package outbound
import ( import (
"context" "context"
"slices"
"sort" "sort"
"strings" "strings"
"sync" "sync/atomic"
"github.com/xtls/xray-core/app/proxyman" "github.com/xtls/xray-core/app/proxyman"
"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"
"github.com/xtls/xray-core/common/utils"
"github.com/xtls/xray-core/core" "github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/features/outbound" "github.com/xtls/xray-core/features/outbound"
"github.com/xtls/xray-core/features/routing"
) )
// Manager is to manage all outbound handlers. // Manager is to manage all outbound handlers.
type Manager struct { type Manager struct {
access sync.RWMutex defaultHandler atomic.Pointer[outbound.Handler]
defaultHandler outbound.Handler taggedHandler *utils.TypedSyncMap[string, outbound.Handler]
taggedHandler map[string]outbound.Handler untaggedHandlers atomic.Pointer[[]outbound.Handler]
untaggedHandlers []outbound.Handler running atomic.Bool
running bool tagsCache *utils.TypedSyncMap[string, []string]
tagsCache *sync.Map balancerPicker routing.BalancerPicker
} }
// New creates a new Manager. // New creates a new Manager.
func New(ctx context.Context, config *proxyman.OutboundConfig) (*Manager, error) { func New(ctx context.Context, config *proxyman.OutboundConfig) (*Manager, error) {
m := &Manager{ m := &Manager{
taggedHandler: make(map[string]outbound.Handler), taggedHandler: utils.NewTypedSyncMap[string, outbound.Handler](),
tagsCache: &sync.Map{},
} }
m.tagsCache = utils.NewTypedSyncMap[string, []string]()
empty := make([]outbound.Handler, 0)
m.untaggedHandlers.Store(&empty)
_ = core.OptionalFeatures(ctx, func(router routing.Router) {
if picker, ok := router.(routing.BalancerPicker); ok {
m.balancerPicker = picker
}
})
return m, nil return m, nil
} }
@@ -39,20 +49,25 @@ func (m *Manager) Type() interface{} {
// Start implements core.Feature // Start implements core.Feature
func (m *Manager) Start() error { func (m *Manager) Start() error {
m.access.Lock() m.running.Store(true)
defer m.access.Unlock()
m.running = true var startErr error
m.taggedHandler.Range(func(_ string, h outbound.Handler) bool {
for _, h := range m.taggedHandler {
if err := h.Start(); err != nil { if err := h.Start(); err != nil {
return err startErr = err
return false
} }
return true
})
if startErr != nil {
return startErr
} }
for _, h := range m.untaggedHandlers { if untagged := m.untaggedHandlers.Load(); untagged != nil {
if err := h.Start(); err != nil { for _, h := range *untagged {
return err if err := h.Start(); err != nil {
return err
}
} }
} }
@@ -61,18 +76,18 @@ func (m *Manager) Start() error {
// Close implements core.Feature // Close implements core.Feature
func (m *Manager) Close() error { func (m *Manager) Close() error {
m.access.Lock() m.running.Store(false)
defer m.access.Unlock()
m.running = false
var errs []error var errs []error
for _, h := range m.taggedHandler { m.taggedHandler.Range(func(_ string, h outbound.Handler) bool {
errs = append(errs, h.Close()) errs = append(errs, h.Close())
} return true
})
for _, h := range m.untaggedHandlers { if untagged := m.untaggedHandlers.Load(); untagged != nil {
errs = append(errs, h.Close()) for _, h := range *untagged {
errs = append(errs, h.Close())
}
} }
return errors.Combine(errs...) return errors.Combine(errs...)
@@ -80,47 +95,70 @@ func (m *Manager) Close() error {
// GetDefaultHandler implements outbound.Manager. // GetDefaultHandler implements outbound.Manager.
func (m *Manager) GetDefaultHandler() outbound.Handler { func (m *Manager) GetDefaultHandler() outbound.Handler {
m.access.RLock() if h := m.defaultHandler.Load(); h != nil {
defer m.access.RUnlock() return *h
if m.defaultHandler == nil {
return nil
}
return m.defaultHandler
}
// GetHandler implements outbound.Manager.
func (m *Manager) GetHandler(tag string) outbound.Handler {
m.access.RLock()
defer m.access.RUnlock()
if handler, found := m.taggedHandler[tag]; found {
return handler
} }
return nil return nil
} }
// GetHandler implements outbound.Manager.
func (m *Manager) GetHandler(tag string) outbound.Handler {
if handler, found := m.taggedHandler.Load(tag); found {
return handler
}
if strings.HasPrefix(tag, "balancer:") && m.balancerPicker != nil {
targetTag := m.tryGetOutboundTagWithBalancer(tag, nil)
if targetTag == "" {
return nil
}
if handler, found := m.taggedHandler.Load(targetTag); found {
return handler
}
}
return nil
}
func (m *Manager) tryGetOutboundTagWithBalancer(tag string, parents []string) string {
balancerTag := tag[len("balancer:"):]
targetTag, err := m.balancerPicker.GetBalancerOutboundTag(balancerTag)
if err != nil {
errors.LogWarning(context.Background(), "failed to pick outbound from balancer [", balancerTag, "]: ", err)
return ""
}
if strings.HasPrefix(targetTag, "balancer:") {
if slices.Contains(parents, balancerTag) {
errors.LogWarning(context.Background(), "detected balancer loop for [", balancerTag, "]")
return ""
}
return m.tryGetOutboundTagWithBalancer(targetTag, append(parents, balancerTag))
}
return targetTag
}
// AddHandler implements outbound.Manager. // AddHandler implements outbound.Manager.
func (m *Manager) AddHandler(ctx context.Context, handler outbound.Handler) error { func (m *Manager) AddHandler(ctx context.Context, handler outbound.Handler) error {
m.access.Lock() m.tagsCache.Clear()
defer m.access.Unlock()
m.tagsCache = &sync.Map{} m.defaultHandler.CompareAndSwap(nil, &handler)
if m.defaultHandler == nil {
m.defaultHandler = handler
}
tag := handler.Tag() tag := handler.Tag()
if len(tag) > 0 { if len(tag) > 0 {
if _, found := m.taggedHandler[tag]; found { if _, found := m.taggedHandler.LoadOrStore(tag, handler); found {
return errors.New("existing tag found: " + tag) return errors.New("existing tag found: " + tag)
} }
m.taggedHandler[tag] = handler
} else { } else {
m.untaggedHandlers = append(m.untaggedHandlers, handler) for {
oldUntagged := m.untaggedHandlers.Load()
newUntagged := make([]outbound.Handler, 0, len(*oldUntagged)+1)
newUntagged = append(newUntagged, *oldUntagged...)
newUntagged = append(newUntagged, handler)
if m.untaggedHandlers.CompareAndSwap(oldUntagged, &newUntagged) {
break
}
}
} }
if m.running { if m.running.Load() {
return handler.Start() return handler.Start()
} }
@@ -132,14 +170,12 @@ func (m *Manager) RemoveHandler(ctx context.Context, tag string) error {
if tag == "" { if tag == "" {
return common.ErrNoClue return common.ErrNoClue
} }
m.access.Lock()
defer m.access.Unlock()
m.tagsCache = &sync.Map{} m.tagsCache.Clear()
delete(m.taggedHandler, tag) m.taggedHandler.Delete(tag)
if m.defaultHandler != nil && m.defaultHandler.Tag() == tag { if cur := m.defaultHandler.Load(); cur != nil && (*cur).Tag() == tag {
m.defaultHandler = nil m.defaultHandler.CompareAndSwap(cur, nil)
} }
return nil return nil
@@ -147,39 +183,37 @@ func (m *Manager) RemoveHandler(ctx context.Context, tag string) error {
// ListHandlers implements outbound.Manager. // ListHandlers implements outbound.Manager.
func (m *Manager) ListHandlers(ctx context.Context) []outbound.Handler { func (m *Manager) ListHandlers(ctx context.Context) []outbound.Handler {
m.access.RLock() var response []outbound.Handler
defer m.access.RUnlock() if untagged := m.untaggedHandlers.Load(); untagged != nil {
response = slices.Clone(*untagged)
response := make([]outbound.Handler, len(m.untaggedHandlers))
copy(response, m.untaggedHandlers)
for _, v := range m.taggedHandler {
response = append(response, v)
} }
m.taggedHandler.Range(func(_ string, v outbound.Handler) bool {
response = append(response, v)
return true
})
return response return response
} }
// Select implements outbound.HandlerSelector. // Select implements outbound.HandlerSelector.
func (m *Manager) Select(selectors []string) []string { func (m *Manager) Select(selectors []string) []string {
key := strings.Join(selectors, ",") key := strings.Join(selectors, ",")
if cache, ok := m.tagsCache.Load(key); ok { if result, ok := m.tagsCache.Load(key); ok {
return cache.([]string) return result
} }
m.access.RLock()
defer m.access.RUnlock()
tags := make([]string, 0, len(selectors)) tags := make([]string, 0, len(selectors))
for tag := range m.taggedHandler { m.taggedHandler.Range(func(tag string, _ outbound.Handler) bool {
for _, selector := range selectors { for _, selector := range selectors {
if strings.HasPrefix(tag, selector) { if strings.HasPrefix(tag, selector) {
tags = append(tags, tag) tags = append(tags, tag)
break break
} }
} }
} return true
})
sort.Strings(tags) sort.Strings(tags)
m.tagsCache.Store(key, tags) m.tagsCache.Store(key, tags)
+8
View File
@@ -165,3 +165,11 @@ func (r *Router) GetOverrideTarget(tag string) (string, error) {
} }
return "", errors.New("cannot find tag") return "", errors.New("cannot find tag")
} }
// PickOutbound implements routing.BalancerPicker
func (r *Router) GetBalancerOutboundTag(balancerTag string) (string, error) {
if b, ok := (*r.balancers.Load())[balancerTag]; ok {
return b.PickOutbound()
}
return "", errors.New("cannot find tag")
}
+5
View File
@@ -8,3 +8,8 @@ type BalancerOverrider interface {
type BalancerPrincipleTarget interface { type BalancerPrincipleTarget interface {
GetPrincipleTarget(tag string) ([]string, error) GetPrincipleTarget(tag string) ([]string, error)
} }
// BalancerPicker return the outbound tag selected by the given balancer
type BalancerPicker interface {
GetBalancerOutboundTag(balancerTag string) (string, error)
}