Exclude the websocket routes

This commit is contained in:
Owen
2026-10-01 12:27:24 -04:00
parent 815f3e506c
commit 98fe56ed13
6 changed files with 221 additions and 60 deletions
+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
+61 -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,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()
}
+31 -8
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()
@@ -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
+19 -32
View File
@@ -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).
+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)
}
}