mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-08-27 04:09:45 +00:00
refactor(wireguard): address review — drop wgIpcSetter, use *device.Device directly
- Remove the wgIpcSetter interface and the test-only ipcOverride field; AddUser/RemoveUser now use the concrete *device.Device via s.device(). - Drop the explanatory comments inside the Server struct. - Remove the two WireGuard UserManager test files.
This commit is contained in:
parent
8e1f63d85b
commit
d737888f6d
3 changed files with 8 additions and 335 deletions
|
|
@ -1,75 +0,0 @@
|
|||
package wireguard
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestPeerConfigAsAccount(t *testing.T) {
|
||||
pc := &PeerConfig{
|
||||
PublicKey: "pubkey1",
|
||||
PreSharedKey: "psk1",
|
||||
AllowedIps: []string{"10.0.0.2/32", "10.0.0.3/32"},
|
||||
}
|
||||
acc, err := pc.AsAccount()
|
||||
require.NoError(t, err)
|
||||
ma := acc.(*MemoryAccount)
|
||||
assert.Equal(t, "pubkey1", ma.PublicKey)
|
||||
assert.Equal(t, "psk1", ma.PreSharedKey)
|
||||
assert.Equal(t, []string{"10.0.0.2/32", "10.0.0.3/32"}, ma.AllowedIPs)
|
||||
}
|
||||
|
||||
func TestMemoryAccountEquals(t *testing.T) {
|
||||
a := &MemoryAccount{PublicKey: "key1"}
|
||||
assert.True(t, a.Equals(&MemoryAccount{PublicKey: "key1"}))
|
||||
assert.False(t, a.Equals(&MemoryAccount{PublicKey: "key2"}))
|
||||
assert.False(t, a.Equals(nil))
|
||||
}
|
||||
|
||||
func TestMemoryAccountToProto(t *testing.T) {
|
||||
orig := &MemoryAccount{PublicKey: "pk", PreSharedKey: "psk", AllowedIPs: []string{"192.168.0.1/24"}}
|
||||
pc, ok := orig.ToProto().(*PeerConfig)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, orig.PublicKey, pc.PublicKey)
|
||||
assert.Equal(t, orig.PreSharedKey, pc.PreSharedKey)
|
||||
assert.Equal(t, orig.AllowedIPs, pc.AllowedIps)
|
||||
}
|
||||
|
||||
func TestBuildPeerIPC(t *testing.T) {
|
||||
a := &MemoryAccount{PublicKey: "pk1", PreSharedKey: "psk1", AllowedIPs: []string{"10.0.0.2/32", "10.0.0.3/32"}}
|
||||
ipc := buildPeerIPC(a)
|
||||
assert.Contains(t, ipc, "public_key=pk1\n")
|
||||
assert.Contains(t, ipc, "preshared_key=psk1\n")
|
||||
assert.Contains(t, ipc, "allowed_ip=10.0.0.2/32\n")
|
||||
|
||||
// no psk
|
||||
a2 := &MemoryAccount{PublicKey: "pk2", AllowedIPs: []string{"10.0.0.4/32"}}
|
||||
ipc2 := buildPeerIPC(a2)
|
||||
assert.NotContains(t, ipc2, "preshared_key")
|
||||
}
|
||||
|
||||
func TestBuildRemovePeerIPC(t *testing.T) {
|
||||
assert.Equal(t, "public_key=somekey\nremove=true\n", buildRemovePeerIPC("somekey"))
|
||||
}
|
||||
|
||||
func TestParseFirstAddr(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
in, want string
|
||||
wantErr bool
|
||||
}{
|
||||
{"10.0.0.2/32", "10.0.0.2", false},
|
||||
{"10.0.0.5", "10.0.0.5", false},
|
||||
{"fd00::1/128", "fd00::1", false},
|
||||
{"not-an-ip", "", true},
|
||||
} {
|
||||
addr, err := parseFirstAddr(tc.in)
|
||||
if tc.wantErr {
|
||||
assert.Error(t, err, "input %q", tc.in)
|
||||
} else {
|
||||
require.NoError(t, err, "input %q", tc.in)
|
||||
assert.Equal(t, tc.want, addr.String())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -27,13 +27,6 @@ import (
|
|||
"gvisor.dev/gvisor/pkg/tcpip/stack"
|
||||
)
|
||||
|
||||
// wgIpcSetter abstracts the IPC surface of *device.Device used by the
|
||||
// UserManager methods. The interface makes the peer management logic
|
||||
// independently testable without a live WireGuard device.
|
||||
type wgIpcSetter interface {
|
||||
IpcSet(string) error
|
||||
}
|
||||
|
||||
type Server struct {
|
||||
conf *DeviceConfig
|
||||
ctx context.Context
|
||||
|
|
@ -52,15 +45,9 @@ type Server struct {
|
|||
dev *device.Device
|
||||
mu sync.Mutex
|
||||
|
||||
// UserManager state: peers indexed by email and by tunnel IP.
|
||||
peers sync.Map // email (or public key) → *protocol.MemoryUser
|
||||
peersByIP sync.Map // netip.Addr → *protocol.MemoryUser
|
||||
peers sync.Map
|
||||
peersByIP sync.Map
|
||||
peerCount atomic.Int64
|
||||
|
||||
// ipcOverride is non-nil only in tests. When set it is used instead of
|
||||
// s.dev for IpcSet calls, allowing UserManager logic to be tested without
|
||||
// a live WireGuard device.
|
||||
ipcOverride wgIpcSetter
|
||||
}
|
||||
|
||||
func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) {
|
||||
|
|
@ -311,13 +298,9 @@ func (s *Server) HandleConnection(conn net.Conn, dest net.Destination) {
|
|||
}
|
||||
}
|
||||
|
||||
// ipc returns the wgIpcSetter used by AddUser / RemoveUser.
|
||||
// In production this is s.dev (set after Start). Tests inject via ipcOverride.
|
||||
// device returns the running *device.Device used by AddUser / RemoveUser.
|
||||
// Callers must NOT hold s.mu.
|
||||
func (s *Server) ipc() (wgIpcSetter, error) {
|
||||
if s.ipcOverride != nil {
|
||||
return s.ipcOverride, nil
|
||||
}
|
||||
func (s *Server) device() (*device.Device, error) {
|
||||
s.mu.Lock()
|
||||
dev := s.dev
|
||||
s.mu.Unlock()
|
||||
|
|
@ -336,11 +319,11 @@ func (s *Server) AddUser(_ context.Context, u *protocol.MemoryUser) error {
|
|||
if !ok {
|
||||
return errors.New("not a WireGuard account")
|
||||
}
|
||||
ipc, err := s.ipc()
|
||||
dev, err := s.device()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := ipc.IpcSet(buildPeerIPC(account)); err != nil {
|
||||
if err := dev.IpcSet(buildPeerIPC(account)); err != nil {
|
||||
return err
|
||||
}
|
||||
email := u.Email
|
||||
|
|
@ -374,11 +357,11 @@ func (s *Server) RemoveUser(_ context.Context, email string) error {
|
|||
s.peersByIP.Delete(addr)
|
||||
}
|
||||
}
|
||||
ipc, err := s.ipc()
|
||||
dev, err := s.device()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return ipc.IpcSet(buildRemovePeerIPC(account.PublicKey))
|
||||
return dev.IpcSet(buildRemovePeerIPC(account.PublicKey))
|
||||
}
|
||||
|
||||
// GetUser implements proxy.UserManager.
|
||||
|
|
|
|||
|
|
@ -1,235 +0,0 @@
|
|||
package wireguard
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/protobuf/proto"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/protocol"
|
||||
)
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// mockIpc implements wgIpcSetter without a real WireGuard device.
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
type mockIpc struct {
|
||||
calls []string
|
||||
err error
|
||||
}
|
||||
|
||||
func (m *mockIpc) IpcSet(cfg string) error {
|
||||
m.calls = append(m.calls, cfg)
|
||||
return m.err
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// helpers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// newTestServer returns a *Server with ipcOverride set to mock and optionally
|
||||
// pre-seeded with static peers. It bypasses NewServer (which requires a live
|
||||
// xray.Instance) by directly populating the UserManager fields.
|
||||
func newTestServer(mock *mockIpc, staticPeers ...*PeerConfig) *Server {
|
||||
s := &Server{ipcOverride: mock}
|
||||
for _, p := range staticPeers {
|
||||
if p.PublicKey == "" {
|
||||
continue
|
||||
}
|
||||
mu := &protocol.MemoryUser{
|
||||
Email: p.PublicKey,
|
||||
Account: &MemoryAccount{
|
||||
PublicKey: p.PublicKey,
|
||||
PreSharedKey: p.PreSharedKey,
|
||||
AllowedIPs: p.AllowedIps,
|
||||
},
|
||||
}
|
||||
s.peers.Store(p.PublicKey, mu)
|
||||
s.peerCount.Add(1)
|
||||
for _, cidr := range p.AllowedIps {
|
||||
if addr, err := parseFirstAddr(cidr); err == nil {
|
||||
s.peersByIP.Store(addr, mu)
|
||||
}
|
||||
}
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
func memUser(email, pubkey, psk string, ips ...string) *protocol.MemoryUser {
|
||||
return &protocol.MemoryUser{
|
||||
Email: email,
|
||||
Account: &MemoryAccount{
|
||||
PublicKey: pubkey,
|
||||
PreSharedKey: psk,
|
||||
AllowedIPs: ips,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
type wrongAccount struct{}
|
||||
|
||||
func (wrongAccount) Equals(protocol.Account) bool { return false }
|
||||
func (wrongAccount) ToProto() proto.Message { return nil }
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Static peer seeding
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestStaticPeersSeededAtStartup(t *testing.T) {
|
||||
s := newTestServer(
|
||||
&mockIpc{},
|
||||
&PeerConfig{PublicKey: "pk1", AllowedIps: []string{"10.0.0.2/32"}},
|
||||
&PeerConfig{PublicKey: "pk2", AllowedIps: []string{"10.0.0.3/32"}},
|
||||
)
|
||||
|
||||
assert.Equal(t, int64(2), s.GetUsersCount(context.Background()))
|
||||
|
||||
u := s.GetUser(context.Background(), "pk1")
|
||||
require.NotNil(t, u)
|
||||
assert.Equal(t, "pk1", u.Account.(*MemoryAccount).PublicKey)
|
||||
|
||||
addr, _ := netip.ParseAddr("10.0.0.2")
|
||||
v, ok := s.peersByIP.Load(addr)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "pk1", v.(*protocol.MemoryUser).Account.(*MemoryAccount).PublicKey)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// AddUser
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestAddUser(t *testing.T) {
|
||||
mock := &mockIpc{}
|
||||
s := newTestServer(mock)
|
||||
|
||||
require.NoError(t, s.AddUser(context.Background(), memUser("alice", "pk-alice", "psk", "10.0.0.2/32")))
|
||||
|
||||
require.Len(t, mock.calls, 1)
|
||||
assert.Contains(t, mock.calls[0], "public_key=pk-alice\n")
|
||||
assert.Contains(t, mock.calls[0], "preshared_key=psk\n")
|
||||
assert.Contains(t, mock.calls[0], "allowed_ip=10.0.0.2/32\n")
|
||||
|
||||
got := s.GetUser(context.Background(), "alice")
|
||||
require.NotNil(t, got)
|
||||
assert.Equal(t, int64(1), s.GetUsersCount(context.Background()))
|
||||
|
||||
addr, _ := netip.ParseAddr("10.0.0.2")
|
||||
v, ok := s.peersByIP.Load(addr)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "alice", v.(*protocol.MemoryUser).Email)
|
||||
}
|
||||
|
||||
func TestAddUserFallsBackToPubKeyWhenEmailEmpty(t *testing.T) {
|
||||
s := newTestServer(&mockIpc{})
|
||||
require.NoError(t, s.AddUser(context.Background(), memUser("", "pk-nomail", "", "10.0.0.5/32")))
|
||||
require.NotNil(t, s.GetUser(context.Background(), "pk-nomail"))
|
||||
}
|
||||
|
||||
func TestAddUserDuplicate(t *testing.T) {
|
||||
s := newTestServer(&mockIpc{})
|
||||
u := memUser("alice", "pk-alice", "", "10.0.0.2/32")
|
||||
require.NoError(t, s.AddUser(context.Background(), u))
|
||||
assert.Error(t, s.AddUser(context.Background(), u))
|
||||
}
|
||||
|
||||
func TestAddUserWrongAccountType(t *testing.T) {
|
||||
s := newTestServer(&mockIpc{})
|
||||
u := &protocol.MemoryUser{Email: "alice", Account: wrongAccount{}}
|
||||
assert.Error(t, s.AddUser(context.Background(), u))
|
||||
}
|
||||
|
||||
func TestAddUserIpcError(t *testing.T) {
|
||||
mock := &mockIpc{err: errors.New("ipc fail")}
|
||||
s := newTestServer(mock)
|
||||
assert.Error(t, s.AddUser(context.Background(), memUser("alice", "pk", "", "10.0.0.2/32")))
|
||||
assert.Nil(t, s.GetUser(context.Background(), "alice"))
|
||||
assert.Equal(t, int64(0), s.GetUsersCount(context.Background()))
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// RemoveUser
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestRemoveUser(t *testing.T) {
|
||||
mock := &mockIpc{}
|
||||
s := newTestServer(mock)
|
||||
require.NoError(t, s.AddUser(context.Background(), memUser("alice", "pk-alice", "", "10.0.0.2/32")))
|
||||
mock.calls = nil
|
||||
|
||||
require.NoError(t, s.RemoveUser(context.Background(), "alice"))
|
||||
|
||||
require.Len(t, mock.calls, 1)
|
||||
assert.Equal(t, "public_key=pk-alice\nremove=true\n", mock.calls[0])
|
||||
assert.Nil(t, s.GetUser(context.Background(), "alice"))
|
||||
assert.Equal(t, int64(0), s.GetUsersCount(context.Background()))
|
||||
|
||||
addr, _ := netip.ParseAddr("10.0.0.2")
|
||||
_, ok := s.peersByIP.Load(addr)
|
||||
assert.False(t, ok)
|
||||
}
|
||||
|
||||
func TestRemoveUserNotFound(t *testing.T) {
|
||||
s := newTestServer(&mockIpc{})
|
||||
assert.Error(t, s.RemoveUser(context.Background(), "nobody"))
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// GetUser / GetUsers / GetUsersCount
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestGetUserNotFound(t *testing.T) {
|
||||
s := newTestServer(&mockIpc{})
|
||||
assert.Nil(t, s.GetUser(context.Background(), "nobody"))
|
||||
}
|
||||
|
||||
func TestGetUsers(t *testing.T) {
|
||||
s := newTestServer(&mockIpc{})
|
||||
require.NoError(t, s.AddUser(context.Background(), memUser("alice", "pk1", "", "10.0.0.2/32")))
|
||||
require.NoError(t, s.AddUser(context.Background(), memUser("bob", "pk2", "", "10.0.0.3/32")))
|
||||
|
||||
users := s.GetUsers(context.Background())
|
||||
assert.Len(t, users, 2)
|
||||
|
||||
seen := map[string]bool{}
|
||||
for _, u := range users {
|
||||
seen[u.Email] = true
|
||||
}
|
||||
assert.True(t, seen["alice"])
|
||||
assert.True(t, seen["bob"])
|
||||
}
|
||||
|
||||
func TestGetUsersCount(t *testing.T) {
|
||||
s := newTestServer(&mockIpc{})
|
||||
assert.Equal(t, int64(0), s.GetUsersCount(context.Background()))
|
||||
|
||||
require.NoError(t, s.AddUser(context.Background(), memUser("a", "pk1", "", "10.0.0.2/32")))
|
||||
assert.Equal(t, int64(1), s.GetUsersCount(context.Background()))
|
||||
|
||||
require.NoError(t, s.AddUser(context.Background(), memUser("b", "pk2", "", "10.0.0.3/32")))
|
||||
assert.Equal(t, int64(2), s.GetUsersCount(context.Background()))
|
||||
|
||||
require.NoError(t, s.RemoveUser(context.Background(), "a"))
|
||||
assert.Equal(t, int64(1), s.GetUsersCount(context.Background()))
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// peersByIP (used by HandleConnection to annotate sessions)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestPeersByIPPopulatedAndCleanedUp(t *testing.T) {
|
||||
s := newTestServer(&mockIpc{})
|
||||
require.NoError(t, s.AddUser(context.Background(), memUser("alice", "pk-alice", "", "10.0.0.2/32")))
|
||||
|
||||
addr, _ := netip.ParseAddr("10.0.0.2")
|
||||
v, ok := s.peersByIP.Load(addr)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "alice", v.(*protocol.MemoryUser).Email)
|
||||
|
||||
require.NoError(t, s.RemoveUser(context.Background(), "alice"))
|
||||
_, ok = s.peersByIP.Load(addr)
|
||||
assert.False(t, ok)
|
||||
}
|
||||
Loading…
Add table
Add a link
Reference in a new issue