Compare commits

..
5 Commits
7 changed files with 447 additions and 39 deletions
+3 -21
View File
@@ -1,7 +1,6 @@
package device
import (
"bytes"
"io"
"net/netip"
"os"
@@ -26,7 +25,7 @@ type FilterRule struct {
// closeAwareDevice wraps a tun.Device along with a flag
// indicating whether its Close method was called.
type closeAwareDevice struct {
isClosed atomic.Bool
isClosed atomic.Bool
tun.Device
closeEventCh chan struct{}
wg sync.WaitGroup
@@ -580,24 +579,7 @@ func (d *MiddleDevice) Read(bufs [][]byte, sizes []int, offset int) (n int, err
// correctly time out instead.
func isLeakedMagicPacket(packet []byte) bool {
payload, ok := extractUDPPayload(packet)
return ok && isMagicPacket(payload)
}
// IsMagicPacket reports whether payload is one of our connectivity-test magic
// packets (a MagicTestRequest or MagicTestResponse). These packets are meant to
// travel directly between physical UDP sockets and must never be encapsulated by
// WireGuard - e.g. if OS routing mistakenly sends one into a WireGuard TUN
// interface (because the destination falls inside a routed tunnel subnet), it
// should be dropped there rather than tunneled, which would otherwise make a
// LAN-local endpoint test falsely appear to succeed over the tunnel.
func isMagicPacket(payload []byte) bool {
if len(payload) >= bind.MagicTestRequestLen && bytes.HasPrefix(payload, bind.MagicTestRequest) {
return true
}
if len(payload) >= bind.MagicTestResponseLen && bytes.HasPrefix(payload, bind.MagicTestResponse) {
return true
}
return false
return ok && bind.IsMagicPacket(payload)
}
// filterDownstreamBufs drops packets going DOWN to the TUN device (from WireGuard)
@@ -728,4 +710,4 @@ func (d *MiddleDevice) WriteToTun(bufs [][]byte, offset int) (int, error) {
return n, err
}
}
}
+1 -1
View File
@@ -32,4 +32,4 @@ require (
)
// To be used ONLY for local development
// replace github.com/fosrl/newt => ../newt
replace github.com/fosrl/newt => ../newt
-2
View File
@@ -1,7 +1,5 @@
github.com/Microsoft/go-winio v0.6.2 h1:F2VQgta7ecxGYO8k3ZZz3RS8fVIXVxONVUPlNERoyfY=
github.com/Microsoft/go-winio v0.6.2/go.mod h1:yd8OoFMLzJbo9gZq8j5qaps8bJ9aShtEA8Ipt1oGCvU=
github.com/fosrl/newt v1.15.0 h1:WpL0whZM1FMjUe2Vy5jSH1bgbxm1O9k1qCyF/mqZT+s=
github.com/fosrl/newt v1.15.0/go.mod h1:l6kWoZPSaXT+ZRUjiyPgwflRqZWYaXpUj9oQ0sOPh4o=
github.com/godbus/dbus/v5 v5.2.2 h1:TUR3TgtSVDmjiXOgAAyaZbYmIeP3DPkld3jgKGV8mXQ=
github.com/godbus/dbus/v5 v5.2.2/go.mod h1:3AAv2+hPq5rdnr5txxxRwiGjPXamgoIHgz9FPBfOp3c=
github.com/google/btree v1.1.3 h1:CVpQJjYgC4VbzxeGVHfvZrv1ctoYCAI8vbl07Fcxlyg=
+13
View File
@@ -51,6 +51,11 @@ func (o *Olm) handleConnect(msg websocket.WSMessage) {
o.updateRegister = nil
}
if o.stopPingRequest != nil {
o.stopPingRequest()
o.stopPingRequest = nil
}
// if there is an existing tunnel then close it
if o.dev != nil {
logger.Info("Got new message. Closing existing tunnel!")
@@ -257,6 +262,14 @@ func (o *Olm) handleConnect(msg websocket.WSMessage) {
network.SetDNSServers([]string{o.dnsProxy.GetProxyIP().String()})
}
if wgData.ExitNode != nil && wgData.ExitNode.Connect {
if err := o.connectExitNode(*wgData.ExitNode); err != nil {
logger.Error("Failed to connect to exit node: %v", err)
}
} else {
logger.Debug("No exit node to connect to (not provided, or connect flag is false)")
}
o.apiServer.SetRegistered(true)
o.registered = true
+287
View File
@@ -0,0 +1,287 @@
package olm
import (
"encoding/json"
"fmt"
"net"
"strings"
"github.com/fosrl/newt/logger"
"github.com/fosrl/newt/network"
"github.com/fosrl/newt/util"
"github.com/fosrl/olm/peers"
"github.com/fosrl/olm/websocket"
)
// exitNodeAliasSiteId is the sentinel siteId used when registering exit node
// aliases with the DNS proxy. It is not a real site, and the JIT handler
// treats siteId 0 as "no JIT lookup", which is correct here since the exit
// node is connected directly rather than on demand.
const exitNodeAliasSiteId = 0
// connectExitNode configures a WireGuard peer connection to an exit node, on the
// same interface and WireGuard device already used for site peers. The exit node
// lives in a different address space than the site tunnel, so a secondary address
// (ExitNodeConfig.TunnelIP) is added to the interface for it - the exit node's own
// WireGuard peer entry only accepts traffic sourced from that address. Nothing here
// is persisted; it's purely in-memory WireGuard/routing state, same as site peers.
func (o *Olm) connectExitNode(cfg ExitNodeConfig) error {
if !o.tunnelRunning {
return fmt.Errorf("tunnel not running")
}
if cfg.PublicKey == "" || cfg.Endpoint == "" || cfg.ServerIP == "" || cfg.TunnelIP == "" {
return fmt.Errorf("incomplete exit node configuration")
}
o.exitNodeMu.Lock()
defer o.exitNodeMu.Unlock()
dev := o.dev
if dev == nil {
return fmt.Errorf("wireguard device not initialized")
}
if o.exitNode != nil && o.exitNode.PublicKey != cfg.PublicKey {
logger.Info("Switching exit nodes, removing previous exit node peer")
if err := o.removeExitNodePeerLocked(); err != nil {
logger.Warn("Failed to remove previous exit node peer: %v", err)
}
}
endpoint := cfg.Endpoint
if !strings.Contains(endpoint, ":") {
relayPort := cfg.RelayPort
if relayPort == 0 {
relayPort = 21820
}
endpoint = fmt.Sprintf("%s:%d", endpoint, relayPort)
}
resolvedEndpoint, err := util.ResolveDomain(endpoint)
if err != nil {
return fmt.Errorf("failed to resolve exit node endpoint: %w", err)
}
persistentKeepalive := 0
if pm := o.getPeerManager(); pm != nil {
persistentKeepalive = pm.PersistentKeepalive
}
allowedIP := strings.Split(cfg.ServerIP, "/")[0] + "/32"
wgConfig := fmt.Sprintf(`public_key=%s
allowed_ip=%s
endpoint=%s
persistent_keepalive_interval=%d`, util.FixKey(cfg.PublicKey), allowedIP, resolvedEndpoint, persistentKeepalive)
if err := dev.IpcSet(wgConfig); err != nil {
return fmt.Errorf("failed to configure exit node peer: %w", err)
}
interfaceName := o.tunnelConfig.InterfaceName
tunnelIP := cfg.TunnelIP
if !strings.Contains(tunnelIP, "/") {
tunnelIP += "/32"
}
if err := network.AddSecondaryAddress(interfaceName, tunnelIP); err != nil {
logger.Warn("Failed to add secondary address %s for exit node: %v", tunnelIP, err)
}
if err := network.AddRouteForServerIP(cfg.ServerIP, interfaceName); err != nil {
logger.Warn("Failed to add route for exit node server IP: %v", err)
}
cfgCopy := cfg
o.exitNode = &cfgCopy
if o.dnsProxy != nil {
serverIP := net.ParseIP(cfg.ServerIP)
if serverIP != nil {
for _, alias := range cfg.Aliases {
logger.Debug("Adding alias %s to the edit node", alias)
if err := o.dnsProxy.AddDNSRecord(alias, serverIP, exitNodeAliasSiteId); err != nil {
logger.Warn("Failed to add DNS record for exit node alias %s: %v", alias, err)
}
}
}
}
logger.Info("Connected to exit node at %s", resolvedEndpoint)
return nil
}
// disconnectExitNode tears down the current exit node peer connection, if any.
func (o *Olm) disconnectExitNode() error {
o.exitNodeMu.Lock()
defer o.exitNodeMu.Unlock()
return o.removeExitNodePeerLocked()
}
// removeExitNodePeerLocked removes the current exit node peer, its secondary
// interface address, and its server IP route. Must be called with exitNodeMu held.
func (o *Olm) removeExitNodePeerLocked() error {
if o.exitNode == nil {
return nil
}
cfg := o.exitNode
o.exitNode = nil
if o.dnsProxy != nil {
serverIP := net.ParseIP(cfg.ServerIP)
if serverIP != nil {
for _, alias := range cfg.Aliases {
o.dnsProxy.RemoveDNSRecordForSite(alias, serverIP, exitNodeAliasSiteId)
}
}
}
if o.dev != nil {
if err := peers.RemovePeer(o.dev, 0, cfg.PublicKey); err != nil {
logger.Warn("Failed to remove exit node peer: %v", err)
}
}
interfaceName := o.tunnelConfig.InterfaceName
if err := network.RemoveRouteForServerIP(cfg.ServerIP, interfaceName); err != nil {
logger.Warn("Failed to remove route for exit node server IP: %v", err)
}
tunnelIP := cfg.TunnelIP
if !strings.Contains(tunnelIP, "/") {
tunnelIP += "/32"
}
if err := network.RemoveSecondaryAddress(interfaceName, tunnelIP); err != nil {
logger.Warn("Failed to remove secondary address %s for exit node: %v", tunnelIP, err)
}
logger.Info("Disconnected from exit node")
return nil
}
// handleExitNodeConnect handles a server-initiated request to connect to (or switch to)
// an exit node, delivered as a full ExitNodeConfig payload.
func (o *Olm) handleExitNodeConnect(msg websocket.WSMessage) {
logger.Debug("Received exit node connect message: %v", msg.Data)
if !o.tunnelRunning {
logger.Debug("Tunnel stopped, ignoring exit node connect message")
return
}
jsonData, err := json.Marshal(msg.Data)
if err != nil {
logger.Error("Error marshaling exit node connect data: %v", err)
return
}
var cfg ExitNodeConfig
if err := json.Unmarshal(jsonData, &cfg); err != nil {
logger.Error("Error unmarshaling exit node connect data: %v", err)
return
}
if !cfg.Connect {
logger.Debug("Exit node connect message has connect=false, disconnecting instead")
if err := o.disconnectExitNode(); err != nil {
logger.Error("Failed to disconnect from exit node: %v", err)
}
return
}
if err := o.connectExitNode(cfg); err != nil {
logger.Error("Failed to connect to exit node: %v", err)
}
}
// handleExitNodeDisconnect handles a server-initiated request to disconnect from the
// currently connected exit node.
func (o *Olm) handleExitNodeDisconnect(msg websocket.WSMessage) {
logger.Debug("Received exit node disconnect message: %v", msg.Data)
if !o.tunnelRunning {
logger.Debug("Tunnel stopped, ignoring exit node disconnect message")
return
}
if err := o.disconnectExitNode(); err != nil {
logger.Error("Failed to disconnect from exit node: %v", err)
}
}
// handleExitNodeUpdateData handles a server-initiated request to change data
// associated with the currently connected exit node, such as its aliases (e.g. a
// resource was renamed). Unlike site aliases, there is no per-alias address to
// track since every exit node alias resolves to the exit node's own ServerIP.
func (o *Olm) handleExitNodeUpdateData(msg websocket.WSMessage) {
logger.Debug("Received exit node update data message: %v", msg.Data)
if !o.tunnelRunning {
logger.Debug("Tunnel stopped, ignoring exit node update data message")
return
}
jsonData, err := json.Marshal(msg.Data)
if err != nil {
logger.Error("Error marshaling exit node update data: %v", err)
return
}
var update ExitNodeUpdateData
if err := json.Unmarshal(jsonData, &update); err != nil {
logger.Error("Error unmarshaling exit node update data: %v", err)
return
}
o.exitNodeMu.Lock()
defer o.exitNodeMu.Unlock()
if o.exitNode == nil {
logger.Debug("Ignoring exit node update data message: no exit node connected")
return
}
serverIP := net.ParseIP(o.exitNode.ServerIP)
// Add new aliases BEFORE removing old ones, same as site aliases, so a rename
// that keeps the same underlying address never has a gap in resolution.
if o.dnsProxy != nil && serverIP != nil {
for _, alias := range update.NewAliases {
if err := o.dnsProxy.AddDNSRecord(alias, serverIP, exitNodeAliasSiteId); err != nil {
logger.Warn("Failed to add DNS record for exit node alias %s: %v", alias, err)
}
}
}
if o.dnsProxy != nil && serverIP != nil {
for _, alias := range update.OldAliases {
o.dnsProxy.RemoveDNSRecordForSite(alias, serverIP, exitNodeAliasSiteId)
}
}
o.exitNode.Aliases = applyStringListUpdate(o.exitNode.Aliases, update.OldAliases, update.NewAliases)
logger.Info("Successfully updated exit node data")
}
// applyStringListUpdate returns list with every entry in removed dropped and every
// entry in added appended, preserving the add-before-remove semantics of the caller.
func applyStringListUpdate(list, removed, added []string) []string {
next := make([]string, 0, len(list)+len(added))
next = append(next, list...)
next = append(next, added...)
removedSet := make(map[string]struct{}, len(removed))
for _, alias := range removed {
removedSet[alias] = struct{}{}
}
filtered := next[:0]
for _, alias := range next {
if _, ok := removedSet[alias]; ok {
continue
}
filtered = append(filtered, alias)
}
return filtered
}
+115 -15
View File
@@ -4,6 +4,7 @@ import (
"context"
"crypto/rand"
"encoding/hex"
"encoding/json"
"fmt"
"net"
"net/http"
@@ -15,6 +16,7 @@ import (
"github.com/fosrl/newt/bind"
"github.com/fosrl/newt/clients/permissions"
"github.com/fosrl/newt/exitnode"
"github.com/fosrl/newt/holepunch"
"github.com/fosrl/newt/logger"
"github.com/fosrl/newt/network"
@@ -55,6 +57,11 @@ type Olm struct {
holePunchManager *holepunch.Manager
peerManager *peers.PeerManager
peerManagerMu sync.RWMutex
// exitNode tracks the currently connected exit node peer, if any. It lives on a
// secondary address on the same interface/WireGuard device as the site peers.
exitNode *ExitNodeConfig
exitNodeMu sync.Mutex
// Power mode management
currentPowerMode string
powerModeMu sync.Mutex
@@ -75,6 +82,11 @@ type Olm struct {
stopRegister func()
updateRegister func(newData any)
// Exit node ping dance, run before registration so the server can pick
// the best exit node (mirrors newt's newt/ping/request flow).
stopPingRequest func()
pendingPingChainId string
stopPeerSends map[string]func()
stopPeerInits map[string]func()
jitPendingSites map[int]string // siteId -> chainId for in-flight JIT requests
@@ -550,6 +562,66 @@ func (o *Olm) StartTunnel(config TunnelConfig) {
o.websocket.RegisterHandler("olm/wg/peer/chain/cancel", o.handleCancelChain)
o.websocket.RegisterHandler("olm/sync", o.handleSync)
// Handlers for the server to direct connecting/disconnecting an exit node after registration
o.websocket.RegisterHandler("olm/wg/exitnode/connect", o.handleExitNodeConnect)
o.websocket.RegisterHandler("olm/wg/exitnode/disconnect", o.handleExitNodeDisconnect)
o.websocket.RegisterHandler("olm/wg/exitnode/data/update", o.handleExitNodeUpdateData)
o.websocket.RegisterHandler("olm/ping/exitNodes", func(msg websocket.WSMessage) {
logger.Debug("Received exit node ping request")
if o.stopPingRequest != nil {
o.stopPingRequest()
o.stopPingRequest = nil
}
if !o.tunnelRunning {
logger.Debug("Tunnel is no longer running, skipping exit node ping")
return
}
var exitNodeData exitnode.ExitNodeData
jsonData, err := json.Marshal(msg.Data)
if err != nil {
logger.Error("Error marshaling exit node data: %v", err)
return
}
if err := json.Unmarshal(jsonData, &exitNodeData); err != nil {
logger.Error("Error unmarshaling exit node data: %v", err)
return
}
if exitNodeData.ChainId != "" {
if exitNodeData.ChainId != o.pendingPingChainId {
logger.Debug("Discarding duplicate/stale olm/ping/exitNodes (chainId=%s, expected=%s)", exitNodeData.ChainId, o.pendingPingChainId)
return
}
o.pendingPingChainId = ""
}
if len(exitNodeData.ExitNodes) == 0 {
logger.Info("No exit nodes provided")
return
}
pingResults := exitnode.PingExitNodes(exitNodeData.ExitNodes, "", false)
publicKey := o.privateKey.PublicKey()
logger.Debug("Sending registration message to server with public key: %s, relay: %v, pingResults: %+v", publicKey, !config.Holepunch, pingResults)
o.stopRegister, o.updateRegister = o.websocket.SendMessageInterval("olm/wg/register", map[string]any{
"publicKey": publicKey.String(),
"relay": !config.Holepunch,
"olmVersion": o.olmConfig.Version,
"olmAgent": o.olmConfig.Agent,
"orgId": config.OrgID,
"userToken": userToken,
"fingerprint": o.fingerprint,
"postures": o.postures,
"pingResults": pingResults,
"chainId": generateChainId(), // use a random chainId for registration updates - it won't be used for cancellation since registration is a one-time message but for tracking the session
}, 2*time.Second, 20) // after 18 tries on the server side we send the error so dont change this without changing that
})
o.websocket.OnConnect(func() error {
logger.Info("Websocket Connected")
@@ -568,8 +640,6 @@ func (o *Olm) StartTunnel(config TunnelConfig) {
return nil
}
publicKey := o.privateKey.PublicKey()
// delay for 500ms to allow for time for the hp to get processed
time.Sleep(500 * time.Millisecond)
@@ -579,19 +649,36 @@ func (o *Olm) StartTunnel(config TunnelConfig) {
return nil
}
if o.stopRegister == nil {
logger.Debug("Sending registration message to server with public key: %s and relay: %v", publicKey, !config.Holepunch)
o.stopRegister, o.updateRegister = o.websocket.SendMessageInterval("olm/wg/register", map[string]any{
"publicKey": publicKey.String(),
"relay": !config.Holepunch,
"olmVersion": o.olmConfig.Version,
"olmAgent": o.olmConfig.Agent,
"orgId": config.OrgID,
"userToken": userToken,
"fingerprint": o.fingerprint,
"postures": o.postures,
"chainId": generateChainId(), // use a random chainId for registration updates - it won't be used for cancellation since registration is a one-time message but for tracking the session
}, 2*time.Second, 20) // after 18 tries on the server side we send the error so dont change this without changing that
if o.stopRegister == nil && o.stopPingRequest == nil {
publicKey := o.privateKey.PublicKey()
pingChainId := generateChainId()
o.pendingPingChainId = pingChainId
logger.Debug("Requesting exit nodes from server for ping selection")
o.stopPingRequest, _ = o.websocket.SendMessageInterval("olm/ping/request", map[string]any{
"chainId": pingChainId,
}, 3*time.Second, 10)
// Backwards-compatible one-shot registration, with no pingResults,
// for servers that predate the exit node ping dance. Servers that
// support it ignore backwardsCompatible register messages (see
// handleOlmRegisterMessage server-side) and wait for the real
// registration sent from the olm/ping/exitNodes handler above.
bcChainId := generateChainId()
if err := o.websocket.SendMessage("olm/wg/register", map[string]any{
"publicKey": publicKey.String(),
"relay": !config.Holepunch,
"olmVersion": o.olmConfig.Version,
"olmAgent": o.olmConfig.Agent,
"orgId": config.OrgID,
"userToken": userToken,
"fingerprint": o.fingerprint,
"postures": o.postures,
"backwardsCompatible": true,
"chainId": bcChainId,
}); err != nil {
logger.Error("Failed to send registration message: %v", err)
}
// Invoke onRegistered callback if configured
if o.olmConfig.OnRegistered != nil {
@@ -688,6 +775,12 @@ func (o *Olm) Close() {
o.stopRegister = nil
}
if o.stopPingRequest != nil {
logger.Debug("Stopping exit node ping request interval")
o.stopPingRequest()
o.stopPingRequest = nil
}
// Stop all pending peer init and send senders before closing websocket
o.peerSendMu.Lock()
for _, stop := range o.stopPeerInits {
@@ -744,6 +837,13 @@ func (o *Olm) Close() {
}
o.peerManagerMu.Unlock()
// The WireGuard device and TUN interface are being torn down below, which takes
// the exit node peer and its secondary address with them - just clear the
// in-memory record so a stale config isn't reused on the next connect.
o.exitNodeMu.Lock()
o.exitNode = nil
o.exitNodeMu.Unlock()
if o.uapiListener != nil {
_ = o.uapiListener.Close()
o.uapiListener = nil
+28
View File
@@ -10,6 +10,34 @@ type WgData struct {
Sites []peers.SiteConfig `json:"sites"`
TunnelIP string `json:"tunnelIP"`
UtilitySubnet string `json:"utilitySubnet"` // this is for things like the DNS server, and alias addresses
ExitNode *ExitNodeConfig `json:"exitNode,omitempty"`
}
// ExitNodeConfig describes an exit node the olm client can connect to for
// resources (e.g. inference) hosted on that node, separate from the site
// peers. It lives in a different address space than the site tunnel - the
// client is assigned TunnelIP (within the exit node's subnet) to reach the
// node at ServerIP. It arrives on the initial "olm/wg/connect" message and can
// also be sent later via "olm/wg/exitnode/connect" / "olm/wg/exitnode/disconnect"
// so the server can direct a client to connect/disconnect after registration.
type ExitNodeConfig struct {
Connect bool `json:"connect"`
Endpoint string `json:"endpoint"`
RelayPort uint16 `json:"relayPort"`
PublicKey string `json:"publicKey"`
ServerIP string `json:"serverIP"`
TunnelIP string `json:"tunnelIP"`
Aliases []string `json:"aliases,omitempty"`
}
// ExitNodeUpdateData describes a change to data associated with the currently
// connected exit node, e.g. when a resource's alias is renamed on the server.
// Aliases have no per-alias address here since every exit node alias resolves
// to the exit node's own ServerIP. More fields can be added here in the
// future as other exit node data becomes updatable.
type ExitNodeUpdateData struct {
OldAliases []string `json:"oldAliases,omitempty"`
NewAliases []string `json:"newAliases,omitempty"`
}
type SyncData struct {