WireGuard inbound: Support dynamic peer management (#6360)

https://github.com/XTLS/Xray-core/pull/6360#issuecomment-4780311547

Closes https://github.com/XTLS/Xray-core/issues/6314

---------

Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com>
Co-authored-by: LjhAUMEM <llnu14702@gmail.com>
This commit is contained in:
bitwiresys
2026-06-27 11:41:22 +00:00
committed by GitHub
co-authored by Claude Sonnet 4.6 LjhAUMEM
parent f496437b84
commit 345c76f9a8
14 changed files with 280 additions and 114 deletions
+156 -16
View File
@@ -2,16 +2,20 @@ package wireguard
import (
"context"
"encoding/hex"
"fmt"
"net/netip"
"reflect"
"strings"
"sync"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
c "github.com/xtls/xray-core/common/ctx"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/log"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/features/policy"
@@ -20,6 +24,7 @@ import (
"github.com/xtls/xray-core/transport"
"github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/transport/internet/stat"
"golang.org/x/crypto/curve25519"
"golang.zx2c4.com/wireguard/device"
"golang.zx2c4.com/wireguard/tun"
"gvisor.dev/gvisor/pkg/tcpip/stack"
@@ -42,6 +47,9 @@ type Server struct {
stack *stack.Stack
dev *device.Device
mu sync.Mutex
pub [32]byte
users *sync.Map
}
func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) {
@@ -72,15 +80,6 @@ func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) {
}
}
if len(conf.Peers) == 0 {
return nil, errors.New("empty peers")
}
for _, peer := range conf.Peers {
if peer.PublicKey == "" {
return nil, errors.New("peer without publickey")
}
}
localAddresses := make([]netip.Addr, 0, len(conf.Endpoint))
for _, localaddress := range conf.Endpoint {
addr, err := netip.ParseAddr(localaddress)
@@ -101,6 +100,19 @@ func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) {
return nil, err
}
pri := common.Must2(ParseKey(conf.SecretKey))
var pub [32]byte
curve25519.ScalarBaseMult(&pub, pri)
users := &sync.Map{}
for _, u := range conf.Users {
user, err := u.ToMemoryUser()
if err != nil {
return nil, err
}
users.Store(user.Account.(*MemoryAccount).Pub, user)
}
return &Server{
conf: conf,
ctx: core.ToBackgroundDetachedContext(ctx),
@@ -116,9 +128,100 @@ func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) {
tun: tun,
stack: stack,
pub: pub,
users: users,
}, nil
}
func (s *Server) AddUser(ctx context.Context, user *protocol.MemoryUser) error {
s.mu.Lock()
defer s.mu.Unlock()
if s.dev == nil {
return errors.New("too early")
}
peer := user.Account.(*MemoryAccount)
if peer.Pub == s.pub {
return errors.New("invalid public key")
}
var sb strings.Builder
sb.WriteString("public_key=" + hex.EncodeToString(peer.Pub[:]) + "\n")
sb.WriteString("replace_allowed_ips=true\n")
for i := range peer.AllowedIPs {
sb.WriteString("allowed_ip=" + peer.AllowedIPs[i].String() + "\n")
}
if peer.PreSharedKey != "" {
sb.WriteString("preshared_key=" + peer.PreSharedKey + "\n")
}
if peer.KeepAlive != "" {
sb.WriteString("persistent_keepalive_interval=" + peer.KeepAlive + "\n")
}
err := s.dev.IpcSet(sb.String())
if err != nil {
return err
}
s.users.Store(peer.Pub, user)
return nil
}
func (s *Server) RemoveUser(ctx context.Context, email string) error {
s.mu.Lock()
defer s.mu.Unlock()
if s.dev == nil {
return errors.New("too early")
}
if user := s.GetUser(ctx, email); user != nil {
peer := user.Account.(*MemoryAccount)
err := s.dev.IpcSet("public_key=" + hex.EncodeToString(peer.Pub[:]) + "\nremove=true\n")
if err != nil {
return err
}
s.users.Delete(peer.Pub)
}
return nil
}
func (s *Server) GetUser(ctx context.Context, email string) (user *protocol.MemoryUser) {
s.users.Range(func(key, value any) bool {
if value.(*protocol.MemoryUser).Email == email {
user = value.(*protocol.MemoryUser)
return false
}
return true
})
return
}
func (s *Server) GetUserByAddr(ctx context.Context, addr netip.Addr) (user *protocol.MemoryUser) {
s.users.Range(func(key, value any) bool {
peer := value.(*protocol.MemoryUser).Account.(*MemoryAccount)
for i := range peer.AllowedIPs {
if peer.AllowedIPs[i].Contains(addr) {
user = value.(*protocol.MemoryUser)
return false
}
}
return true
})
return
}
func (s *Server) GetUsers(ctx context.Context) (users []*protocol.MemoryUser) {
s.users.Range(func(key, value interface{}) bool {
users = append(users, value.(*protocol.MemoryUser))
return true
})
return
}
func (s *Server) GetUsersCount(context.Context) (count int64) {
s.users.Range(func(key, value interface{}) bool {
count++
return true
})
return
}
// Network implements proxy.Inbound.Network.
func (*Server) Network() []net.Network {
return []net.Network{}
@@ -196,18 +299,20 @@ func (s *Server) Start() error {
dev := device.NewDevice(s.tun, bind, logger)
var cfg strings.Builder
cfg.WriteString("private_key=" + s.conf.SecretKey + "\n")
for _, peer := range s.conf.Peers {
cfg.WriteString("public_key=" + peer.PublicKey + "\n")
s.users.Range(func(key, value any) bool {
peer := value.(*protocol.MemoryUser).Account.(*MemoryAccount)
cfg.WriteString("public_key=" + hex.EncodeToString(peer.Pub[:]) + "\n")
for i := range peer.AllowedIPs {
cfg.WriteString("allowed_ip=" + peer.AllowedIPs[i].String() + "\n")
}
if peer.PreSharedKey != "" {
cfg.WriteString("preshared_key=" + peer.PreSharedKey + "\n")
}
for _, ip := range peer.AllowedIps {
cfg.WriteString("allowed_ip=" + ip + "\n")
}
if peer.KeepAlive != "" {
cfg.WriteString("persistent_keepalive_interval=" + peer.KeepAlive + "\n")
}
}
return true
})
err := dev.IpcSet(cfg.String())
if err != nil {
return err
@@ -227,12 +332,36 @@ func (s *Server) HandleConnection(conn net.Conn, dest net.Destination) {
defer cancel()
ctx = c.ContextWithID(ctx, session.NewID())
source := net.DestinationFromAddr(conn.RemoteAddr())
remote := conn.RemoteAddr()
if remote == nil {
errors.LogError(context.Background(), "nil remote")
return
}
var addr netip.Addr
switch v := remote.(type) {
case *net.TCPAddr:
addr, _ = netip.AddrFromSlice(v.IP)
case *net.UDPAddr:
addr, _ = netip.AddrFromSlice(v.IP)
default:
errors.LogError(context.Background(), "invalid addr type ", reflect.TypeOf(v))
return
}
user := s.GetUserByAddr(context.TODO(), addr)
if user == nil {
errors.LogError(context.Background(), "nil user form ", remote, " to ", dest)
return
}
source := net.DestinationFromAddr(remote)
inbound := session.Inbound{
Name: "wireguard",
Tag: s.tag,
CanSpliceCopy: 3,
Source: source,
User: user,
}
ctx = session.ContextWithInbound(ctx, &inbound)
@@ -257,3 +386,14 @@ func (s *Server) HandleConnection(conn net.Conn, dest net.Destination) {
errors.LogError(ctx, errors.New("connection closed").Base(err))
}
}
func ParseKey(str string) (*[32]byte, error) {
slice, err := hex.DecodeString(str)
if err != nil {
return nil, err
}
if len(slice) != 32 {
return nil, errors.New("len(slice) != 32")
}
return (*[32]byte)(slice), nil
}