Files
kube-vip/pkg/wireguard/wireguard.go
Daniel 50993b63f1 Enhance WireGuard nftables and endpoint handling (#1469)
* do not masquerade for local endpoints

Signed-off-by: Daniel Nägele <daniel@naegele.dev>

* fix: do not add VIP to lo in wg mode

Signed-off-by: Daniel Nägele <daniel@naegele.dev>

* fix: setup policy routing for wg interface

Signed-off-by: Daniel Nägele <daniel@naegele.dev>

* refactor: use k8s API types for protocol

Signed-off-by: Daniel Nägele <daniel@naegele.dev>

* conservatively apply packet mark

only apply the ct mark as packet mark if it matches our calculated
fwmark

Signed-off-by: Daniel Nägele <daniel@naegele.dev>

* use new nftable setup

the nftable setup now uses only one set of chains per tunnel and
utilizes named maps and sets to match NAT the connections properly

Signed-off-by: Daniel Nägele <daniel@naegele.dev>

* ensure proper cleanup

Signed-off-by: Daniel Nägele <daniel@naegele.dev>

* watch kubernetes endpoints

Signed-off-by: Daniel Nägele <daniel@naegele.dev>

* refactor wireguard nftables implementation

Signed-off-by: Daniel Nägele <daniel@naegele.dev>

* use helper for if name determination

Signed-off-by: Daniel Nägele <daniel@naegele.dev>

---------

Signed-off-by: Daniel Nägele <daniel@naegele.dev>
2026-03-24 14:25:05 +01:00

246 lines
6.7 KiB
Go

package wireguard
import (
"fmt"
"net"
"time"
log "log/slog"
"github.com/vishvananda/netlink"
"golang.zx2c4.com/wireguard/wgctrl"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
)
type WGConfig struct {
InterfaceName string // e.g., "wg0"
PrivateKey string
PeerPublicKey string
PeerEndpoint string
Address string
AllowedIPs []string // CIDRs allowed through tunnel
ListenPort int
KeepAlive time.Duration
}
type WireGuard struct {
cfg WGConfig
linkIndex *int
client *wgctrl.Client
}
func NewWireGuard(cfg WGConfig) *WireGuard {
return &WireGuard{
cfg: cfg,
}
}
// Up brings up the interface, configures peers, and sets routes
func (w *WireGuard) Up() error {
log.Info("bringing up wireguard interface", "interface", w.cfg.InterfaceName)
if err := w.createInterface(); err != nil {
_ = w.teardown()
return err
}
if err := w.configurePeer(); err != nil {
_ = w.teardown()
return err
}
if err := w.addRoutes(); err != nil {
_ = w.teardown()
return err
}
return nil
}
// Down tears down the interface and routes
func (w *WireGuard) Down() error {
return w.teardown()
}
// CreateInterface creates the wg interface in the current netns if it doesn't exist
func (w *WireGuard) createInterface() error {
// Validate ListenPort is within valid range for uint32
if w.cfg.ListenPort < 0 || w.cfg.ListenPort > 65535 {
return fmt.Errorf("ListenPort %d is out of valid range (0-65535)", w.cfg.ListenPort)
}
genericLink := &netlink.GenericLink{
LinkAttrs: netlink.LinkAttrs{Name: w.cfg.InterfaceName},
LinkType: "wireguard",
}
err := netlink.LinkAdd(genericLink)
if err != nil && !isExistErr(err) {
return fmt.Errorf("failed to create wireguard interface: %v", err)
}
// Set up interface
link, err := netlink.LinkByName(w.cfg.InterfaceName)
if err != nil {
return fmt.Errorf("cannot find interface after creation: %v", err)
}
if err := netlink.LinkSetUp(link); err != nil {
return fmt.Errorf("failed to bring interface up: %v", err)
}
// Add the VIP address to the WireGuard interface.
// When WireGuard mode is active, the VIP is ONLY on the tunnel interface
// (not also on lo), which prevents the kernel from treating incoming packets
// as loopback traffic and ensures proper INPUT chain processing for nftables DNAT.
addr, err := netlink.ParseAddr(w.cfg.Address)
if err != nil {
return fmt.Errorf("failed to parse address: %s, %v", w.cfg.Address, err)
}
err = netlink.AddrAdd(link, addr)
if err != nil && !isExistErr(err) {
return fmt.Errorf("could not add address to link: %s, %v", w.cfg.Address, err)
}
log.Info("assigned VIP address to WireGuard interface", "interface", w.cfg.InterfaceName, "address", w.cfg.Address)
if err = netlink.LinkSetMTU(link, 1420); err != nil {
return fmt.Errorf("failed to set mtu %w", err)
}
// NOTE: We intentionally do NOT add a blanket "not fwmark" rule here.
// That would route ALL node traffic through the WireGuard tunnel.
// Instead, we use connmark in nftables to mark only response packets
// to incoming WireGuard connections, and policy routing routes only
// those marked packets back through the tunnel.
// See nftables.SetupPolicyRouting() and ApplyDNAT() for details.
w.linkIndex = &link.Attrs().Index
log.Info("created link", "interface", w.cfg.InterfaceName)
return nil
}
// ConfigurePeer sets the private key, peer, allowed IPs, endpoint, keepalive
func (w *WireGuard) configurePeer() error {
client, err := wgctrl.New()
if err != nil {
return fmt.Errorf("failed to open wgctrl client: %v", err)
}
privKey, err := wgtypes.ParseKey(w.cfg.PrivateKey)
if err != nil {
return fmt.Errorf("invalid private key: %v", err)
}
pubKey, err := wgtypes.ParseKey(w.cfg.PeerPublicKey)
if err != nil {
return fmt.Errorf("invalid peer public key: %v", err)
}
addr, err := net.ResolveUDPAddr("udp", w.cfg.PeerEndpoint)
if err != nil {
return fmt.Errorf("failed to resolve UDP address: %v", err)
}
peer := wgtypes.PeerConfig{
PublicKey: pubKey,
Endpoint: addr,
PersistentKeepaliveInterval: &w.cfg.KeepAlive,
ReplaceAllowedIPs: true,
}
for _, cidr := range w.cfg.AllowedIPs {
_, ipnet, err := net.ParseCIDR(cidr)
if err != nil {
return fmt.Errorf("invalid AllowedIP %s: %v", cidr, err)
}
peer.AllowedIPs = append(peer.AllowedIPs, *ipnet)
}
conf := wgtypes.Config{
PrivateKey: &privKey,
ListenPort: &w.cfg.ListenPort,
ReplacePeers: true,
Peers: []wgtypes.PeerConfig{peer},
FirewallMark: &w.cfg.ListenPort,
}
if err = client.ConfigureDevice(w.cfg.InterfaceName, conf); err != nil {
return fmt.Errorf("failed to configure wireguard device: %v", err)
}
w.client = client
log.Info("wireguard device configured", "interface", w.cfg.InterfaceName)
return nil
}
// AddRoutes creates the routes for the VIP/public IPs via the wg interface
func (w *WireGuard) addRoutes() error {
for _, cidr := range w.cfg.AllowedIPs {
_, dstNet, err := net.ParseCIDR(cidr)
if err != nil {
log.Error("invalid route CIDR, skipping", "cidr", cidr, "err", err)
continue
}
route := netlink.Route{
LinkIndex: *w.linkIndex,
Table: w.cfg.ListenPort,
Dst: dstNet,
}
log.Info("Added route", "dst", dstNet.String())
if err := netlink.RouteAdd(&route); err != nil && !isExistErr(err) {
return fmt.Errorf("failed to add route %s: %v", cidr, err)
}
}
return nil
}
// RemoveRoutes deletes the previously added routes
func (w *WireGuard) removeRoutes() {
for _, cidr := range w.cfg.AllowedIPs {
_, dstNet, err := net.ParseCIDR(cidr)
if err != nil {
log.Error("invalid route CIDR, skipping", "cidr", cidr, "err", err)
continue
}
if w.linkIndex == nil {
continue
}
route := netlink.Route{
LinkIndex: *w.linkIndex,
Dst: dstNet,
Table: w.cfg.ListenPort,
}
err = netlink.RouteDel(&route) // best effort
if err != nil {
log.Error("failed to remove route", "cidr", cidr, "err", err)
}
}
log.Info("routes removed", "interface", w.cfg.InterfaceName)
}
// Teardown deletes the interface entirely (called on leadership loss)
func (w *WireGuard) teardown() error {
if w.client != nil {
if err := w.client.Close(); err != nil {
return fmt.Errorf("failed to close wgctrl client: %v", err)
}
w.client = nil
}
w.removeRoutes()
if link, err := netlink.LinkByName(w.cfg.InterfaceName); err == nil {
if err = netlink.LinkDel(link); err != nil {
log.Error("failed to delete link", "err", err)
}
}
w.linkIndex = nil
log.Info("tear down complete", "interface", w.cfg.InterfaceName)
return nil
}
// helper: check if error indicates interface/route exists
func isExistErr(err error) bool {
if err == nil {
return false
}
return err.Error() == "file exists" || err.Error() == "Link exists"
}