Merge pull request #158 from fosrl/dev

Sync routes on host with when interface changes and include ws ips
This commit is contained in:
Owen Schwartz
2026-10-01 16:07:29 -04:00
committed by GitHub
10 changed files with 357 additions and 72 deletions
+9 -1
View File
@@ -713,6 +713,14 @@ func (p *DNSProxy) dialTunnel(network, addr string) (net.Conn, uint16, error) {
return nil, 0, err
}
// The tunnel netstack only has an IPv4 address and route (the WireGuard
// interface IP), so an IPv6 upstream server can't be reached through it.
// To4() is nil for those, and converting that to a [4]byte below would panic.
raddrIP := raddr.IP.To4()
if raddrIP == nil {
return nil, 0, fmt.Errorf("upstream DNS server %s is not an IPv4 address, only IPv4 is supported over the tunnel", addr)
}
// Use tunnel IP as source
ipBytes := p.tunnelIP.As4()
@@ -725,7 +733,7 @@ func (p *DNSProxy) dialTunnel(network, addr string) (net.Conn, uint16, error) {
raddrTcpip := &tcpip.FullAddress{
NIC: 1,
Addr: tcpip.AddrFrom4([4]byte(raddr.IP.To4())),
Addr: tcpip.AddrFrom4([4]byte(raddrIP)),
Port: uint16(raddr.Port),
}
+16
View File
@@ -2,9 +2,11 @@ package dns
import (
"net"
"net/netip"
"testing"
"github.com/miekg/dns"
"gvisor.dev/gvisor/pkg/tcpip/stack"
)
func TestCheckLocalRecordsNODATAForAAAA(t *testing.T) {
@@ -176,3 +178,17 @@ func TestCheckLocalRecordsNODATAWildcard(t *testing.T) {
t.Fatalf("Expected 1 answer, got %d", len(response.Answer))
}
}
func TestDialTunnelRejectsIPv6Upstream(t *testing.T) {
proxy := &DNSProxy{
tunnelStack: stack.New(stack.Options{}),
tunnelIP: netip.MustParseAddr("100.90.128.1"),
tunnelActivePorts: make(map[uint16]bool),
}
defer proxy.tunnelStack.Close()
// Must return an error rather than panic: the tunnel netstack is IPv4-only
if _, _, err := proxy.dialTunnel("udp", "[2606:4700:4700::1111]:53"); err == nil {
t.Fatal("Expected error dialing an IPv6 upstream through the tunnel, got nil")
}
}
+2 -2
View File
@@ -4,7 +4,7 @@ go 1.26.0
require (
github.com/Microsoft/go-winio v0.6.2
github.com/fosrl/newt v1.18.0
github.com/fosrl/newt v1.18.1
github.com/godbus/dbus/v5 v5.2.2
github.com/google/nftables v0.3.0
github.com/gorilla/websocket v1.5.3
@@ -29,7 +29,7 @@ require (
golang.org/x/sync v0.22.0 // indirect
golang.org/x/time v0.12.0 // indirect
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 // indirect
golang.zx2c4.com/wireguard/windows v1.0.1 // indirect
golang.zx2c4.com/wireguard/windows v1.1.1 // indirect
)
// To be used ONLY for local development
+4 -4
View File
@@ -1,7 +1,7 @@
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.18.0 h1:fFBksIV6BoI+BGmv6fwo87tXBk1WlxbVoruvupeD6hk=
github.com/fosrl/newt v1.18.0/go.mod h1:CwcuQtifgDQeSWSEB3yfqOgheFv9yltVBYQNgDB/DoM=
github.com/fosrl/newt v1.18.1 h1:gKWBDKSqQLSgD6z0jPzZB86U19a2xNfOTV0S1KdYSx0=
github.com/fosrl/newt v1.18.1/go.mod h1:gTr7YDmLqjVrqZJ3SQU0MR9wKfqbUZOJF8hSFIi/M+I=
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=
@@ -42,8 +42,8 @@ golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb h1:whnFRlWMcXI9d+Z
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb/go.mod h1:rpwXGsirqLqN2L0JDJQlwOboGHmptD5ZD6T2VmcqhTw=
golang.zx2c4.com/wireguard/wgctrl v0.0.0-20241231184526-a9ab2273dd10 h1:3GDAcqdIg1ozBNLgPy4SLT84nfcBjr6rhGtXYtrkWLU=
golang.zx2c4.com/wireguard/wgctrl v0.0.0-20241231184526-a9ab2273dd10/go.mod h1:T97yPqesLiNrOYxkwmhMI0ZIlJDm+p0PMR8eRVeR5tQ=
golang.zx2c4.com/wireguard/windows v1.0.1 h1:eOxiDVbywPC+ZQqvdCK7x+ZwWXKbYv50TtH8ysFIbw8=
golang.zx2c4.com/wireguard/windows v1.0.1/go.mod h1:+fbT3FFdX4zzYDLwJh5+HPEcNN/3HyNdzhNSVsQM+zs=
golang.zx2c4.com/wireguard/windows v1.1.1 h1:8/H97U1v1PNDNcBsMZgU3KFuND9MQdTsU2NOwmCXArE=
golang.zx2c4.com/wireguard/windows v1.1.1/go.mod h1:+fbT3FFdX4zzYDLwJh5+HPEcNN/3HyNdzhNSVsQM+zs=
gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c h1:m/r7OM+Y2Ty1sgBQ7Qb27VgIMBW8ZZhT4gLnUyDIhzI=
gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c/go.mod h1:3r5CMtNQMKIvBlrmM9xWUNamjKBYPOWyXOjmg5Kts3g=
software.sslmate.com/src/go-pkcs12 v0.7.3 h1:JBQD3FDqYjTeyDAeZQklj2ar88ykBLtALloPJHyAauU=
+1
View File
@@ -274,6 +274,7 @@ func (o *Olm) handleConnect(msg websocket.WSMessage) {
// endpoints it already recorded now that there's somewhere to put them.
o.flushPendingHolepunchBypassEndpoints()
o.flushPendingDNSBypassEndpoints()
o.flushPendingControlBypassEndpoints()
if o.dnsProxy != nil {
if err := o.dnsProxy.Start(); err != nil { // start DNS proxy first so there is no downtime
+68 -14
View File
@@ -3,7 +3,7 @@ package olm
import (
"encoding/json"
"fmt"
"net/url"
"net"
"github.com/fosrl/newt/logger"
"github.com/fosrl/olm/peers"
@@ -55,8 +55,7 @@ func (o *Olm) DisableGateway() error {
return nil
}
// applySelectGateway resolves the control-plane endpoint host and delegates
// to the peer manager. Shared by SelectGateway (API-invoked, already
// applySelectGateway delegates to the peer manager. Shared by SelectGateway (API-invoked, already
// registered-checked by the caller) and applyPendingGatewayConfig
// (StartTunnel-time, called after registration completes).
func (o *Olm) applySelectGateway(siteResourceId int, siteIds []int) error {
@@ -64,7 +63,7 @@ func (o *Olm) applySelectGateway(siteResourceId int, siteIds []int) error {
if pm == nil {
return fmt.Errorf("cannot select gateway: tunnel not running")
}
if err := pm.SetGateway(siteResourceId, siteIds, extractControlEndpointHost(o.tunnelConfig.Endpoint)); err != nil {
if err := pm.SetGateway(siteResourceId, siteIds); err != nil {
return err
}
o.apiServer.SetGatewayStatus(true, siteResourceId, siteIds)
@@ -274,15 +273,70 @@ func (o *Olm) flushPendingDNSBypassEndpoints() {
}
}
// extractControlEndpointHost returns the bare host (no scheme/port) of the
// Pangolin server olm is registered against, for gateway bypass-route
// purposes. Falls back to the raw endpoint string on parse failure -
// resolveEndpointIPLocked/net.SplitHostPort tolerate a bare host - rather
// than failing the whole gateway activation over a cosmetic parse issue.
func extractControlEndpointHost(endpoint string) string {
u, err := url.Parse(endpoint)
if err != nil || u.Hostname() == "" {
return endpoint
// updateControlBypassEndpoints is the websocket client's OnDialTargets
// callback: ips are every address the Pangolin server host resolved to,
// about to be dialed for the token request or websocket. It diffs them
// against the currently-registered set and adds/removes gateway bypass routes
// for the difference, via the same AddGatewayBypassEndpoint machinery used
// for hole-punch and DNS endpoints above. Runs synchronously before the dial,
// so the control connection is pinned to the physical network before its
// first packet - whether or not gateway mode is active yet. Replacing the set
// on each dial is safe because the websocket client only dials while it has
// no live connection (initial connect, or reconnect after tearing the old one
// down); a DNS change while connected has no effect on the established
// connection and is picked up on the next dial. Only IPv4 addresses are
// registered: the gateway route only captures IPv4, and bypass routes are
// IPv4 host routes. Safe to call before the peer manager exists (see
// flushPendingControlBypassEndpoints).
func (o *Olm) updateControlBypassEndpoints(ips []string) {
pm := o.getPeerManager()
newBypassEndpoints := make(map[string]bool, len(ips))
for _, ip := range ips {
if parsed := net.ParseIP(ip); parsed != nil && parsed.To4() != nil {
newBypassEndpoints[ip] = true
}
}
o.controlBypassMu.Lock()
defer o.controlBypassMu.Unlock()
if pm != nil {
for ip := range newBypassEndpoints {
if !o.controlBypassEndpoints[ip] {
pm.AddGatewayBypassEndpoint(ip)
}
}
for ip := range o.controlBypassEndpoints {
if !newBypassEndpoints[ip] {
pm.RemoveGatewayBypassEndpoint(ip)
}
}
}
o.controlBypassEndpoints = newBypassEndpoints
// A (re)connect is a good moment to check the bypass routes: the
// previous connection may have dropped because the network changed and
// took them with it.
if pm != nil {
pm.ReconcileGatewayBypassRoutes()
}
}
// flushPendingControlBypassEndpoints re-registers every currently-known
// control-plane bypass endpoint with the peer manager. Mirrors
// flushPendingHolepunchBypassEndpoints: the websocket's first dial happens
// well before handleConnect creates the peer manager, so whatever was
// recorded needs to be pushed in once it becomes available.
// AddGatewayBypassEndpoint is idempotent.
func (o *Olm) flushPendingControlBypassEndpoints() {
pm := o.getPeerManager()
if pm == nil {
return
}
o.controlBypassMu.Lock()
defer o.controlBypassMu.Unlock()
for ip := range o.controlBypassEndpoints {
pm.AddGatewayBypassEndpoint(ip)
}
return u.Hostname()
}
+38 -9
View File
@@ -93,6 +93,16 @@ type Olm struct {
dnsBypassEndpoints map[string]bool
dnsBypassMu sync.Mutex
// controlBypassEndpoints tracks the Pangolin server (control-plane) IPs
// currently registered as gateway bypass targets: every IPv4 address the
// server host resolved to on the websocket client's most recent dial (see
// websocket.Client.OnDialTargets). Replaced on every dial, before any
// packet is sent, so the websocket and token requests are never captured
// by the gateway default-route-equivalent - including after the server's
// DNS record changes. Mirrors hpBypassEndpoints.
controlBypassEndpoints map[string]bool
controlBypassMu sync.Mutex
// primaryTunnelIP is the site tunnel's own address (wgData.TunnelIP), set once
// per connect in handleConnect. It's the interface's first/primary address -
// on macOS/iOS NetworkExtension, an unbound outbound socket's source gets
@@ -279,14 +289,15 @@ func Init(ctx context.Context, config OlmConfig) (*Olm, error) {
apiServer.SetAgent(config.Agent)
newOlm := &Olm{
logFile: logFile,
olmCtx: ctx,
apiServer: apiServer,
olmConfig: config,
stopPeerSends: make(map[string]func()),
stopPeerInits: make(map[string]func()),
jitPendingSites: make(map[int]string),
hpBypassEndpoints: make(map[string]bool),
logFile: logFile,
olmCtx: ctx,
apiServer: apiServer,
olmConfig: config,
stopPeerSends: make(map[string]func()),
stopPeerInits: make(map[string]func()),
jitPendingSites: make(map[int]string),
hpBypassEndpoints: make(map[string]bool),
controlBypassEndpoints: make(map[string]bool),
}
newOlm.registerAPICallbacks()
@@ -554,7 +565,7 @@ func (o *Olm) StartTunnel(config TunnelConfig) {
// Fall back to hardcoded DNS if the system monitor could not detect any.
if len(o.tunnelConfig.PublicDNS) == 0 {
if o.tunnelConfig.DNS != "" {
o.tunnelConfig.PublicDNS = []string{o.tunnelConfig.DNS + ":53"}
o.tunnelConfig.PublicDNS = []string{net.JoinHostPort(o.tunnelConfig.DNS, "53")}
} else {
o.tunnelConfig.PublicDNS = []string{"8.8.8.8:53"}
}
@@ -617,6 +628,12 @@ func (o *Olm) StartTunnel(config TunnelConfig) {
"postures": o.postures,
}
}),
// Resolve the server host via the physical network's DNS, the same
// way WireGuard endpoints are, rather than through olm's own DNS
// override.
websocket.WithPublicDNSProvider(func() []string {
return o.tunnelConfig.PublicDNS
}),
)
if err != nil {
logger.Error("Failed to create olm: %v", err)
@@ -790,6 +807,8 @@ func (o *Olm) StartTunnel(config TunnelConfig) {
return nil
})
o.websocket.OnDialTargets(o.updateControlBypassEndpoints)
o.websocket.OnTokenUpdate(func(token string, exitNodes []websocket.ExitNode) {
// Check if tunnel is still running and hole punch manager exists
if !o.tunnelRunning || o.holePunchManager == nil {
@@ -995,6 +1014,10 @@ func (o *Olm) Close() {
o.hpBypassEndpoints = make(map[string]bool)
o.hpBypassMu.Unlock()
o.controlBypassMu.Lock()
o.controlBypassEndpoints = make(map[string]bool)
o.controlBypassMu.Unlock()
if o.uapiListener != nil {
_ = o.uapiListener.Close()
o.uapiListener = nil
@@ -1383,6 +1406,12 @@ func (o *Olm) RebindSocket() error {
logger.Info("Successfully rebound UDP socket on port %d", newPort)
// A rebind is requested on network changes; the bypass routes may need
// to follow the new physical path too.
if pm := o.getPeerManager(); pm != nil {
pm.ReconcileGatewayBypassRoutes()
}
// Check if we're in low power mode before triggering hole punch
o.powerModeMu.Lock()
isLowPower := o.currentPowerMode == "low"
+110 -36
View File
@@ -1,6 +1,8 @@
package peers
import (
"context"
"errors"
"fmt"
"net"
"sort"
@@ -105,10 +107,7 @@ type PeerManager struct {
// allowedIPClaims (see gatewayCIDR). gatewayExcludedIPs is a refcount so a
// destination referenced by more than one thing (or re-added while
// already excluded) is only actually un-excluded once nothing references
// it any more. gatewayControlIP is the control-plane (Pangolin server)
// endpoint, resolved once at activation and excluded for the lifetime of
// the gateway so the management connection doesn't depend on the current
// gateway-owner site's own uplink. gatewaySiteResourceId is the numeric
// it any more. gatewaySiteResourceId is the numeric
// ID (not the niceId, which can be renamed) of the gateway-mode site
// resource the candidate set was selected from, so server-pushed
// add/remove/disable updates (see UpdateGatewaySites) are only applied when
@@ -118,20 +117,29 @@ type PeerManager struct {
gatewaySiteResourceId int
gatewaySiteIds map[int]bool
gatewayExcludedIPs map[string]int
gatewayControlIP string
// gatewayExtraEndpoints tracks "host:port" (or already-resolved "ip:port")
// endpoints registered by callers outside the normal site-peer lifecycle -
// currently hole-punch exit node probing endpoints and a connected exit
// node's own WireGuard endpoint (see olm's OnTokenUpdate handler and
// connectExitNode) - that must stay off the gateway route the same way a
// currently hole-punch exit node probing endpoints, a connected exit
// node's own WireGuard endpoint, upstream DNS servers, and the control-plane
// (Pangolin server) IPs the websocket dials (see olm's OnTokenUpdate and
// OnDialTargets handlers and connectExitNode) - that must stay off the gateway route the same way a
// site peer's own endpoint does, so hole punching / the exit node
// connection still originates from the local network rather than looping
// through the tunnel. Business intent, independent of gatewayActive - see
// AddGatewayBypassEndpoint/RemoveGatewayBypassEndpoint.
gatewayExtraEndpoints map[string]bool
// gatewayRouteWatchCancel stops the goroutines started by
// startBypassRouteWatcherLocked; nil while gateway mode is inactive.
gatewayRouteWatchCancel context.CancelFunc
}
// bypassRouteReconcileInterval is how often bypass routes are re-checked
// while gateway mode is active, as a fallback for route changes the OS
// watcher (network.WatchRouteChanges) misses.
const bypassRouteReconcileInterval = 30 * time.Second
// gatewayCIDR is the WireGuard AllowedIPs claim key for "this site is the
// gateway (full-tunnel/default-route)". It is never added to
// SiteConfig.RemoteSubnets/AllowedIps and never sent over the wire - it only
@@ -323,7 +331,7 @@ func (pm *PeerManager) excludeEndpointLocked(ip string) {
return
}
if pm.gatewayExcludedIPs[ip] == 0 {
if err := network.AddBypassRouteForDestination(ip); err != nil {
if err := network.AddBypassRouteForDestination(ip, pm.interfaceName); err != nil {
logger.Error("Gateway: failed to add bypass route for %s: %v", ip, err)
}
}
@@ -348,19 +356,13 @@ func (pm *PeerManager) unexcludeEndpointLocked(ip string) {
}
// activateGatewayLocked performs the one-time setup for gateway mode:
// resolving and excluding the control-plane endpoint and every
// currently-tracked peer's active endpoint (so none of them can be captured
// by the default-route-equivalent installed at the end), then installing
// that route. Must be called with pm.mu held, and only once (guarded by
// pm.gatewayActive in the caller).
func (pm *PeerManager) activateGatewayLocked(controlEndpointHost string) error {
if ip, ok := pm.resolveEndpointIPLocked(controlEndpointHost); ok {
pm.gatewayControlIP = ip
pm.excludeEndpointLocked(ip)
} else if controlEndpointHost != "" {
logger.Warn("Gateway: failed to resolve control endpoint %q for bypass route", controlEndpointHost)
}
// resolving and excluding every currently-tracked peer's active endpoint and
// every extra bypass endpoint (see gatewayExtraEndpoints), so none of them
// can be captured by the default-route-equivalent installed at the end, then
// installing that route and starting to watch for network changes that need
// the bypass routes reconciled. Must be called with pm.mu held, and only once
// (guarded by pm.gatewayActive in the caller).
func (pm *PeerManager) activateGatewayLocked() error {
for _, peer := range pm.peers {
if ip, ok := pm.resolveActiveEndpointIPLocked(peer); ok {
pm.excludeEndpointLocked(ip)
@@ -377,14 +379,89 @@ func (pm *PeerManager) activateGatewayLocked(controlEndpointHost string) error {
return fmt.Errorf("failed to install gateway route: %v", err)
}
pm.startBypassRouteWatcherLocked()
return nil
}
// deactivateGatewayLocked reverses activateGatewayLocked: removes the
// default-route-equivalent, then every remaining bypass route (including the
// control endpoint), and resets gateway exclusion state. Must be called with
// startBypassRouteWatcherLocked keeps the bypass routes in step with the
// physical network for as long as gateway mode is active: the OS removes
// them along with the interface or address they were on when the network
// changes (e.g. switching Wi-Fi networks, Wi-Fi to Ethernet), and without
// them the tunnel's own traffic would be captured by the gateway route.
// Reconciles on every OS route change notification, and periodically as a
// fallback. A no-op where bypasses are applied by the host app from
// NetworkSettings (mobile, macOS NetworkExtension), which tracks the physical
// network itself. Must be called with pm.mu held.
func (pm *PeerManager) startBypassRouteWatcherLocked() {
if !network.ManagesHostRoutes() || pm.gatewayRouteWatchCancel != nil {
return
}
ctx, cancel := context.WithCancel(context.Background())
pm.gatewayRouteWatchCancel = cancel
if err := network.WatchRouteChanges(ctx, pm.ReconcileGatewayBypassRoutes); err != nil {
logger.Warn("Gateway: failed to watch for route changes, bypass routes will only be re-checked every %v: %v", bypassRouteReconcileInterval, err)
}
go func() {
ticker := time.NewTicker(bypassRouteReconcileInterval)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
pm.ReconcileGatewayBypassRoutes()
}
}
}()
}
// stopBypassRouteWatcherLocked reverses startBypassRouteWatcherLocked. A
// reconcile already waiting on pm.mu may still run once afterwards, which is
// harmless: it is a no-op once gateway mode is inactive. Must be called with
// pm.mu held.
func (pm *PeerManager) stopBypassRouteWatcherLocked() {
if pm.gatewayRouteWatchCancel != nil {
pm.gatewayRouteWatchCancel()
pm.gatewayRouteWatchCancel = nil
}
}
// ReconcileGatewayBypassRoutes re-checks every bypass route while gateway
// mode is active, re-adding any the OS has dropped and moving any that no
// longer follow the current physical path (see network.ReconcileBypassRoute).
// Also repairs routes that failed to install in the first place (e.g. added
// while offline). Idempotent and cheap when nothing has changed; safe to call
// at any time.
func (pm *PeerManager) ReconcileGatewayBypassRoutes() {
pm.mu.Lock()
defer pm.mu.Unlock()
if !pm.gatewayActive {
return
}
for ip := range pm.gatewayExcludedIPs {
changed, err := network.ReconcileBypassRoute(ip, pm.interfaceName)
switch {
case errors.Is(err, network.ErrNoPhysicalRoute):
logger.Debug("Gateway: no physical route for bypass %s yet: %v", ip, err)
case err != nil:
logger.Warn("Gateway: failed to reconcile bypass route for %s: %v", ip, err)
case changed:
logger.Info("Gateway: restored bypass route for %s after a network change", ip)
}
}
}
// deactivateGatewayLocked reverses activateGatewayLocked: stops watching
// for network changes, removes the default-route-equivalent, then every
// remaining bypass route, and resets gateway exclusion state. Must be called
// with pm.mu held.
func (pm *PeerManager) deactivateGatewayLocked() {
pm.stopBypassRouteWatcherLocked()
if err := network.RemoveGatewayDefaultRoute(pm.interfaceName); err != nil {
logger.Error("Gateway: failed to remove gateway route: %v", err)
}
@@ -395,7 +472,6 @@ func (pm *PeerManager) deactivateGatewayLocked() {
}
}
pm.gatewayExcludedIPs = make(map[string]int)
pm.gatewayControlIP = ""
}
// claimGatewayClaimLocked registers siteId's claim to the gateway CIDR via
@@ -450,13 +526,12 @@ func (pm *PeerManager) releaseGatewayClaimLocked(siteId int) {
// UpdateGatewaySites). Every siteId must already be a tracked peer, or the
// call is rejected outright (no partial application). On first activation this
// installs the OS-level gateway route plus every bypass route needed so the
// tunnel's own traffic (control-plane endpoint, every tracked peer's active
// endpoint) isn't captured by it; subsequent calls only change which sites
// may own the "0.0.0.0/0" WireGuard AllowedIP, via the existing generic
// claim/optimizer machinery - exactly like remote subnets. controlEndpointHost
// is the Pangolin server host olm is registered against (bare host, port
// optional); always excluded regardless of which sites are selected.
func (pm *PeerManager) SetGateway(siteResourceId int, siteIds []int, controlEndpointHost string) error {
// tunnel's own traffic (every tracked peer's active endpoint, plus every
// extra bypass endpoint such as the control plane) isn't captured by it;
// subsequent calls only change which sites may own the "0.0.0.0/0" WireGuard
// AllowedIP, via the existing generic claim/optimizer machinery - exactly
// like remote subnets.
func (pm *PeerManager) SetGateway(siteResourceId int, siteIds []int) error {
pm.mu.Lock()
defer pm.mu.Unlock()
@@ -478,7 +553,7 @@ func (pm *PeerManager) SetGateway(siteResourceId int, siteIds []int, controlEndp
}
if !pm.gatewayActive {
if err := pm.activateGatewayLocked(controlEndpointHost); err != nil {
if err := pm.activateGatewayLocked(); err != nil {
return err
}
pm.gatewayActive = true
@@ -649,8 +724,7 @@ func (pm *PeerManager) ClearGateway() error {
// AddGatewayBypassEndpoint registers hostport (a "host:port" string, or an
// already-resolved "ip:port") as needing protection from the gateway
// default-route-equivalent, for endpoints outside the normal site-peer
// lifecycle - hole-punch exit node probing endpoints and a connected exit
// node's own WireGuard endpoint. If gateway mode is currently active, the
// lifecycle - see gatewayExtraEndpoints. If gateway mode is currently active, the
// bypass route is installed immediately; otherwise this only records intent,
// applied the next time gateway activates. Safe to call repeatedly with the
// same hostport (idempotent).
+71 -6
View File
@@ -3,12 +3,14 @@ package websocket
import (
"bytes"
"compress/gzip"
"context"
"crypto/tls"
"crypto/x509"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"net/http"
"net/url"
"os"
@@ -19,6 +21,7 @@ import (
"software.sslmate.com/src/go-pkcs12"
"github.com/fosrl/newt/logger"
"github.com/fosrl/newt/util"
"github.com/gorilla/websocket"
)
@@ -96,6 +99,8 @@ type Client struct {
onConnect func() error
onTokenUpdate func(token string, exitNodes []ExitNode)
onAuthError func(statusCode int, message string) // Callback for auth errors
onDialTargets func(ips []string) // Callback with the server IPs about to be dialed, see dialContext
publicDNS func() []string // Provides the DNS servers used to resolve the server host, see dialContext
writeMux sync.Mutex
clientType string // Type of client (e.g., "newt", "olm")
tlsConfig TLSConfig
@@ -156,6 +161,16 @@ func WithPingDataProvider(fn func() map[string]any) ClientOption {
}
}
// WithPublicDNSProvider sets a callback returning the DNS servers ("ip:port")
// used to resolve the server host when dialing (see dialContext). Called on
// every dial, so it may return a value that changes over time. If unset, or
// if every server fails, the platform resolver is used.
func WithPublicDNSProvider(fn func() []string) ClientOption {
return func(c *Client) {
c.publicDNS = fn
}
}
func (c *Client) OnConnect(callback func() error) {
c.onConnect = callback
}
@@ -168,6 +183,53 @@ func (c *Client) OnAuthError(callback func(statusCode int, message string)) {
c.onAuthError = callback
}
// OnDialTargets sets a callback invoked synchronously with every IP address
// the server host resolved to, immediately before each dial (token fetch and
// websocket). The callback runs before any packet is sent to those addresses,
// so it can install routes for them - see dialContext.
func (c *Client) OnDialTargets(callback func(ips []string)) {
c.onDialTargets = callback
}
// dialContext is the dial function for every connection to the server (token
// fetch and websocket, including through a proxy). It resolves the host
// itself rather than leaving it to net.Dialer, reports the resulting IPs via
// onDialTargets, and then dials only those IPs - so the caller knows exactly
// which addresses the connection can use and can keep them off the gateway
// (full-tunnel) default route before the first packet is sent. Each dial
// resolves again, so DNS changes are picked up on the next (re)connect; the
// address of a connection that is already established never changes.
func (c *Client) dialContext(ctx context.Context, network, addr string) (net.Conn, error) {
host, port, err := net.SplitHostPort(addr)
if err != nil {
return nil, err
}
var publicDNS []string
if c.publicDNS != nil {
publicDNS = c.publicDNS()
}
ips, err := util.ResolveDomainAllUpstream(host, publicDNS)
if err != nil {
return nil, fmt.Errorf("failed to resolve %s: %w", host, err)
}
if c.onDialTargets != nil {
c.onDialTargets(ips)
}
dialer := net.Dialer{Timeout: 10 * time.Second}
var lastErr error
for _, ip := range ips {
conn, err := dialer.DialContext(ctx, network, net.JoinHostPort(ip, port))
if err == nil {
return conn, nil
}
lastErr = err
}
return nil, lastErr
}
// NewClient creates a new websocket client
func NewClient(ID, secret, userToken, orgId, endpoint string, pingInterval time.Duration, opts ...ClientOption) (*Client, error) {
config := &Config{
@@ -483,12 +545,12 @@ func (c *Client) getToken() (string, []ExitNode, string, error) {
logger.Debug("websocket: Requesting token from %s with body: %s", req.URL.String(), string(jsonData))
// Make the request
client := &http.Client{}
transport := http.DefaultTransport.(*http.Transport).Clone()
transport.DialContext = c.dialContext
if tlsConfig != nil {
client.Transport = &http.Transport{
TLSClientConfig: tlsConfig,
}
transport.TLSClientConfig = tlsConfig
}
client := &http.Client{Transport: transport}
resp, err := client.Do(req)
if err != nil {
return "", nil, "", fmt.Errorf("failed to request new token: %w", err)
@@ -610,8 +672,11 @@ func (c *Client) establishConnection() error {
}
u.RawQuery = q.Encode()
// Connect to WebSocket
dialer := websocket.DefaultDialer
// Connect to WebSocket. Copy DefaultDialer rather than mutating the
// shared global.
dialerCopy := *websocket.DefaultDialer
dialer := &dialerCopy
dialer.NetDialContext = c.dialContext
// Use new TLS configuration method
if c.tlsConfig.ClientCertFile != "" || c.tlsConfig.ClientKeyFile != "" || len(c.tlsConfig.CAFiles) > 0 || c.tlsConfig.PKCS12File != "" {
+38
View File
@@ -1,6 +1,8 @@
package websocket
import (
"context"
"net"
"net/http"
"net/http/httptest"
"strings"
@@ -111,3 +113,39 @@ func TestReconnectClearsCurrentConnection(t *testing.T) {
t.Fatalf("expected conn to be cleared after reconnect(a) when a was still current, got %v", got)
}
}
// TestDialContextReportsTargetsBeforeDialing checks that dialContext reports
// the resolved server IPs via onDialTargets before connecting, and then
// connects to one of exactly those IPs - the contract olm relies on to install
// gateway bypass routes for the control connection.
func TestDialContextReportsTargetsBeforeDialing(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {}))
t.Cleanup(srv.Close)
_, port, err := net.SplitHostPort(strings.TrimPrefix(srv.URL, "http://"))
if err != nil {
t.Fatal(err)
}
c := newTestClient()
var reported []string
c.OnDialTargets(func(ips []string) {
if len(reported) != 0 {
t.Errorf("onDialTargets called more than once")
}
reported = append([]string(nil), ips...)
})
conn, err := c.dialContext(context.Background(), "tcp", net.JoinHostPort("127.0.0.1", port))
if err != nil {
t.Fatalf("dialContext: %v", err)
}
defer conn.Close()
if len(reported) != 1 || reported[0] != "127.0.0.1" {
t.Fatalf("onDialTargets got %v, want [127.0.0.1]", reported)
}
remoteHost, _, _ := net.SplitHostPort(conn.RemoteAddr().String())
if remoteHost != reported[0] {
t.Fatalf("dialed %s, but reported %v", remoteHost, reported)
}
}