mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-09 12:02:40 +00:00
refactor
This commit is contained in:
parent
d737888f6d
commit
404365211e
6 changed files with 259 additions and 222 deletions
|
|
@ -7,6 +7,9 @@ import (
|
|||
"strings"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/protocol"
|
||||
"github.com/xtls/xray-core/common/serial"
|
||||
"github.com/xtls/xray-core/common/task"
|
||||
"github.com/xtls/xray-core/proxy/wireguard"
|
||||
"google.golang.org/protobuf/proto"
|
||||
)
|
||||
|
|
@ -17,9 +20,12 @@ type WireGuardPeerConfig struct {
|
|||
Endpoint string `json:"endpoint"`
|
||||
KeepAlive uint32 `json:"keepAlive"`
|
||||
AllowedIPs []string `json:"allowedIPs,omitempty"`
|
||||
|
||||
Level uint32 `json:"level"`
|
||||
Email string `json:"email"`
|
||||
}
|
||||
|
||||
func (c *WireGuardPeerConfig) Build() (proto.Message, error) {
|
||||
func (c *WireGuardPeerConfig) Build() (*wireguard.PeerConfig, error) {
|
||||
var err error
|
||||
config := new(wireguard.PeerConfig)
|
||||
|
||||
|
|
@ -78,14 +84,32 @@ func (c *WireGuardConfig) Build() (proto.Message, error) {
|
|||
config.Endpoint = c.Address
|
||||
}
|
||||
|
||||
if c.Peers != nil {
|
||||
if c.IsClient {
|
||||
config.Peers = make([]*wireguard.PeerConfig, len(c.Peers))
|
||||
for i, p := range c.Peers {
|
||||
msg, err := p.Build()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
config.Peers[i] = msg.(*wireguard.PeerConfig)
|
||||
config.Peers[i] = msg
|
||||
}
|
||||
} else {
|
||||
config.Users = make([]*protocol.User, len(c.Peers))
|
||||
processUser := func(idx int) error {
|
||||
p := c.Peers[idx]
|
||||
m, err := p.Build()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
config.Users[idx] = &protocol.User{
|
||||
Email: p.Email,
|
||||
Level: p.Level,
|
||||
Account: serial.ToTypedMessage(m),
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if err := task.ParallelForN(len(c.Peers), processUser); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,72 +0,0 @@
|
|||
package wireguard
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"strings"
|
||||
|
||||
"github.com/xtls/xray-core/common/protocol"
|
||||
"google.golang.org/protobuf/proto"
|
||||
)
|
||||
|
||||
// MemoryAccount is the runtime representation of a WireGuard peer credential.
|
||||
// It is produced by PeerConfig.AsAccount and consumed by the UserManager methods
|
||||
// on Server (AddUser / RemoveUser / GetUser / GetUsers / GetUsersCount).
|
||||
type MemoryAccount struct {
|
||||
PublicKey string
|
||||
PreSharedKey string
|
||||
AllowedIPs []string
|
||||
}
|
||||
|
||||
// Equals implements protocol.Account.
|
||||
func (a *MemoryAccount) Equals(other protocol.Account) bool {
|
||||
b, ok := other.(*MemoryAccount)
|
||||
return ok && a.PublicKey == b.PublicKey
|
||||
}
|
||||
|
||||
// ToProto implements protocol.Account.
|
||||
func (a *MemoryAccount) ToProto() proto.Message {
|
||||
return &PeerConfig{
|
||||
PublicKey: a.PublicKey,
|
||||
PreSharedKey: a.PreSharedKey,
|
||||
AllowedIps: a.AllowedIPs,
|
||||
}
|
||||
}
|
||||
|
||||
// AsAccount implements protocol.AsAccount so that PeerConfig can be used as a
|
||||
// typed account in an AddUserOperation. API callers set the @type field to
|
||||
// "xray.proxy.wireguard.PeerConfig".
|
||||
func (p *PeerConfig) AsAccount() (protocol.Account, error) {
|
||||
return &MemoryAccount{
|
||||
PublicKey: p.PublicKey,
|
||||
PreSharedKey: p.PreSharedKey,
|
||||
AllowedIPs: p.AllowedIps,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// buildPeerIPC returns a WireGuard IPC string that adds or updates a single peer.
|
||||
func buildPeerIPC(a *MemoryAccount) string {
|
||||
var b strings.Builder
|
||||
b.WriteString("public_key=" + a.PublicKey + "\n")
|
||||
if a.PreSharedKey != "" {
|
||||
b.WriteString("preshared_key=" + a.PreSharedKey + "\n")
|
||||
}
|
||||
for _, ip := range a.AllowedIPs {
|
||||
b.WriteString("allowed_ip=" + ip + "\n")
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
|
||||
// buildRemovePeerIPC returns a WireGuard IPC string that removes a peer.
|
||||
func buildRemovePeerIPC(publicKey string) string {
|
||||
return "public_key=" + publicKey + "\nremove=true\n"
|
||||
}
|
||||
|
||||
// parseFirstAddr extracts the host address from a CIDR string or plain address
|
||||
// (e.g. "10.0.0.2/32" → 10.0.0.2, "fd00::1" → fd00::1).
|
||||
func parseFirstAddr(cidr string) (netip.Addr, error) {
|
||||
prefix, err := netip.ParsePrefix(cidr)
|
||||
if err != nil {
|
||||
return netip.ParseAddr(cidr)
|
||||
}
|
||||
return prefix.Addr(), nil
|
||||
}
|
||||
|
|
@ -1 +1,60 @@
|
|||
package wireguard
|
||||
|
||||
import (
|
||||
"encoding/hex"
|
||||
"net/netip"
|
||||
|
||||
"github.com/xtls/xray-core/common/protocol"
|
||||
"google.golang.org/protobuf/proto"
|
||||
)
|
||||
|
||||
func (p *PeerConfig) AsAccount() (protocol.Account, error) {
|
||||
pub, err := ParseKey(p.PublicKey)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
allowedIPs := make([]netip.Prefix, 0, len(p.AllowedIps))
|
||||
for _, ip := range p.AllowedIps {
|
||||
p, err := netip.ParsePrefix(ip)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
allowedIPs = append(allowedIPs, p)
|
||||
}
|
||||
|
||||
return &MemoryAccount{
|
||||
Pub: *pub,
|
||||
AllowedIPs: allowedIPs,
|
||||
PreSharedKey: p.PreSharedKey,
|
||||
KeepAlive: p.KeepAlive,
|
||||
}, nil
|
||||
}
|
||||
|
||||
type MemoryAccount struct {
|
||||
Pub [32]byte
|
||||
AllowedIPs []netip.Prefix
|
||||
PreSharedKey string
|
||||
KeepAlive string
|
||||
}
|
||||
|
||||
func (a *MemoryAccount) Equals(other protocol.Account) bool {
|
||||
if b, ok := other.(*MemoryAccount); ok {
|
||||
return a.Pub == b.Pub
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (a *MemoryAccount) ToProto() proto.Message {
|
||||
allowedIPs := make([]string, 0, len(a.AllowedIPs))
|
||||
for _, ip := range a.AllowedIPs {
|
||||
allowedIPs = append(allowedIPs, ip.String())
|
||||
}
|
||||
|
||||
return &PeerConfig{
|
||||
PublicKey: hex.EncodeToString(a.Pub[:]),
|
||||
AllowedIps: allowedIPs,
|
||||
PreSharedKey: a.PreSharedKey,
|
||||
KeepAlive: a.KeepAlive,
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@
|
|||
package wireguard
|
||||
|
||||
import (
|
||||
protocol "github.com/xtls/xray-core/common/protocol"
|
||||
protoreflect "google.golang.org/protobuf/reflect/protoreflect"
|
||||
protoimpl "google.golang.org/protobuf/runtime/protoimpl"
|
||||
reflect "reflect"
|
||||
|
|
@ -157,6 +158,7 @@ type DeviceConfig struct {
|
|||
SecretKey string `protobuf:"bytes,1,opt,name=secret_key,json=secretKey,proto3" json:"secret_key,omitempty"`
|
||||
Endpoint []string `protobuf:"bytes,2,rep,name=endpoint,proto3" json:"endpoint,omitempty"`
|
||||
Peers []*PeerConfig `protobuf:"bytes,3,rep,name=peers,proto3" json:"peers,omitempty"`
|
||||
Users []*protocol.User `protobuf:"bytes,5,rep,name=users,proto3" json:"users,omitempty"`
|
||||
Mtu int32 `protobuf:"varint,4,opt,name=mtu,proto3" json:"mtu,omitempty"`
|
||||
Reserved []byte `protobuf:"bytes,6,opt,name=reserved,proto3" json:"reserved,omitempty"`
|
||||
DomainStrategy DeviceConfig_DomainStrategy `protobuf:"varint,7,opt,name=domain_strategy,json=domainStrategy,proto3,enum=xray.proxy.wireguard.DeviceConfig_DomainStrategy" json:"domain_strategy,omitempty"`
|
||||
|
|
@ -217,6 +219,13 @@ func (x *DeviceConfig) GetPeers() []*PeerConfig {
|
|||
return nil
|
||||
}
|
||||
|
||||
func (x *DeviceConfig) GetUsers() []*protocol.User {
|
||||
if x != nil {
|
||||
return x.Users
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *DeviceConfig) GetMtu() int32 {
|
||||
if x != nil {
|
||||
return x.Mtu
|
||||
|
|
@ -256,7 +265,7 @@ var File_proxy_wireguard_config_proto protoreflect.FileDescriptor
|
|||
|
||||
const file_proxy_wireguard_config_proto_rawDesc = "" +
|
||||
"\n" +
|
||||
"\x1cproxy/wireguard/config.proto\x12\x14xray.proxy.wireguard\"\xad\x01\n" +
|
||||
"\x1cproxy/wireguard/config.proto\x12\x14xray.proxy.wireguard\x1a\x1acommon/protocol/user.proto\"\xad\x01\n" +
|
||||
"\n" +
|
||||
"PeerConfig\x12\x1d\n" +
|
||||
"\n" +
|
||||
|
|
@ -266,12 +275,13 @@ const file_proxy_wireguard_config_proto_rawDesc = "" +
|
|||
"\n" +
|
||||
"keep_alive\x18\x04 \x01(\tR\tkeepAlive\x12\x1f\n" +
|
||||
"\vallowed_ips\x18\x05 \x03(\tR\n" +
|
||||
"allowedIps\"\xaa\x03\n" +
|
||||
"allowedIps\"\xdc\x03\n" +
|
||||
"\fDeviceConfig\x12\x1d\n" +
|
||||
"\n" +
|
||||
"secret_key\x18\x01 \x01(\tR\tsecretKey\x12\x1a\n" +
|
||||
"\bendpoint\x18\x02 \x03(\tR\bendpoint\x126\n" +
|
||||
"\x05peers\x18\x03 \x03(\v2 .xray.proxy.wireguard.PeerConfigR\x05peers\x12\x10\n" +
|
||||
"\x05peers\x18\x03 \x03(\v2 .xray.proxy.wireguard.PeerConfigR\x05peers\x120\n" +
|
||||
"\x05users\x18\x05 \x03(\v2\x1a.xray.common.protocol.UserR\x05users\x12\x10\n" +
|
||||
"\x03mtu\x18\x04 \x01(\x05R\x03mtu\x12\x1a\n" +
|
||||
"\breserved\x18\x06 \x01(\fR\breserved\x12Z\n" +
|
||||
"\x0fdomain_strategy\x18\a \x01(\x0e21.xray.proxy.wireguard.DeviceConfig.DomainStrategyR\x0edomainStrategy\x12\x1b\n" +
|
||||
|
|
@ -305,15 +315,17 @@ var file_proxy_wireguard_config_proto_goTypes = []any{
|
|||
(DeviceConfig_DomainStrategy)(0), // 0: xray.proxy.wireguard.DeviceConfig.DomainStrategy
|
||||
(*PeerConfig)(nil), // 1: xray.proxy.wireguard.PeerConfig
|
||||
(*DeviceConfig)(nil), // 2: xray.proxy.wireguard.DeviceConfig
|
||||
(*protocol.User)(nil), // 3: xray.common.protocol.User
|
||||
}
|
||||
var file_proxy_wireguard_config_proto_depIdxs = []int32{
|
||||
1, // 0: xray.proxy.wireguard.DeviceConfig.peers:type_name -> xray.proxy.wireguard.PeerConfig
|
||||
0, // 1: xray.proxy.wireguard.DeviceConfig.domain_strategy:type_name -> xray.proxy.wireguard.DeviceConfig.DomainStrategy
|
||||
2, // [2:2] is the sub-list for method output_type
|
||||
2, // [2:2] is the sub-list for method input_type
|
||||
2, // [2:2] is the sub-list for extension type_name
|
||||
2, // [2:2] is the sub-list for extension extendee
|
||||
0, // [0:2] is the sub-list for field type_name
|
||||
3, // 1: xray.proxy.wireguard.DeviceConfig.users:type_name -> xray.common.protocol.User
|
||||
0, // 2: xray.proxy.wireguard.DeviceConfig.domain_strategy:type_name -> xray.proxy.wireguard.DeviceConfig.DomainStrategy
|
||||
3, // [3:3] is the sub-list for method output_type
|
||||
3, // [3:3] is the sub-list for method input_type
|
||||
3, // [3:3] is the sub-list for extension type_name
|
||||
3, // [3:3] is the sub-list for extension extendee
|
||||
0, // [0:3] is the sub-list for field type_name
|
||||
}
|
||||
|
||||
func init() { file_proxy_wireguard_config_proto_init() }
|
||||
|
|
|
|||
|
|
@ -6,6 +6,8 @@ option go_package = "github.com/xtls/xray-core/proxy/wireguard";
|
|||
option java_package = "com.xray.proxy.wireguard";
|
||||
option java_multiple_files = true;
|
||||
|
||||
import "common/protocol/user.proto";
|
||||
|
||||
message PeerConfig {
|
||||
string public_key = 1;
|
||||
string pre_shared_key = 2;
|
||||
|
|
@ -25,6 +27,7 @@ message DeviceConfig {
|
|||
string secret_key = 1;
|
||||
repeated string endpoint = 2;
|
||||
repeated PeerConfig peers = 3;
|
||||
repeated xray.common.protocol.User users = 5;
|
||||
int32 mtu = 4;
|
||||
|
||||
bytes reserved = 6;
|
||||
|
|
|
|||
|
|
@ -2,11 +2,12 @@ package wireguard
|
|||
|
||||
import (
|
||||
"context"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"net/netip"
|
||||
"reflect"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
|
||||
"github.com/xtls/xray-core/common/buf"
|
||||
c "github.com/xtls/xray-core/common/ctx"
|
||||
|
|
@ -22,6 +23,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"
|
||||
|
|
@ -45,9 +47,9 @@ type Server struct {
|
|||
dev *device.Device
|
||||
mu sync.Mutex
|
||||
|
||||
peers sync.Map
|
||||
peersByIP sync.Map
|
||||
peerCount atomic.Int64
|
||||
pub [32]byte
|
||||
users *sync.Map
|
||||
emails *sync.Map
|
||||
}
|
||||
|
||||
func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) {
|
||||
|
|
@ -78,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)
|
||||
|
|
@ -107,7 +100,25 @@ func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) {
|
|||
return nil, err
|
||||
}
|
||||
|
||||
s := &Server{
|
||||
pri, err := ParseKey(conf.SecretKey)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var pub [32]byte
|
||||
curve25519.ScalarBaseMult(&pub, pri)
|
||||
|
||||
users := &sync.Map{}
|
||||
emails := &sync.Map{}
|
||||
for _, u := range conf.Users {
|
||||
user, err := u.ToMemoryUser()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
users.Store(user.Account.(*MemoryAccount).Pub, user)
|
||||
emails.Store(user.Email, user)
|
||||
}
|
||||
|
||||
return &Server{
|
||||
conf: conf,
|
||||
ctx: core.ToBackgroundDetachedContext(ctx),
|
||||
policyManager: p,
|
||||
|
|
@ -122,32 +133,83 @@ func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) {
|
|||
|
||||
tun: tun,
|
||||
stack: stack,
|
||||
}
|
||||
|
||||
// Seed peer maps from the static config so that GetUser / GetUsers work
|
||||
// for peers configured at startup, not only for dynamically added ones.
|
||||
for _, peer := range conf.Peers {
|
||||
if peer.PublicKey == "" {
|
||||
continue
|
||||
}
|
||||
mu := &protocol.MemoryUser{
|
||||
Email: peer.PublicKey,
|
||||
Account: &MemoryAccount{
|
||||
PublicKey: peer.PublicKey,
|
||||
PreSharedKey: peer.PreSharedKey,
|
||||
AllowedIPs: peer.AllowedIps,
|
||||
},
|
||||
}
|
||||
s.peers.Store(peer.PublicKey, mu)
|
||||
s.peerCount.Add(1)
|
||||
for _, cidr := range peer.AllowedIps {
|
||||
if addr, err := parseFirstAddr(cidr); err == nil {
|
||||
s.peersByIP.Store(addr, mu)
|
||||
}
|
||||
}
|
||||
}
|
||||
pub: pub,
|
||||
users: users,
|
||||
emails: emails,
|
||||
}, nil
|
||||
}
|
||||
|
||||
return s, 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")
|
||||
for _, ip := range peer.AllowedIPs {
|
||||
sb.WriteString("allowed_ip=" + ip.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)
|
||||
s.emails.Store(user.Email, 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 value, ok := s.emails.Load(email); ok {
|
||||
peer := value.(*protocol.MemoryUser).Account.(*MemoryAccount)
|
||||
err := s.dev.IpcSet("public_key=" + hex.EncodeToString(peer.Pub[:]) + "\n" + "remove=true\n")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
s.emails.Delete(email)
|
||||
s.users.Delete(peer.Pub)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Server) GetUser(ctx context.Context, email string) *protocol.MemoryUser {
|
||||
if value, ok := s.emails.Load(email); ok {
|
||||
return value.(*protocol.MemoryUser)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
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.
|
||||
|
|
@ -227,18 +289,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 _, ip := range peer.AllowedIPs {
|
||||
cfg.WriteString("allowed_ip=" + ip.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
|
||||
|
|
@ -258,21 +322,47 @@ func (s *Server) HandleConnection(conn net.Conn, dest net.Destination) {
|
|||
defer cancel()
|
||||
ctx = c.ContextWithID(ctx, session.NewID())
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
var user *protocol.MemoryUser
|
||||
s.users.Range(func(key, value any) bool {
|
||||
peer := value.(*protocol.MemoryUser).Account.(*MemoryAccount)
|
||||
for _, ip := range peer.AllowedIPs {
|
||||
if ip.Contains(addr) {
|
||||
user = value.(*protocol.MemoryUser)
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
})
|
||||
|
||||
if user == nil {
|
||||
errors.LogError(context.Background(), "nil user for ", remote)
|
||||
return
|
||||
}
|
||||
|
||||
source := net.DestinationFromAddr(conn.RemoteAddr())
|
||||
inbound := session.Inbound{
|
||||
Name: "wireguard",
|
||||
Tag: s.tag,
|
||||
CanSpliceCopy: 3,
|
||||
Source: source,
|
||||
}
|
||||
|
||||
// Tag the session with the peer's user identity when the tunnel source IP
|
||||
// maps to a known peer. This makes the user visible in access logs and
|
||||
// enables per-user traffic accounting via the stats manager.
|
||||
if addrPort, err := netip.ParseAddrPort(conn.RemoteAddr().String()); err == nil {
|
||||
if v, ok := s.peersByIP.Load(addrPort.Addr()); ok {
|
||||
inbound.User = v.(*protocol.MemoryUser)
|
||||
}
|
||||
User: user,
|
||||
}
|
||||
|
||||
ctx = session.ContextWithInbound(ctx, &inbound)
|
||||
|
|
@ -298,92 +388,13 @@ func (s *Server) HandleConnection(conn net.Conn, dest net.Destination) {
|
|||
}
|
||||
}
|
||||
|
||||
// device returns the running *device.Device used by AddUser / RemoveUser.
|
||||
// Callers must NOT hold s.mu.
|
||||
func (s *Server) device() (*device.Device, error) {
|
||||
s.mu.Lock()
|
||||
dev := s.dev
|
||||
s.mu.Unlock()
|
||||
if dev == nil {
|
||||
return nil, errors.New("wireguard server is not running")
|
||||
}
|
||||
return dev, nil
|
||||
}
|
||||
|
||||
// AddUser implements proxy.UserManager.
|
||||
// The peer is installed into the running WireGuard device via IPC and
|
||||
// tracked in the in-memory maps so that subsequent API calls and log
|
||||
// annotations see it immediately.
|
||||
func (s *Server) AddUser(_ context.Context, u *protocol.MemoryUser) error {
|
||||
account, ok := u.Account.(*MemoryAccount)
|
||||
if !ok {
|
||||
return errors.New("not a WireGuard account")
|
||||
}
|
||||
dev, err := s.device()
|
||||
func ParseKey(str string) (*[32]byte, error) {
|
||||
slice, err := hex.DecodeString(str)
|
||||
if err != nil {
|
||||
return err
|
||||
return nil, err
|
||||
}
|
||||
if err := dev.IpcSet(buildPeerIPC(account)); err != nil {
|
||||
return err
|
||||
if len(slice) != 32 {
|
||||
return nil, errors.New("invalid length")
|
||||
}
|
||||
email := u.Email
|
||||
if email == "" {
|
||||
email = account.PublicKey
|
||||
}
|
||||
u.Email = email
|
||||
if _, loaded := s.peers.LoadOrStore(email, u); loaded {
|
||||
return errors.New("peer ", email, " already exists")
|
||||
}
|
||||
s.peerCount.Add(1)
|
||||
for _, cidr := range account.AllowedIPs {
|
||||
if addr, err := parseFirstAddr(cidr); err == nil {
|
||||
s.peersByIP.Store(addr, u)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemoveUser implements proxy.UserManager.
|
||||
func (s *Server) RemoveUser(_ context.Context, email string) error {
|
||||
v, ok := s.peers.LoadAndDelete(email)
|
||||
if !ok {
|
||||
return errors.New("peer ", email, " not found")
|
||||
}
|
||||
s.peerCount.Add(-1)
|
||||
mu := v.(*protocol.MemoryUser)
|
||||
account := mu.Account.(*MemoryAccount)
|
||||
for _, cidr := range account.AllowedIPs {
|
||||
if addr, err := parseFirstAddr(cidr); err == nil {
|
||||
s.peersByIP.Delete(addr)
|
||||
}
|
||||
}
|
||||
dev, err := s.device()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return dev.IpcSet(buildRemovePeerIPC(account.PublicKey))
|
||||
}
|
||||
|
||||
// GetUser implements proxy.UserManager.
|
||||
func (s *Server) GetUser(_ context.Context, email string) *protocol.MemoryUser {
|
||||
v, ok := s.peers.Load(email)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
return v.(*protocol.MemoryUser)
|
||||
}
|
||||
|
||||
// GetUsers implements proxy.UserManager.
|
||||
func (s *Server) GetUsers(_ context.Context) []*protocol.MemoryUser {
|
||||
var out []*protocol.MemoryUser
|
||||
s.peers.Range(func(_, v any) bool {
|
||||
out = append(out, v.(*protocol.MemoryUser))
|
||||
return true
|
||||
})
|
||||
return out
|
||||
}
|
||||
|
||||
// GetUsersCount implements proxy.UserManager.
|
||||
func (s *Server) GetUsersCount(_ context.Context) int64 {
|
||||
return s.peerCount.Load()
|
||||
return (*[32]byte)(slice), nil
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue