diff --git a/olm/connect.go b/olm/connect.go index 4b4409d..48748f7 100644 --- a/olm/connect.go +++ b/olm/connect.go @@ -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 diff --git a/olm/gateway.go b/olm/gateway.go index 962e3c0..02661db 100644 --- a/olm/gateway.go +++ b/olm/gateway.go @@ -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,63 @@ 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 +} + +// 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() } diff --git a/olm/olm.go b/olm/olm.go index f1a79b5..ef1215f 100644 --- a/olm/olm.go +++ b/olm/olm.go @@ -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() @@ -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 diff --git a/peers/manager.go b/peers/manager.go index fc0f866..19d2f69 100644 --- a/peers/manager.go +++ b/peers/manager.go @@ -105,10 +105,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,13 +115,13 @@ 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 @@ -323,7 +320,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 +345,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 +// resolving and excluding every currently-tracked peer's active endpoint and +// every extra bypass endpoint (including the control plane - see +// gatewayExtraEndpoints) (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) - } - +func (pm *PeerManager) activateGatewayLocked() error { for _, peer := range pm.peers { if ip, ok := pm.resolveActiveEndpointIPLocked(peer); ok { pm.excludeEndpointLocked(ip) @@ -381,8 +372,7 @@ func (pm *PeerManager) activateGatewayLocked(controlEndpointHost string) error { } // 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 +// default-route-equivalent, then every remaining bypass route, and resets gateway exclusion state. Must be called with // pm.mu held. func (pm *PeerManager) deactivateGatewayLocked() { if err := network.RemoveGatewayDefaultRoute(pm.interfaceName); err != nil { @@ -395,7 +385,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 +439,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 +466,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 +637,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). diff --git a/websocket/client.go b/websocket/client.go index ecf04d8..71a9e73 100644 --- a/websocket/client.go +++ b/websocket/client.go @@ -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 != "" { diff --git a/websocket/client_test.go b/websocket/client_test.go index 121dfbc..569dbc2 100644 --- a/websocket/client_test.go +++ b/websocket/client_test.go @@ -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) + } +}