WIP
This commit is contained in:
84
peer/forwarder.go
Normal file
84
peer/forwarder.go
Normal file
@@ -0,0 +1,84 @@
|
|||||||
|
package peer
|
||||||
|
|
||||||
|
import (
|
||||||
|
"log"
|
||||||
|
"net"
|
||||||
|
"net/netip"
|
||||||
|
"sync/atomic"
|
||||||
|
)
|
||||||
|
|
||||||
|
const ForwarderBasePort = 5000
|
||||||
|
|
||||||
|
// 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
|
||||||
|
// peer's physical WireGuard port. Port assignment: ForwarderBasePort +
|
||||||
|
// peer_vpn_ip_byte (peer .11 → port 5011).
|
||||||
|
//
|
||||||
|
// Incoming packets are authenticated for free by the relay's outer WireGuard
|
||||||
|
// session before they reach the forwarder.
|
||||||
|
type Forwarder struct {
|
||||||
|
vpnIP netip.Addr
|
||||||
|
peers [256]atomic.Pointer[fwdPeer]
|
||||||
|
}
|
||||||
|
|
||||||
|
// fwdPeer is immutable once created. Endpoint changes produce a new fwdPeer.
|
||||||
|
type fwdPeer struct {
|
||||||
|
listenConn *net.UDPConn
|
||||||
|
endpoint netip.AddrPort
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewForwarder creates a Forwarder that will bind listeners to vpnIP.
|
||||||
|
func NewForwarder(vpnIP netip.Addr) *Forwarder {
|
||||||
|
return &Forwarder{vpnIP: vpnIP}
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetEndpoint registers or updates the physical WireGuard endpoint for the
|
||||||
|
// 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.
|
||||||
|
func (f *Forwarder) SetEndpoint(peerIPByte byte, ep netip.AddrPort) {
|
||||||
|
if old := f.peers[peerIPByte].Swap(nil); old != nil {
|
||||||
|
old.listenConn.Close()
|
||||||
|
}
|
||||||
|
|
||||||
|
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.
|
||||||
|
func (f *Forwarder) RemoveEndpoint(peerIPByte byte) {
|
||||||
|
if old := f.peers[peerIPByte].Swap(nil); old != nil {
|
||||||
|
old.listenConn.Close()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *fwdPeer) run() {
|
||||||
|
network := "udp4"
|
||||||
|
if p.endpoint.Addr().Is6() {
|
||||||
|
network = "udp6"
|
||||||
|
}
|
||||||
|
sendConn, err := net.DialUDP(network, nil, net.UDPAddrFromAddrPort(p.endpoint))
|
||||||
|
if err != nil {
|
||||||
|
log.Printf("[Forwarder] dial %v: %v", p.endpoint, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
defer sendConn.Close()
|
||||||
|
|
||||||
|
buf := make([]byte, 1<<16)
|
||||||
|
for {
|
||||||
|
n, err := p.listenConn.Read(buf)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if _, err := sendConn.Write(buf[:n]); err != nil {
|
||||||
|
log.Printf("[Forwarder] write to %v: %v", p.endpoint, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
155
peer/forwarder_test.go
Normal file
155
peer/forwarder_test.go
Normal file
@@ -0,0 +1,155 @@
|
|||||||
|
package peer
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// newTestForwarder creates a Forwarder using 127.0.0.1 as the VPN IP so that
|
||||||
|
// no real WireGuard interface is needed.
|
||||||
|
func newTestForwarder(t *testing.T) *Forwarder {
|
||||||
|
t.Helper()
|
||||||
|
f := NewForwarder(netip.MustParseAddr("127.0.0.1"))
|
||||||
|
t.Cleanup(func() {
|
||||||
|
for b := range 256 {
|
||||||
|
f.RemoveEndpoint(byte(b))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
return f
|
||||||
|
}
|
||||||
|
|
||||||
|
// listenRandom opens a UDP listener on 127.0.0.1 with a kernel-assigned port.
|
||||||
|
func listenRandom(t *testing.T) *net.UDPConn {
|
||||||
|
t.Helper()
|
||||||
|
conn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("listenRandom: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() { conn.Close() })
|
||||||
|
return conn
|
||||||
|
}
|
||||||
|
|
||||||
|
func addrPort(conn *net.UDPConn) netip.AddrPort {
|
||||||
|
ap, _ := netip.ParseAddrPort(conn.LocalAddr().String())
|
||||||
|
return ap
|
||||||
|
}
|
||||||
|
|
||||||
|
func recvFrom(t *testing.T, conn *net.UDPConn, timeout time.Duration) []byte {
|
||||||
|
t.Helper()
|
||||||
|
conn.SetReadDeadline(time.Now().Add(timeout))
|
||||||
|
buf := make([]byte, 1<<16)
|
||||||
|
n, _, err := conn.ReadFromUDP(buf)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("recvFrom: %v", err)
|
||||||
|
}
|
||||||
|
return buf[:n]
|
||||||
|
}
|
||||||
|
|
||||||
|
func sendTo(t *testing.T, dst netip.AddrPort, payload []byte) {
|
||||||
|
t.Helper()
|
||||||
|
conn, err := net.DialUDP("udp4", nil, net.UDPAddrFromAddrPort(dst))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("sendTo dial: %v", err)
|
||||||
|
}
|
||||||
|
defer conn.Close()
|
||||||
|
if _, err := conn.Write(payload); err != nil {
|
||||||
|
t.Fatalf("sendTo write: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestForwarderBasic(t *testing.T) {
|
||||||
|
f := newTestForwarder(t)
|
||||||
|
target := listenRandom(t)
|
||||||
|
|
||||||
|
const peerByte = byte(11)
|
||||||
|
f.SetEndpoint(peerByte, addrPort(target))
|
||||||
|
|
||||||
|
fwdAddr := netip.AddrPortFrom(netip.MustParseAddr("127.0.0.1"), ForwarderBasePort+uint16(peerByte))
|
||||||
|
sendTo(t, fwdAddr, []byte("hello wireguard"))
|
||||||
|
|
||||||
|
got := recvFrom(t, target, time.Second)
|
||||||
|
if string(got) != "hello wireguard" {
|
||||||
|
t.Fatalf("got %q, want %q", got, "hello wireguard")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestForwarderEndpointUpdate(t *testing.T) {
|
||||||
|
f := newTestForwarder(t)
|
||||||
|
target1 := listenRandom(t)
|
||||||
|
target2 := listenRandom(t)
|
||||||
|
|
||||||
|
const peerByte = byte(22)
|
||||||
|
fwdAddr := netip.AddrPortFrom(netip.MustParseAddr("127.0.0.1"), ForwarderBasePort+uint16(peerByte))
|
||||||
|
|
||||||
|
f.SetEndpoint(peerByte, addrPort(target1))
|
||||||
|
sendTo(t, fwdAddr, []byte("first"))
|
||||||
|
recvFrom(t, target1, time.Second)
|
||||||
|
|
||||||
|
// Update endpoint to target2.
|
||||||
|
f.SetEndpoint(peerByte, addrPort(target2))
|
||||||
|
sendTo(t, fwdAddr, []byte("second"))
|
||||||
|
|
||||||
|
got := recvFrom(t, target2, time.Second)
|
||||||
|
if string(got) != "second" {
|
||||||
|
t.Fatalf("got %q, want %q", got, "second")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Nothing should have arrived at target1 after the update.
|
||||||
|
target1.SetReadDeadline(time.Now().Add(50 * time.Millisecond))
|
||||||
|
buf := make([]byte, 16)
|
||||||
|
if n, _, err := target1.ReadFromUDP(buf); err == nil {
|
||||||
|
t.Fatalf("unexpected packet at old target: %q", buf[:n])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestForwarderRemoveEndpoint(t *testing.T) {
|
||||||
|
f := newTestForwarder(t)
|
||||||
|
target := listenRandom(t)
|
||||||
|
|
||||||
|
const peerByte = byte(33)
|
||||||
|
fwdAddr := netip.AddrPortFrom(netip.MustParseAddr("127.0.0.1"), ForwarderBasePort+uint16(peerByte))
|
||||||
|
|
||||||
|
f.SetEndpoint(peerByte, addrPort(target))
|
||||||
|
sendTo(t, fwdAddr, []byte("before"))
|
||||||
|
recvFrom(t, target, time.Second)
|
||||||
|
|
||||||
|
f.RemoveEndpoint(peerByte)
|
||||||
|
|
||||||
|
// After removal the forwarder port is closed. Send a packet and verify
|
||||||
|
// nothing arrives at target within a short window.
|
||||||
|
sendTo(t, fwdAddr, []byte("after"))
|
||||||
|
target.SetReadDeadline(time.Now().Add(100 * time.Millisecond))
|
||||||
|
buf := make([]byte, 16)
|
||||||
|
if n, _, err := target.ReadFromUDP(buf); err == nil {
|
||||||
|
t.Fatalf("packet arrived after RemoveEndpoint: %q", buf[:n])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestForwarderMultiplePeers(t *testing.T) {
|
||||||
|
f := newTestForwarder(t)
|
||||||
|
targetA := listenRandom(t)
|
||||||
|
targetB := listenRandom(t)
|
||||||
|
|
||||||
|
const (
|
||||||
|
byteA = byte(10)
|
||||||
|
byteB = byte(11)
|
||||||
|
)
|
||||||
|
vpnIP := netip.MustParseAddr("127.0.0.1")
|
||||||
|
fwdA := netip.AddrPortFrom(vpnIP, ForwarderBasePort+uint16(byteA))
|
||||||
|
fwdB := netip.AddrPortFrom(vpnIP, ForwarderBasePort+uint16(byteB))
|
||||||
|
|
||||||
|
f.SetEndpoint(byteA, addrPort(targetA))
|
||||||
|
f.SetEndpoint(byteB, addrPort(targetB))
|
||||||
|
|
||||||
|
sendTo(t, fwdA, []byte("for-a"))
|
||||||
|
sendTo(t, fwdB, []byte("for-b"))
|
||||||
|
|
||||||
|
if got := string(recvFrom(t, targetA, time.Second)); got != "for-a" {
|
||||||
|
t.Fatalf("targetA got %q, want %q", got, "for-a")
|
||||||
|
}
|
||||||
|
if got := string(recvFrom(t, targetB, time.Second)); got != "for-b" {
|
||||||
|
t.Fatalf("targetB got %q, want %q", got, "for-b")
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user