mirror of
https://github.com/fosrl/olm.git
synced 2026-10-02 02:39:10 +02:00
Merge pull request #158 from fosrl/dev
Sync routes on host with when interface changes and include ws ips
This commit is contained in:
+9
-1
@@ -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),
|
||||
}
|
||||
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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=
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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 != "" {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user