This commit is contained in:
jdl
2026-06-08 18:36:32 +02:00
parent 465881708f
commit c258d877e8
2 changed files with 62 additions and 52 deletions

View File

@@ -1,6 +1,7 @@
package peer package peer
import ( import (
"fmt"
"log" "log"
"net" "net"
"net/netip" "net/netip"
@@ -10,75 +11,85 @@ import (
const ForwarderBasePort = 5000 const ForwarderBasePort = 5000
// Forwarder is a relay-side UDP forwarding service. It listens on one UDP port // Forwarder is a relay-side UDP forwarding service. It listens on one UDP port
// per peer on the relay's VPN interface and forwards received bytes to each // per possible peer on the relay's VPN interface and forwards received bytes to
// peer's physical WireGuard port. Port assignment: ForwarderBasePort + // each peer's physical WireGuard port. Port assignment: ForwarderBasePort +
// peer_vpn_ip_byte (peer .11 → port 5011). // peer_vpn_ip_byte (peer .11 → port 5011).
// //
// Incoming packets are authenticated for free by the relay's outer WireGuard // All 256 listeners are opened at construction time and kept open for the
// session before they reach the forwarder. // lifetime of the Forwarder. SetEndpoint and RemoveEndpoint only update the
// dialed send connection for a given peer byte.
type Forwarder struct { type Forwarder struct {
vpnIP netip.Addr vpnIP netip.Addr
peers [256]atomic.Pointer[fwdPeer] sendConns [256]atomic.Pointer[net.UDPConn] // updated atomically; read by run goroutines
listenConns [256]*net.UDPConn // written once in NewForwarder; read by Close
} }
// fwdPeer is immutable once created. Endpoint changes produce a new fwdPeer. // NewForwarder opens all 256 forwarding listeners on vpnIP and starts their
type fwdPeer struct { // goroutines. It returns an error if any listener fails to bind.
listenConn *net.UDPConn func NewForwarder(vpnIP netip.Addr) (*Forwarder, error) {
endpoint netip.AddrPort f := &Forwarder{vpnIP: vpnIP}
} for i := range 256 {
listenAddr := netip.AddrPortFrom(vpnIP, ForwarderBasePort+uint16(i))
// NewForwarder creates a Forwarder that will bind listeners to vpnIP. listenConn, err := net.ListenUDP("udp4", net.UDPAddrFromAddrPort(listenAddr))
func NewForwarder(vpnIP netip.Addr) *Forwarder { if err != nil {
return &Forwarder{vpnIP: vpnIP} for j := range i {
} f.listenConns[j].Close()
}
// SetEndpoint registers or updates the physical WireGuard endpoint for the return nil, fmt.Errorf("forwarder listen %v: %w", listenAddr, err)
// peer identified by peerIPByte (last octet of its VPN IP). The old fwdPeer, }
// if any, is closed before the new one is opened on the same port. f.listenConns[i] = listenConn
func (f *Forwarder) SetEndpoint(peerIPByte byte, ep netip.AddrPort) { go f.run(byte(i), listenConn)
if old := f.peers[peerIPByte].Swap(nil); old != nil {
old.listenConn.Close()
} }
return f, nil
listenAddr := netip.AddrPortFrom(f.vpnIP, ForwarderBasePort+uint16(peerIPByte))
listenConn, err := net.ListenUDP("udp4", net.UDPAddrFromAddrPort(listenAddr))
if err != nil {
log.Printf("[Forwarder] listen %v: %v", listenAddr, err)
return
}
p := &fwdPeer{listenConn: listenConn, endpoint: ep}
f.peers[peerIPByte].Store(p)
go p.run()
} }
// RemoveEndpoint stops forwarding for the given peer and closes its listener. // Close shuts down all listeners and open send connections.
func (f *Forwarder) RemoveEndpoint(peerIPByte byte) { func (f *Forwarder) Close() {
if old := f.peers[peerIPByte].Swap(nil); old != nil { for i := range 256 {
old.listenConn.Close() if old := f.sendConns[i].Swap(nil); old != nil {
old.Close()
}
f.listenConns[i].Close()
} }
} }
func (p *fwdPeer) run() { // SetEndpoint registers or replaces the physical WireGuard endpoint for the
// peer identified by b (last octet of its VPN IP).
func (f *Forwarder) SetEndpoint(b byte, ep netip.AddrPort) {
network := "udp4" network := "udp4"
if p.endpoint.Addr().Is6() { if ep.Addr().Is6() {
network = "udp6" network = "udp6"
} }
sendConn, err := net.DialUDP(network, nil, net.UDPAddrFromAddrPort(p.endpoint)) sendConn, err := net.DialUDP(network, nil, net.UDPAddrFromAddrPort(ep))
if err != nil { if err != nil {
log.Printf("[Forwarder] dial %v: %v", p.endpoint, err) log.Printf("[Forwarder] dial %v: %v", ep, err)
return return
} }
defer sendConn.Close() if old := f.sendConns[b].Swap(sendConn); old != nil {
old.Close()
}
}
// RemoveEndpoint stops forwarding for the given peer.
func (f *Forwarder) RemoveEndpoint(b byte) {
if old := f.sendConns[b].Swap(nil); old != nil {
old.Close()
}
}
func (f *Forwarder) run(b byte, listenConn *net.UDPConn) {
buf := make([]byte, 1<<16) buf := make([]byte, 1<<16)
for { for {
n, err := p.listenConn.Read(buf) n, err := listenConn.Read(buf)
if err != nil { if err != nil {
return return
} }
if _, err := sendConn.Write(buf[:n]); err != nil { conn := f.sendConns[b].Load()
log.Printf("[Forwarder] write to %v: %v", p.endpoint, err) if conn == nil {
continue
}
if _, err := conn.Write(buf[:n]); err != nil {
log.Printf("[Forwarder] write: %v", err)
} }
} }
} }

View File

@@ -11,12 +11,11 @@ import (
// no real WireGuard interface is needed. // no real WireGuard interface is needed.
func newTestForwarder(t *testing.T) *Forwarder { func newTestForwarder(t *testing.T) *Forwarder {
t.Helper() t.Helper()
f := NewForwarder(netip.MustParseAddr("127.0.0.1")) f, err := NewForwarder(netip.MustParseAddr("127.0.0.1"))
t.Cleanup(func() { if err != nil {
for b := range 256 { t.Fatalf("NewForwarder: %v", err)
f.RemoveEndpoint(byte(b)) }
} t.Cleanup(f.Close)
})
return f return f
} }