diff --git a/proxy/wireguard/bind.go b/proxy/wireguard/bind.go index 0e71e1539..fda540d4d 100644 --- a/proxy/wireguard/bind.go +++ b/proxy/wireguard/bind.go @@ -52,9 +52,12 @@ func (b *bind) Open(port uint16) (fns []conn.ReceiveFunc, actualPort uint16, err case <-ch: default: errors.LogErrorInner(context.Background(), err, "unexpected closed") - if b.downFunc != nil { + b.mu.Lock() + downFunc := b.downFunc + b.mu.Unlock() + if downFunc != nil { go func() { - common.Must(b.downFunc()) + common.Must(downFunc()) }() } } @@ -76,6 +79,13 @@ func (b *bind) Open(port uint16) (fns []conn.ReceiveFunc, actualPort uint16, err }, uint16(c.LocalAddr().(*net.UDPAddr).Port), nil } +// setDownFunc sets downFunc after the device is created, since the device may already be using the bind. +func (b *bind) setDownFunc(f func() error) { + b.mu.Lock() + defer b.mu.Unlock() + b.downFunc = f +} + func (b *bind) Close() error { b.mu.Lock() defer b.mu.Unlock() diff --git a/proxy/wireguard/client.go b/proxy/wireguard/client.go index 9cfff7265..a1046c361 100644 --- a/proxy/wireguard/client.go +++ b/proxy/wireguard/client.go @@ -288,7 +288,13 @@ func (h *Handler) init(ctx context.Context) error { } return pktConn, nil } - bind := &bind{} + // device.NewDevice may use the bind right away (Up -> BindUpdate -> Open), + // so everything it reads must be set before creating the device. + bind := &bind{ + resolveFunc: resolveFunc, + listenFunc: listenFunc, + reserved: h.conf.Reserved, + } logger := &device.Logger{ Verbosef: func(format string, args ...any) { log.Record(&log.GeneralMessage{ @@ -304,10 +310,7 @@ func (h *Handler) init(ctx context.Context) error { }, } dev := device.NewDevice(h.tun, bind, logger) - bind.resolveFunc = resolveFunc - bind.listenFunc = listenFunc - bind.downFunc = dev.Down - bind.reserved = h.conf.Reserved + bind.setDownFunc(dev.Down) var cfg strings.Builder cfg.WriteString("private_key=" + h.conf.SecretKey + "\n") for _, peer := range h.conf.Peers { diff --git a/proxy/wireguard/server.go b/proxy/wireguard/server.go index f8bdf218b..5eef90e7a 100644 --- a/proxy/wireguard/server.go +++ b/proxy/wireguard/server.go @@ -113,7 +113,7 @@ func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) { users.Store(user.Account.(*MemoryAccount).Pub, user) } - return &Server{ + s := &Server{ conf: conf, ctx: core.ToBackgroundDetachedContext(ctx), policyManager: p, @@ -131,7 +131,10 @@ func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) { pub: pub, users: users, - }, nil + } + // Install the stack's protocol handlers before the device can deliver packets to it (Start -> dev.Up). + CreateForwarder(stack, s.HandleConnection) + return s, nil } func (s *Server) AddUser(ctx context.Context, user *protocol.MemoryUser) error { @@ -320,7 +323,6 @@ func (s *Server) Start() error { return err } s.dev = dev - CreateForwarder(s.stack, s.HandleConnection) return nil }