diff --git a/clients/clients.go b/clients/clients.go index 4568e09..42a5a54 100644 --- a/clients/clients.go +++ b/clients/clients.go @@ -190,13 +190,12 @@ func NewWireGuardService(interfaceName string, port uint16, mtu int, host string // Register websocket handlers wsClient.RegisterHandler("newt/wg/receive-config", service.handleConfig) - wsClient.RegisterHandler("newt/wg/peer/add", service.handleAddPeer) - wsClient.RegisterHandler("newt/wg/peer/remove", service.handleRemovePeer) - wsClient.RegisterHandler("newt/wg/peer/update", service.handleUpdatePeer) - wsClient.RegisterHandler("newt/wg/targets/add", service.handleAddTarget) - wsClient.RegisterHandler("newt/wg/targets/remove", service.handleRemoveTarget) - wsClient.RegisterHandler("newt/wg/targets/update", service.handleUpdateTarget) - wsClient.RegisterHandler("newt/wg/sync", service.handleSyncConfig) + // wsClient.RegisterHandler("newt/wg/peer/add", service.handleAddPeer) + // wsClient.RegisterHandler("newt/wg/peer/remove", service.handleRemovePeer) + // wsClient.RegisterHandler("newt/wg/peer/update", service.handleUpdatePeer) + // wsClient.RegisterHandler("newt/wg/targets/add", service.handleAddTarget) + // wsClient.RegisterHandler("newt/wg/targets/remove", service.handleRemoveTarget) + // wsClient.RegisterHandler("newt/wg/targets/update", service.handleUpdateTarget) return service, nil } @@ -568,37 +567,15 @@ func (s *WireGuardService) handleConfig(msg websocket.WSMessage) { logger.Info("Client connectivity setup. Ready to accept connections from clients!") } -// SyncConfig represents the configuration sent from server for syncing -type SyncConfig struct { - Targets []Target `json:"targets"` - Peers []Peer `json:"peers"` -} - -func (s *WireGuardService) handleSyncConfig(msg websocket.WSMessage) { - var syncConfig SyncConfig - - logger.Debug("Received sync message: %v", msg) - logger.Info("Received sync configuration from remote server") - - jsonData, err := json.Marshal(msg.Data) - if err != nil { - logger.Error("Error marshaling sync data: %v", err) - return +// Sync synchronizes the clients WireGuard peers and targets with the desired state +// received as part of the main newt/sync message. +func (s *WireGuardService) Sync(peers []Peer, targets []Target) { + if err := s.syncPeers(peers); err != nil { + logger.Error("Failed to sync client peers: %v", err) } - if err := json.Unmarshal(jsonData, &syncConfig); err != nil { - logger.Error("Error unmarshaling sync data: %v", err) - return - } - - // Sync peers - if err := s.syncPeers(syncConfig.Peers); err != nil { - logger.Error("Failed to sync peers: %v", err) - } - - // Sync targets - if err := s.syncTargets(syncConfig.Targets); err != nil { - logger.Error("Failed to sync targets: %v", err) + if err := s.syncTargets(targets); err != nil { + logger.Error("Failed to sync client targets: %v", err) } } @@ -665,8 +642,12 @@ func (s *WireGuardService) syncPeers(desiredPeers []Peer) error { return nil } -// syncTargets synchronizes the current targets with the desired state -// It removes targets not in the desired list and adds missing ones +// syncTargets synchronizes the current targets with the desired state. +// A sync represents the full authoritative state from the server, so rather +// than diffing against the currently installed rules (which can miss +// changes to a rule's contents when its source/dest prefix key is +// unchanged - e.g. a RewriteTo update), we just rebuild the entire rule set +// from scratch on every sync. This guarantees no stale rule can survive. func (s *WireGuardService) syncTargets(desiredTargets []Target) error { if s.tnet == nil { // Native interface mode - proxy features not available, skip silently @@ -674,70 +655,31 @@ func (s *WireGuardService) syncTargets(desiredTargets []Target) error { return nil } - // Get current rules from the proxy handler - currentRules := s.tnet.GetProxySubnetRules() - - // Build a map of current rules by source+dest prefix - type ruleKey struct { - sourcePrefix string - destPrefix string - } - currentRuleMap := make(map[ruleKey]bool) - for _, rule := range currentRules { - key := ruleKey{ - sourcePrefix: rule.SourcePrefix.String(), - destPrefix: rule.DestPrefix.String(), - } - currentRuleMap[key] = true - } - - // Build a map of desired targets - desiredTargetMap := make(map[ruleKey]Target) + var rules []netstack2.SubnetRule for _, target := range desiredTargets { - key := ruleKey{ - sourcePrefix: target.SourcePrefix, - destPrefix: target.DestPrefix, + destPrefix, err := netip.ParsePrefix(target.DestPrefix) + if err != nil { + logger.Warn("Invalid dest prefix %s during sync: %v", target.DestPrefix, err) + continue } - desiredTargetMap[key] = target - } - // Remove targets that are not in the desired list - for _, rule := range currentRules { - key := ruleKey{ - sourcePrefix: rule.SourcePrefix.String(), - destPrefix: rule.DestPrefix.String(), + var portRanges []netstack2.PortRange + for _, pr := range target.PortRange { + portRanges = append(portRanges, netstack2.PortRange{ + Min: pr.Min, + Max: pr.Max, + Protocol: pr.Protocol, + }) } - if _, exists := desiredTargetMap[key]; !exists { - s.tnet.RemoveProxySubnetRule(rule.SourcePrefix, rule.DestPrefix) - logger.Info("Removed target %s -> %s during sync", rule.SourcePrefix.String(), rule.DestPrefix.String()) - } - } - // Add targets that are missing - for key, target := range desiredTargetMap { - if _, exists := currentRuleMap[key]; !exists { - sourcePrefix, err := netip.ParsePrefix(target.SourcePrefix) + for _, sp := range resolveSourcePrefixes(target) { + sourcePrefix, err := netip.ParsePrefix(sp) if err != nil { - logger.Warn("Invalid source prefix %s during sync: %v", target.SourcePrefix, err) + logger.Warn("Invalid source prefix %s during sync: %v", sp, err) continue } - destPrefix, err := netip.ParsePrefix(target.DestPrefix) - if err != nil { - logger.Warn("Invalid dest prefix %s during sync: %v", target.DestPrefix, err) - continue - } - - var portRanges []netstack2.PortRange - for _, pr := range target.PortRange { - portRanges = append(portRanges, netstack2.PortRange{ - Min: pr.Min, - Max: pr.Max, - Protocol: pr.Protocol, - }) - } - - s.tnet.AddProxySubnetRule(netstack2.SubnetRule{ + rules = append(rules, netstack2.SubnetRule{ SourcePrefix: sourcePrefix, DestPrefix: destPrefix, RewriteTo: target.RewriteTo, @@ -749,10 +691,12 @@ func (s *WireGuardService) syncTargets(desiredTargets []Target) error { TLSCert: target.TLSCert, TLSKey: target.TLSKey, }) - logger.Info("Added target %s -> %s during sync", target.SourcePrefix, target.DestPrefix) } } + s.tnet.ReplaceProxySubnetRules(rules) + logger.Info("Synced targets: %d rules installed", len(rules)) + return nil } diff --git a/netstack2/handlers.go b/netstack2/handlers.go index e6ea62e..8426685 100644 --- a/netstack2/handlers.go +++ b/netstack2/handlers.go @@ -166,6 +166,13 @@ func (h *TCPHandler) handleTCPConn(netstackConn *gonet.TCPConn, id stack.Transpo defer netstackConn.Close() + // Release this connection's NAT state once it fully closes, so a rule + // change (e.g. RewriteTo) takes effect for the next connection on this + // tuple instead of being masked by a stale cached resolution forever. + if h.proxyHandler != nil { + defer h.proxyHandler.releaseConnectionState(srcIP, srcPort, dstIP, dstPort, uint8(tcp.ProtocolNumber)) + } + logger.Info("TCP Forwarder: Handling connection %s:%d -> %s:%d", srcIP, srcPort, dstIP, dstPort) // Check if there's a destination rewrite for this connection (e.g., localhost targets) @@ -315,6 +322,13 @@ func (h *UDPHandler) handleUDPConn(netstackConn *gonet.UDPConn, id stack.Transpo dstIP := id.LocalAddress.String() dstPort := id.LocalPort + // Release this session's NAT state once it fully closes (session end or + // idle timeout), so a rule change takes effect for the next session on + // this tuple instead of being masked by a stale cached resolution. + if h.proxyHandler != nil { + defer h.proxyHandler.releaseConnectionState(srcIP, srcPort, dstIP, dstPort, uint8(udp.ProtocolNumber)) + } + logger.Info("UDP Forwarder: Handling connection %s:%d -> %s:%d", srcIP, srcPort, dstIP, dstPort) // Drop connection if blocking is enabled diff --git a/netstack2/proxy.go b/netstack2/proxy.go index 00d6763..be5f23f 100644 --- a/netstack2/proxy.go +++ b/netstack2/proxy.go @@ -272,6 +272,18 @@ func (p *ProxyHandler) RemoveSubnetRule(sourcePrefix, destPrefix netip.Prefix) { p.subnetLookup.RemoveSubnet(sourcePrefix, destPrefix) } +// ReplaceAllSubnetRules atomically replaces the full set of subnet rules. +// Intended for full-state syncs where the desired rule set is authoritative, +// so any stale rule is guaranteed to be cleared even if its key +// (SourcePrefix, DestPrefix) matches a still-desired rule but its contents +// (e.g. RewriteTo) have changed. +func (p *ProxyHandler) ReplaceAllSubnetRules(rules []SubnetRule) { + if p == nil || !p.enabled { + return + } + p.subnetLookup.ReplaceAll(rules) +} + // GetAllRules returns all subnet rules from the proxy handler func (p *ProxyHandler) GetAllRules() []SubnetRule { if p == nil || !p.enabled { @@ -335,6 +347,51 @@ func (p *ProxyHandler) SetHTTPRequestLogSender(fn SendFunc) { p.httpRequestLogger.SetSendFunc(fn) } +// releaseConnectionState removes the per-connection NAT state for a single +// (srcIP, srcPort, dstIP, dstPort, proto) tuple. Callers must only invoke +// this once that exact connection has fully closed (both directions torn +// down), since a new connection can never be accepted on the same 5-tuple +// before then. +// +// This intentionally does NOT touch destRewriteTable/resourceTable: those +// are keyed without srcPort (destKey), so they are shared across every +// concurrent connection from the same source to the same destination +// service. They don't need connection-scoped cleanup - each new connection's +// first packet already refreshes them via HandleIncomingPacket - and +// deleting them here on a single connection's close could break other +// connections still in flight to the same destination. +func (p *ProxyHandler) releaseConnectionState(srcIP string, srcPort uint16, dstIP string, dstPort uint16, proto uint8) { + if p == nil || !p.enabled { + return + } + + key := connKey{ + srcIP: srcIP, + srcPort: srcPort, + dstIP: dstIP, + dstPort: dstPort, + proto: proto, + } + + p.natMu.Lock() + defer p.natMu.Unlock() + + entry, ok := p.natTable[key] + if !ok { + return + } + delete(p.natTable, key) + + reverseKey := reverseConnKey{ + rewrittenTo: entry.rewrittenTo.String(), + originalSrcIP: srcIP, + originalSrcPort: srcPort, + originalDstPort: dstPort, + proto: proto, + } + delete(p.reverseNatTable, reverseKey) +} + // LookupDestinationRewrite looks up the rewritten destination for a connection // This is used by TCP/UDP handlers to find the actual target address func (p *ProxyHandler) LookupDestinationRewrite(srcIP, dstIP string, dstPort uint16, proto uint8) (netip.Addr, bool) { diff --git a/netstack2/subnet_lookup.go b/netstack2/subnet_lookup.go index 757908a..c07162c 100644 --- a/netstack2/subnet_lookup.go +++ b/netstack2/subnet_lookup.go @@ -52,6 +52,13 @@ func (sl *SubnetLookup) AddSubnet(rule SubnetRule) { sl.mu.Lock() defer sl.mu.Unlock() + sl.addSubnetLocked(rule) +} + +// addSubnetLocked is the lock-free body of AddSubnet, factored out so +// ReplaceAll can insert many rules under a single lock acquisition. +// Callers must hold sl.mu for writing. +func (sl *SubnetLookup) addSubnetLocked(rule SubnetRule) { rulePtr := &rule // Canonicalize source prefix to handle host bits correctly @@ -89,6 +96,21 @@ func (sl *SubnetLookup) AddSubnet(rule SubnetRule) { destTriePtr.rules = newRules } +// ReplaceAll atomically replaces the entire rule set with the given rules. +// This guarantees no stale rule can survive a sync, even when a rule's key +// (SourcePrefix, DestPrefix) is unchanged but other fields (e.g. RewriteTo) +// differ - a case that an add/remove diff keyed only on prefixes would miss. +func (sl *SubnetLookup) ReplaceAll(rules []SubnetRule) { + sl.mu.Lock() + defer sl.mu.Unlock() + + sl.sourceTrie = &bart.Table[*destTrie]{} + + for _, rule := range rules { + sl.addSubnetLocked(rule) + } +} + // RemoveSubnet removes a subnet rule from the lookup table func (sl *SubnetLookup) RemoveSubnet(sourcePrefix, destPrefix netip.Prefix) { sl.mu.Lock() diff --git a/netstack2/tun.go b/netstack2/tun.go index d104f2e..49f74d4 100644 --- a/netstack2/tun.go +++ b/netstack2/tun.go @@ -364,6 +364,15 @@ func (net *Net) RemoveProxySubnetRule(sourcePrefix, destPrefix netip.Prefix) { } } +// ReplaceProxySubnetRules atomically replaces the full set of subnet rules +// on the proxy handler with the given rules. +func (net *Net) ReplaceProxySubnetRules(rules []SubnetRule) { + tun := (*netTun)(net) + if tun.proxyHandler != nil { + tun.proxyHandler.ReplaceAllSubnetRules(rules) + } +} + // GetProxySubnetRules returns all subnet rules from the proxy handler func (net *Net) GetProxySubnetRules() []SubnetRule { tun := (*netTun)(net) diff --git a/newt/connect.go b/newt/connect.go index 050725a..1413abd 100644 --- a/newt/connect.go +++ b/newt/connect.go @@ -261,42 +261,91 @@ persistent_keepalive_interval=5`, util.FixKey(n.privateKey.String()), util.FixKe } if len(n.wgData.BrowserGatewayTargets) > 0 { + // The netstack is fresh on (re)connect, so any previously running + // gateway listener is bound to a now-defunct interface - tear it down. if n.browserGatewayStop != nil { n.browserGatewayStop() n.browserGatewayStop = nil + n.browserGateway = nil } - bgTargets := make([]browsergateway.Target, 0, len(n.wgData.BrowserGatewayTargets)) - for _, t := range n.wgData.BrowserGatewayTargets { - bgTargets = append(bgTargets, browsergateway.Target{ - ID: t.ID, - Type: t.Type, - Destination: t.Destination, - DestinationPort: t.DestinationPort, - AuthToken: t.AuthToken, - }) - } - - n.browserGateway = browsergateway.New(browsergateway.Config{SSHCredentials: n.sshCredStore}) - n.browserGateway.SetTargets(bgTargets) - - var ln net.Listener - var bgErr error - if n.config.UseNativeMainInterface { - ln, bgErr = net.Listen("tcp", fmt.Sprintf("%s:%d", n.wgData.TunnelIP, browsergateway.ListenPort)) + if err := n.startBrowserGateway(); err != nil { + logger.Error("Failed to start browser gateway listener: %v", err) } else { - ln, bgErr = n.tnet.ListenTCP(&net.TCPAddr{Port: browsergateway.ListenPort}) - } - if bgErr != nil { - logger.Error("Failed to start browser gateway listener: %v", bgErr) - } else { - n.browserGatewayStop = func() { _ = ln.Close() } - go func() { - logger.Debug("Browser gateway started on port %d", browsergateway.ListenPort) - if startErr := n.browserGateway.Start(ln); startErr != nil { - logger.Error("Browser gateway stopped with error: %v", startErr) - } - }() + n.browserGateway.SetTargets(toBrowserGatewayTargets(n.wgData.BrowserGatewayTargets)) } } } + +// startBrowserGateway creates the browser gateway and its listener if one +// isn't already running. Callers that need to rebind to a fresh netstack +// (e.g. on reconnect) must stop and clear any existing gateway first. +func (n *Newt) startBrowserGateway() error { + if n.browserGateway != nil { + return nil + } + if n.tnet == nil && !n.config.UseNativeMainInterface { + return fmt.Errorf("netstack not ready") + } + + gateway := browsergateway.New(browsergateway.Config{SSHCredentials: n.sshCredStore}) + + var ln net.Listener + var err error + if n.config.UseNativeMainInterface { + ln, err = net.Listen("tcp", fmt.Sprintf("%s:%d", n.wgData.TunnelIP, browsergateway.ListenPort)) + } else { + ln, err = n.tnet.ListenTCP(&net.TCPAddr{Port: browsergateway.ListenPort}) + } + if err != nil { + return err + } + + n.browserGateway = gateway + n.browserGatewayStop = func() { _ = ln.Close() } + go func() { + logger.Debug("Browser gateway started on port %d", browsergateway.ListenPort) + if startErr := gateway.Start(ln); startErr != nil { + logger.Error("Browser gateway stopped with error: %v", startErr) + } + }() + + return nil +} + +// syncBrowserGatewayTargets reconciles the browser gateway's allowed +// destinations with the desired state received from a sync message. +// It lazily starts the gateway if targets are present and it isn't running +// yet, and clears the allow-list (without tearing down the listener) when +// no targets are desired. +func (n *Newt) syncBrowserGatewayTargets(targets []BrowserGatewayTarget) { + bgTargets := toBrowserGatewayTargets(targets) + + if len(bgTargets) == 0 { + if n.browserGateway != nil { + n.browserGateway.SetTargets(nil) + } + return + } + + if err := n.startBrowserGateway(); err != nil { + logger.Error("Failed to start browser gateway: %v", err) + return + } + + n.browserGateway.SetTargets(bgTargets) +} + +func toBrowserGatewayTargets(targets []BrowserGatewayTarget) []browsergateway.Target { + bgTargets := make([]browsergateway.Target, 0, len(targets)) + for _, t := range targets { + bgTargets = append(bgTargets, browsergateway.Target{ + ID: t.ID, + Type: t.Type, + Destination: t.Destination, + DestinationPort: t.DestinationPort, + AuthToken: t.AuthToken, + }) + } + return bgTargets +} diff --git a/newt/data.go b/newt/data.go index 311f7da..719b707 100644 --- a/newt/data.go +++ b/newt/data.go @@ -142,6 +142,14 @@ func (n *Newt) handleSync(msg websocket.WSMessage) { n.updateRemoteExitNodeSubnets(syncData.RemoteExitNodeSubnets) } + // Sync clients WireGuard peers and targets, if clients are set up + if n.wgService != nil { + n.wgService.Sync(syncData.Peers, syncData.ClientTargets) + } + + // Sync browser gateway targets + n.syncBrowserGatewayTargets(syncData.BrowserGatewayTargets) + // Sync health check targets if err := n.healthMonitor.SyncTargets(syncData.HealthCheckTargets); err != nil { logger.Error("Failed to sync health check targets: %v", err) diff --git a/newt/handlers.go b/newt/handlers.go index b679a2f..0a448f1 100644 --- a/newt/handlers.go +++ b/newt/handlers.go @@ -276,6 +276,8 @@ func (n *Newt) registerHandlers(ctx context.Context) { logger.Debug("Sent exit node ping results to cloud for selection: pingResults=%+v", pingResults) }) + n.client.RegisterHandler("newt/sync", n.handleSync) + n.client.RegisterHandler("newt/tcp/add", func(msg websocket.WSMessage) { logger.Debug(fmtReceivedMsg, msg) @@ -458,8 +460,6 @@ func (n *Newt) registerHandlers(ctx context.Context) { logger.Info("Removed %d remote exit node subnets", len(data.Subnets)) }) - n.client.RegisterHandler("newt/sync", n.handleSync) - n.client.RegisterHandler("newt/socket/check", func(msg websocket.WSMessage) { logger.Debug("Received Docker socket check request") diff --git a/newt/types.go b/newt/types.go index d2ff65c..b534713 100644 --- a/newt/types.go +++ b/newt/types.go @@ -1,6 +1,9 @@ package newt -import "github.com/fosrl/newt/healthcheck" +import ( + wgclients "github.com/fosrl/newt/clients" + "github.com/fosrl/newt/healthcheck" +) type BrowserGatewayTarget struct { ID int `json:"id"` @@ -62,7 +65,10 @@ type BlueprintResult struct { // Define the sync data structure type SyncData struct { - Targets TargetsByType `json:"targets"` - HealthCheckTargets []healthcheck.Config `json:"healthCheckTargets"` - RemoteExitNodeSubnets []string `json:"remoteExitNodeSubnets"` + Targets TargetsByType `json:"proxyTargets"` + HealthCheckTargets []healthcheck.Config `json:"healthCheckTargets"` + RemoteExitNodeSubnets []string `json:"remoteExitNodeSubnets"` + Peers []wgclients.Peer `json:"peers"` + ClientTargets []wgclients.Target `json:"clientTargets"` + BrowserGatewayTargets []BrowserGatewayTarget `json:"browserGatewayTargets"` }