diff --git a/protocol/wireguard/endpoint.go b/protocol/wireguard/endpoint.go index 14e3ecb53..6412a4862 100644 --- a/protocol/wireguard/endpoint.go +++ b/protocol/wireguard/endpoint.go @@ -75,23 +75,6 @@ func NewEndpoint(ctx context.Context, router adapter.Router, logger log.ContextL udpTimeout = C.UDPTimeout } networkManager := service.FromContext[adapter.NetworkManager](ctx) - var egressPool *tun.UDPEgressPool - udpListener, isUDPListener := common.Cast[dialer.UDPListener](outboundDialer) - if isUDPListener { - anchorControl, egressEnabled := udpListener.UDPListenerControl() - if egressEnabled { - egressPool = tun.NewUDPEgressPool(tun.UDPEgressPoolOptions{ - Logger: logger, - Control: anchorControl, - InterfaceFinder: networkManager.InterfaceFinder(), - InterfaceMonitor: networkManager.InterfaceMonitor(), - ExcludeInterface: options.Name, - IsExempt: func() bool { - return networkManager.AutoRedirectOutputMark() != 0 - }, - }) - } - } wgEndpoint, err := wireguard.NewEndpoint(wireguard.EndpointOptions{ Context: ctx, Logger: logger, @@ -103,8 +86,16 @@ func NewEndpoint(ctx context.Context, router adapter.Router, logger log.ContextL UDPFiltering: tun.NATFiltering(options.UDPFiltering), UDPNATMax: options.UDPNATMax, InterfaceFinder: networkManager.InterfaceFinder(), - EgressPool: egressPool, - Dialer: outboundDialer, + EgressPoolOptions: tun.UDPEgressPoolOptions{ + Logger: logger, + InterfaceFinder: networkManager.InterfaceFinder(), + InterfaceMonitor: networkManager.InterfaceMonitor(), + ExcludeInterface: options.Name, + IsExempt: func() bool { + return networkManager.AutoRedirectOutputMark() != 0 + }, + }, + Dialer: outboundDialer, CreateDialer: func(interfaceName string) N.Dialer { return common.Must1(dialer.NewDefault(ctx, option.DialerOptions{ BindInterface: interfaceName, diff --git a/transport/wireguard/endpoint.go b/transport/wireguard/endpoint.go index 9761783cf..8b2e87ee1 100644 --- a/transport/wireguard/endpoint.go +++ b/transport/wireguard/endpoint.go @@ -13,6 +13,7 @@ import ( "unsafe" "github.com/sagernet/sing-box/common/dialer" + "github.com/sagernet/sing-tun" "github.com/sagernet/sing/common" E "github.com/sagernet/sing/common/exceptions" F "github.com/sagernet/sing/common/format" @@ -35,6 +36,7 @@ type Endpoint struct { returnDevice *returnDeviceWrapper device *device.Device allowedIPs *device.AllowedIPs + egressPool *tun.UDPEgressPool pause pause.Manager pauseCallback *list.Element[pause.Callback] } @@ -154,13 +156,16 @@ func (e *Endpoint) Start(resolve bool) error { var bind conn.Bind udpListener, isUDPListener := common.Cast[dialer.UDPListener](e.options.Dialer) if isUDPListener { - listenerControl, _ := udpListener.UDPListenerControl() + listenerControl, egressEnabled := udpListener.UDPListenerControl() standardBind := conn.NewStdNetBind(listenerControl).(*conn.StdNetBind) if e.options.ListenPort == 0 && len(e.peers) == 1 && e.peers[0].endpoint.IsValid() { standardBind.SetSinglePeerMode() } - if e.options.EgressPool != nil { - standardBind.SetEgressProvider(e.options.EgressPool) + if egressEnabled { + egressPoolOptions := e.options.EgressPoolOptions + egressPoolOptions.Control = listenerControl + e.egressPool = tun.NewUDPEgressPool(egressPoolOptions) + standardBind.SetEgressProvider(e.egressPool) } bind = standardBind } else { @@ -235,8 +240,9 @@ func (e *Endpoint) Close() error { e.pause.UnregisterCallback(e.pauseCallback) e.pauseCallback = nil } - if e.options.EgressPool != nil { - e.options.EgressPool.Close() + if e.egressPool != nil { + e.egressPool.Close() + e.egressPool = nil } if e.device != nil { e.device.Down() diff --git a/transport/wireguard/endpoint_options.go b/transport/wireguard/endpoint_options.go index 0a3e7d997..8a2b5f505 100644 --- a/transport/wireguard/endpoint_options.go +++ b/transport/wireguard/endpoint_options.go @@ -23,18 +23,18 @@ type EndpointOptions struct { UDPFiltering tun.NATFiltering UDPNATMax uint32 - InterfaceFinder control.InterfaceFinder - EgressPool *tun.UDPEgressPool - Dialer N.Dialer - CreateDialer func(interfaceName string) N.Dialer - Name string - MTU uint32 - Address []netip.Prefix - PrivateKey string - ListenPort uint16 - ResolvePeer func(domain string) (netip.Addr, error) - Peers []PeerOptions - Workers int + InterfaceFinder control.InterfaceFinder + EgressPoolOptions tun.UDPEgressPoolOptions + Dialer N.Dialer + CreateDialer func(interfaceName string) N.Dialer + Name string + MTU uint32 + Address []netip.Prefix + PrivateKey string + ListenPort uint16 + ResolvePeer func(domain string) (netip.Addr, error) + Peers []PeerOptions + Workers int } type PeerOptions struct {