mirror of
https://github.com/fosrl/newt.git
synced 2026-08-04 02:25:13 -05:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
46f58dd59f | ||
|
|
01cb6f39ca | ||
|
|
69d3925167 | ||
|
|
385dafa857 | ||
|
|
f9d57acd3e | ||
|
|
6610655376 | ||
|
|
8d582b4ea5 | ||
|
|
2bad244186 | ||
|
|
10adb416f8 | ||
|
|
dde44d6666 | ||
|
|
e21d608bd8 | ||
|
|
6d8aca9c7c |
@@ -683,6 +683,23 @@ func (b *SharedBind) receiveIPv4Simple(conn *net.UDPConn, bufs [][]byte, sizes [
|
||||
}
|
||||
}
|
||||
|
||||
// IsMagicPacket reports whether payload is one of our connectivity-test magic
|
||||
// packets (a MagicTestRequest or MagicTestResponse). These packets are meant to
|
||||
// travel directly between physical UDP sockets and must never be encapsulated by
|
||||
// WireGuard - e.g. if OS routing mistakenly sends one into a WireGuard TUN
|
||||
// interface (because the destination falls inside a routed tunnel subnet), it
|
||||
// should be dropped there rather than tunneled, which would otherwise make a
|
||||
// LAN-local endpoint test falsely appear to succeed over the tunnel.
|
||||
func IsMagicPacket(payload []byte) bool {
|
||||
if len(payload) >= MagicTestRequestLen && bytes.HasPrefix(payload, MagicTestRequest) {
|
||||
return true
|
||||
}
|
||||
if len(payload) >= MagicTestResponseLen && bytes.HasPrefix(payload, MagicTestResponse) {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// handleMagicPacket checks if the packet is a magic test packet and responds if so.
|
||||
// Returns true if the packet was a magic packet and was handled (should not be passed to WireGuard).
|
||||
func (b *SharedBind) handleMagicPacket(data []byte, addr *net.UDPAddr) bool {
|
||||
|
||||
+147
-15
@@ -35,10 +35,11 @@ import (
|
||||
)
|
||||
|
||||
type WgConfig struct {
|
||||
IpAddress string `json:"ipAddress"`
|
||||
Peers []Peer `json:"peers"`
|
||||
Targets []Target `json:"targets"`
|
||||
ChainId string `json:"chainId"`
|
||||
IpAddress string `json:"ipAddress"`
|
||||
Peers []Peer `json:"peers"`
|
||||
Targets []Target `json:"targets"`
|
||||
Certs []CertData `json:"certs"`
|
||||
ChainId string `json:"chainId"`
|
||||
}
|
||||
|
||||
type Target struct {
|
||||
@@ -53,6 +54,23 @@ type Target struct {
|
||||
HTTPTargets []netstack2.HTTPTarget `json:"httpTargets,omitempty"` // for http protocol, list of downstream services to load balance across
|
||||
TLSCert string `json:"tlsCert,omitempty"` // PEM-encoded certificate for incoming HTTPS termination
|
||||
TLSKey string `json:"tlsKey,omitempty"` // PEM-encoded private key for incoming HTTPS termination
|
||||
TLSCertID string `json:"tlsCertId,omitempty"` // references an entry in the sync message's Certs list instead of inlining TLSCert/TLSKey
|
||||
}
|
||||
|
||||
// CertData is a single shared TLS certificate/key pair, referenced by ID from
|
||||
// one or more Targets via TLSCertID. Sent once per sync message so that many
|
||||
// targets backed by the same certificate (e.g. a wildcard cert) don't each
|
||||
// carry a full copy of the PEM data.
|
||||
type CertData struct {
|
||||
ID string `json:"id"`
|
||||
Cert string `json:"cert"`
|
||||
Key string `json:"key"`
|
||||
}
|
||||
|
||||
// CertPair holds the resolved PEM certificate/key material for a CertData entry.
|
||||
type CertPair struct {
|
||||
Cert string
|
||||
Key string
|
||||
}
|
||||
|
||||
type PortRange struct {
|
||||
@@ -122,6 +140,11 @@ type WireGuardService struct {
|
||||
|
||||
// connection blocking: when true, all new incoming connections are dropped
|
||||
blocked atomic.Bool
|
||||
|
||||
// certs resolves TLSCertID references on incoming Targets to their PEM
|
||||
// cert/key material. Replaced wholesale on every full sync.
|
||||
certs map[string]CertPair
|
||||
certsMu sync.RWMutex
|
||||
}
|
||||
|
||||
// generateChainId generates a random chain ID for deduplicating round-trip messages.
|
||||
@@ -196,6 +219,8 @@ func NewWireGuardService(interfaceName string, port uint16, mtu int, host string
|
||||
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/certs/add", service.handleAddCerts)
|
||||
wsClient.RegisterHandler("newt/certs/remove", service.handleRemoveCerts)
|
||||
|
||||
return service, nil
|
||||
}
|
||||
@@ -504,9 +529,10 @@ func (s *WireGuardService) LoadRemoteConfig() error {
|
||||
chainId := generateChainId()
|
||||
s.pendingConfigChainId = chainId
|
||||
s.stopGetConfig = s.client.SendMessageInterval("newt/wg/get-config", map[string]interface{}{
|
||||
"publicKey": s.key.PublicKey().String(),
|
||||
"port": s.Port,
|
||||
"chainId": chainId,
|
||||
"publicKey": s.key.PublicKey().String(),
|
||||
"port": s.Port,
|
||||
"chainId": chainId,
|
||||
"localEndpoints": network.GetLocalEndpoints(s.Port, s.interfaceName),
|
||||
}, 2*time.Second)
|
||||
|
||||
logger.Debug("Requesting WireGuard configuration from remote server")
|
||||
@@ -543,6 +569,7 @@ func (s *WireGuardService) handleConfig(msg websocket.WSMessage) {
|
||||
}
|
||||
|
||||
s.config = config
|
||||
s.SetCerts(config.Certs)
|
||||
|
||||
if s.stopGetConfig != nil {
|
||||
s.stopGetConfig()
|
||||
@@ -567,6 +594,107 @@ func (s *WireGuardService) handleConfig(msg websocket.WSMessage) {
|
||||
logger.Info("Client connectivity setup. Ready to accept connections from clients!")
|
||||
}
|
||||
|
||||
// SetCerts replaces the TLSCertID lookup table used by resolveTLS. The server
|
||||
// sends the complete set of referenced certs on every full sync, so this is a
|
||||
// wholesale replacement rather than an incremental merge.
|
||||
func (s *WireGuardService) SetCerts(certs []CertData) {
|
||||
m := make(map[string]CertPair, len(certs))
|
||||
for _, c := range certs {
|
||||
m[c.ID] = CertPair{Cert: c.Cert, Key: c.Key}
|
||||
}
|
||||
s.certsMu.Lock()
|
||||
s.certs = m
|
||||
s.certsMu.Unlock()
|
||||
}
|
||||
|
||||
// resolveTLS returns the PEM cert/key to use for target's incoming HTTPS
|
||||
// termination: target.TLSCertID looked up in the certs table if set, falling
|
||||
// back to the target's own inline TLSCert/TLSKey otherwise.
|
||||
func (s *WireGuardService) resolveTLS(target Target) (cert, key string) {
|
||||
if target.TLSCertID == "" {
|
||||
return target.TLSCert, target.TLSKey
|
||||
}
|
||||
s.certsMu.RLock()
|
||||
pair, ok := s.certs[target.TLSCertID]
|
||||
s.certsMu.RUnlock()
|
||||
if !ok {
|
||||
logger.Warn("No cert found for tlsCertId %s, falling back to inline cert on target", target.TLSCertID)
|
||||
return target.TLSCert, target.TLSKey
|
||||
}
|
||||
return pair.Cert, pair.Key
|
||||
}
|
||||
|
||||
// AddCerts upserts the given certs into the lookup table used by resolveTLS,
|
||||
// without discarding any certs already present. Used for incremental cert
|
||||
// pushes (e.g. after a renewal) outside of a full newt/sync or
|
||||
// newt/wg/receive-config, which replace the table wholesale via SetCerts.
|
||||
func (s *WireGuardService) AddCerts(certs []CertData) {
|
||||
if len(certs) == 0 {
|
||||
return
|
||||
}
|
||||
s.certsMu.Lock()
|
||||
if s.certs == nil {
|
||||
s.certs = make(map[string]CertPair, len(certs))
|
||||
}
|
||||
for _, c := range certs {
|
||||
s.certs[c.ID] = CertPair{Cert: c.Cert, Key: c.Key}
|
||||
}
|
||||
s.certsMu.Unlock()
|
||||
}
|
||||
|
||||
// RemoveCerts deletes the given cert IDs from the lookup table, e.g. once the
|
||||
// server knows no target references them anymore.
|
||||
func (s *WireGuardService) RemoveCerts(ids []string) {
|
||||
if len(ids) == 0 {
|
||||
return
|
||||
}
|
||||
s.certsMu.Lock()
|
||||
for _, id := range ids {
|
||||
delete(s.certs, id)
|
||||
}
|
||||
s.certsMu.Unlock()
|
||||
}
|
||||
|
||||
// handleAddCerts processes a "newt/certs/add" message: an array of CertData
|
||||
// to upsert into the cert lookup table.
|
||||
func (s *WireGuardService) handleAddCerts(msg websocket.WSMessage) {
|
||||
jsonData, err := json.Marshal(msg.Data)
|
||||
if err != nil {
|
||||
logger.Info("Error marshaling cert add data: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
var certs []CertData
|
||||
if err := json.Unmarshal(jsonData, &certs); err != nil {
|
||||
logger.Warn("Error unmarshaling cert add data: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
s.AddCerts(certs)
|
||||
logger.Info("Added %d certs", len(certs))
|
||||
}
|
||||
|
||||
// handleRemoveCerts processes a "newt/certs/remove" message: {ids: [...]}
|
||||
// naming the cert IDs to drop from the lookup table.
|
||||
func (s *WireGuardService) handleRemoveCerts(msg websocket.WSMessage) {
|
||||
jsonData, err := json.Marshal(msg.Data)
|
||||
if err != nil {
|
||||
logger.Info("Error marshaling cert remove data: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
var data struct {
|
||||
IDs []string `json:"ids"`
|
||||
}
|
||||
if err := json.Unmarshal(jsonData, &data); err != nil {
|
||||
logger.Warn("Error unmarshaling cert remove data: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
s.RemoveCerts(data.IDs)
|
||||
logger.Info("Removed %d certs", len(data.IDs))
|
||||
}
|
||||
|
||||
// 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) {
|
||||
@@ -679,6 +807,7 @@ func (s *WireGuardService) syncTargets(desiredTargets []Target) error {
|
||||
continue
|
||||
}
|
||||
|
||||
tlsCert, tlsKey := s.resolveTLS(target)
|
||||
rules = append(rules, netstack2.SubnetRule{
|
||||
SourcePrefix: sourcePrefix,
|
||||
DestPrefix: destPrefix,
|
||||
@@ -688,8 +817,8 @@ func (s *WireGuardService) syncTargets(desiredTargets []Target) error {
|
||||
ResourceId: target.ResourceId,
|
||||
Protocol: target.Protocol,
|
||||
HTTPTargets: target.HTTPTargets,
|
||||
TLSCert: target.TLSCert,
|
||||
TLSKey: target.TLSKey,
|
||||
TLSCert: tlsCert,
|
||||
TLSKey: tlsKey,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -975,6 +1104,7 @@ func (s *WireGuardService) ensureTargets(targets []Target) error {
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid CIDR %s: %v", sp, err)
|
||||
}
|
||||
tlsCert, tlsKey := s.resolveTLS(target)
|
||||
s.tnet.AddProxySubnetRule(netstack2.SubnetRule{
|
||||
SourcePrefix: sourcePrefix,
|
||||
DestPrefix: destPrefix,
|
||||
@@ -984,8 +1114,8 @@ func (s *WireGuardService) ensureTargets(targets []Target) error {
|
||||
ResourceId: target.ResourceId,
|
||||
Protocol: target.Protocol,
|
||||
HTTPTargets: target.HTTPTargets,
|
||||
TLSCert: target.TLSCert,
|
||||
TLSKey: target.TLSKey,
|
||||
TLSCert: tlsCert,
|
||||
TLSKey: tlsKey,
|
||||
})
|
||||
logger.Info("Added target subnet from %s to %s rewrite to %s with port ranges: %v", sp, target.DestPrefix, target.RewriteTo, target.PortRange)
|
||||
}
|
||||
@@ -1379,6 +1509,7 @@ func (s *WireGuardService) handleAddTarget(msg websocket.WSMessage) {
|
||||
logger.Info("Invalid CIDR %s: %v", sp, err)
|
||||
continue
|
||||
}
|
||||
tlsCert, tlsKey := s.resolveTLS(target)
|
||||
s.tnet.AddProxySubnetRule(netstack2.SubnetRule{
|
||||
SourcePrefix: sourcePrefix,
|
||||
DestPrefix: destPrefix,
|
||||
@@ -1388,8 +1519,8 @@ func (s *WireGuardService) handleAddTarget(msg websocket.WSMessage) {
|
||||
ResourceId: target.ResourceId,
|
||||
Protocol: target.Protocol,
|
||||
HTTPTargets: target.HTTPTargets,
|
||||
TLSCert: target.TLSCert,
|
||||
TLSKey: target.TLSKey,
|
||||
TLSCert: tlsCert,
|
||||
TLSKey: tlsKey,
|
||||
})
|
||||
logger.Info("Added target subnet from %s to %s rewrite to %s with port ranges: %v", sp, target.DestPrefix, target.RewriteTo, target.PortRange)
|
||||
}
|
||||
@@ -1508,6 +1639,7 @@ func (s *WireGuardService) handleUpdateTarget(msg websocket.WSMessage) {
|
||||
logger.Info("Invalid CIDR %s: %v", sp, err)
|
||||
continue
|
||||
}
|
||||
tlsCert, tlsKey := s.resolveTLS(target)
|
||||
s.tnet.AddProxySubnetRule(netstack2.SubnetRule{
|
||||
SourcePrefix: sourcePrefix,
|
||||
DestPrefix: destPrefix,
|
||||
@@ -1517,8 +1649,8 @@ func (s *WireGuardService) handleUpdateTarget(msg websocket.WSMessage) {
|
||||
ResourceId: target.ResourceId,
|
||||
Protocol: target.Protocol,
|
||||
HTTPTargets: target.HTTPTargets,
|
||||
TLSCert: target.TLSCert,
|
||||
TLSKey: target.TLSKey,
|
||||
TLSCert: tlsCert,
|
||||
TLSKey: tlsKey,
|
||||
})
|
||||
logger.Info("Added target subnet from %s to %s rewrite to %s with port ranges: %v", sp, target.DestPrefix, target.RewriteTo, target.PortRange)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,176 @@
|
||||
// Package exitnode implements the exit-node ping dance run before
|
||||
// registering with the server: request the candidate exit nodes, ping each
|
||||
// one over HTTP, and report the results so the server can pick the best one.
|
||||
// It is shared between newt and olm, which both register the same way.
|
||||
package exitnode
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/fosrl/newt/logger"
|
||||
)
|
||||
|
||||
// ExitNodeData is the payload the server sends in response to a
|
||||
// "*/ping/request" message.
|
||||
type ExitNodeData struct {
|
||||
ExitNodes []ExitNode `json:"exitNodes"`
|
||||
ChainId string `json:"chainId"`
|
||||
}
|
||||
|
||||
// ExitNode is a candidate exit node offered by the server for ping selection.
|
||||
type ExitNode struct {
|
||||
ID int `json:"exitNodeId"`
|
||||
Name string `json:"exitNodeName"`
|
||||
Endpoint string `json:"endpoint"`
|
||||
Weight float64 `json:"weight"`
|
||||
WasPreviouslyConnected bool `json:"wasPreviouslyConnected"`
|
||||
}
|
||||
|
||||
// ExitNodePingResult is the measured latency (or error) for one exit node,
|
||||
// sent back to the server in the "*/wg/register" message's pingResults field.
|
||||
type ExitNodePingResult struct {
|
||||
ExitNodeID int `json:"exitNodeId"`
|
||||
LatencyMs int64 `json:"latencyMs"`
|
||||
Weight float64 `json:"weight"`
|
||||
Error string `json:"error,omitempty"`
|
||||
Name string `json:"exitNodeName"`
|
||||
Endpoint string `json:"endpoint"`
|
||||
WasPreviouslyConnected bool `json:"wasPreviouslyConnected"`
|
||||
}
|
||||
|
||||
// PingExitNodes pings the given exit nodes over HTTP and returns a per-node
|
||||
// ExitNodePingResult suitable for inclusion in a wg/register message's
|
||||
// pingResults field, so the server can select the best exit node.
|
||||
//
|
||||
// If there's only one exit node, or preferEndpoint names one of them, the
|
||||
// matching node is returned immediately with LatencyMs 0 and no pinging is
|
||||
// done. Otherwise every node is pinged pingAttempts times over HTTP GET
|
||||
// <endpoint>/ping and the average latency of successful attempts is used.
|
||||
//
|
||||
// When alreadyConnected is true, a node flagged WasPreviouslyConnected is
|
||||
// excluded from the results as long as at least one other healthy node is
|
||||
// available, biasing reconnects toward switching away from a possibly
|
||||
// degraded node.
|
||||
func PingExitNodes(exitNodes []ExitNode, preferEndpoint string, alreadyConnected bool) []ExitNodePingResult {
|
||||
if len(exitNodes) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
if len(exitNodes) == 1 || preferEndpoint != "" {
|
||||
selected := exitNodes[0]
|
||||
if preferEndpoint != "" {
|
||||
for _, node := range exitNodes {
|
||||
if node.Endpoint == preferEndpoint {
|
||||
selected = node
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
logger.Debug("Only one exit node available, using it directly: %s", selected.Endpoint)
|
||||
|
||||
return []ExitNodePingResult{
|
||||
{
|
||||
ExitNodeID: selected.ID,
|
||||
LatencyMs: 0,
|
||||
Weight: selected.Weight,
|
||||
Error: "",
|
||||
Name: selected.Name,
|
||||
Endpoint: selected.Endpoint,
|
||||
WasPreviouslyConnected: selected.WasPreviouslyConnected,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
type nodeResult struct {
|
||||
Node ExitNode
|
||||
Latency time.Duration
|
||||
Err error
|
||||
}
|
||||
|
||||
results := make([]nodeResult, len(exitNodes))
|
||||
const pingAttempts = 3
|
||||
for i, node := range exitNodes {
|
||||
var totalLatency time.Duration
|
||||
var lastErr error
|
||||
successes := 0
|
||||
httpClient := &http.Client{
|
||||
Timeout: 5 * time.Second,
|
||||
}
|
||||
url := node.Endpoint
|
||||
if !strings.HasPrefix(url, "http://") && !strings.HasPrefix(url, "https://") {
|
||||
url = "http://" + url
|
||||
}
|
||||
if !strings.HasSuffix(url, "/ping") {
|
||||
url = strings.TrimRight(url, "/") + "/ping"
|
||||
}
|
||||
for j := 0; j < pingAttempts; j++ {
|
||||
start := time.Now()
|
||||
resp, err := httpClient.Get(url)
|
||||
latency := time.Since(start)
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
logger.Warn("Failed to ping exit node %d (%s) attempt %d: %v", node.ID, url, j+1, err)
|
||||
continue
|
||||
}
|
||||
resp.Body.Close()
|
||||
totalLatency += latency
|
||||
successes++
|
||||
}
|
||||
var avgLatency time.Duration
|
||||
if successes > 0 {
|
||||
avgLatency = totalLatency / time.Duration(successes)
|
||||
}
|
||||
if successes == 0 {
|
||||
results[i] = nodeResult{Node: node, Latency: 0, Err: lastErr}
|
||||
} else {
|
||||
results[i] = nodeResult{Node: node, Latency: avgLatency, Err: nil}
|
||||
}
|
||||
}
|
||||
|
||||
var pingResults []ExitNodePingResult
|
||||
for _, res := range results {
|
||||
errMsg := ""
|
||||
if res.Err != nil {
|
||||
errMsg = res.Err.Error()
|
||||
}
|
||||
pingResults = append(pingResults, ExitNodePingResult{
|
||||
ExitNodeID: res.Node.ID,
|
||||
LatencyMs: res.Latency.Milliseconds(),
|
||||
Weight: res.Node.Weight,
|
||||
Error: errMsg,
|
||||
Name: res.Node.Name,
|
||||
Endpoint: res.Node.Endpoint,
|
||||
WasPreviouslyConnected: res.Node.WasPreviouslyConnected,
|
||||
})
|
||||
}
|
||||
|
||||
if alreadyConnected {
|
||||
var filteredPingResults []ExitNodePingResult
|
||||
previouslyConnectedNodeIdx := -1
|
||||
for i, res := range pingResults {
|
||||
if res.WasPreviouslyConnected {
|
||||
previouslyConnectedNodeIdx = i
|
||||
}
|
||||
}
|
||||
goodNodeCount := 0
|
||||
for i, res := range pingResults {
|
||||
if i != previouslyConnectedNodeIdx && res.LatencyMs > 0 && res.Error == "" {
|
||||
goodNodeCount++
|
||||
}
|
||||
}
|
||||
if previouslyConnectedNodeIdx != -1 && goodNodeCount > 0 {
|
||||
for i, res := range pingResults {
|
||||
if i != previouslyConnectedNodeIdx {
|
||||
filteredPingResults = append(filteredPingResults, res)
|
||||
}
|
||||
}
|
||||
pingResults = filteredPingResults
|
||||
logger.Info("Excluding previously connected exit node from ping results due to other available nodes")
|
||||
}
|
||||
}
|
||||
|
||||
return pingResults
|
||||
}
|
||||
@@ -3,9 +3,15 @@ package netstack2
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"crypto/ecdsa"
|
||||
"crypto/elliptic"
|
||||
"crypto/rand"
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"crypto/x509/pkix"
|
||||
"errors"
|
||||
"fmt"
|
||||
"math/big"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httputil"
|
||||
@@ -65,6 +71,13 @@ type HTTPHandler struct {
|
||||
// of the PEM certificate and key. Parsing a keypair is relatively expensive
|
||||
// and the same cert is likely reused across many connections.
|
||||
tlsCache sync.Map // map[string]*tls.Config
|
||||
|
||||
// fallbackTLSOnce/fallbackTLSCfg hold a lazily-generated self-signed
|
||||
// certificate used when a rule's configured cert/key fails to parse, so
|
||||
// that a misconfigured rule degrades to a browser cert warning instead of
|
||||
// silently dropping every connection.
|
||||
fallbackTLSOnce sync.Once
|
||||
fallbackTLSCfg *tls.Config
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
@@ -262,7 +275,13 @@ func (h *HTTPHandler) getTLSConfig(rule *SubnetRule) (*tls.Config, error) {
|
||||
|
||||
cert, err := tls.X509KeyPair([]byte(rule.TLSCert), []byte(rule.TLSKey))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to parse TLS keypair: %w", err)
|
||||
// A misconfigured rule (bad/missing PEM data) must not take the whole
|
||||
// connection down: fall back to a self-signed cert so the handshake
|
||||
// still completes and the request reaches handleRequest, which routes
|
||||
// independently of the cert. Clients will see a cert warning instead
|
||||
// of a silent connection reset.
|
||||
logger.Warn("HTTP handler: falling back to self-signed cert for rule (invalid configured keypair): %v", err)
|
||||
return h.getFallbackTLSConfig(), nil
|
||||
}
|
||||
cfg := &tls.Config{
|
||||
Certificates: []tls.Certificate{cert},
|
||||
@@ -273,6 +292,57 @@ func (h *HTTPHandler) getTLSConfig(rule *SubnetRule) (*tls.Config, error) {
|
||||
return actual.(*tls.Config), nil
|
||||
}
|
||||
|
||||
// getFallbackTLSConfig returns a *tls.Config backed by a self-signed
|
||||
// certificate, generated once and reused for the lifetime of the handler.
|
||||
func (h *HTTPHandler) getFallbackTLSConfig() *tls.Config {
|
||||
h.fallbackTLSOnce.Do(func() {
|
||||
cert, err := generateSelfSignedCert()
|
||||
if err != nil {
|
||||
// Generation of an in-memory self-signed cert has no external
|
||||
// dependencies and should never fail; if it somehow does, there
|
||||
// is no sensible fallback left, so surface it loudly.
|
||||
logger.Error("HTTP handler: failed to generate fallback self-signed cert: %v", err)
|
||||
return
|
||||
}
|
||||
h.fallbackTLSCfg = &tls.Config{Certificates: []tls.Certificate{cert}}
|
||||
})
|
||||
return h.fallbackTLSCfg
|
||||
}
|
||||
|
||||
// generateSelfSignedCert creates a fresh, in-memory self-signed TLS
|
||||
// certificate/key pair valid for one year, used as a fallback when a rule's
|
||||
// configured certificate cannot be parsed.
|
||||
func generateSelfSignedCert() (tls.Certificate, error) {
|
||||
priv, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
||||
if err != nil {
|
||||
return tls.Certificate{}, fmt.Errorf("failed to generate private key: %w", err)
|
||||
}
|
||||
|
||||
serialNumber, err := rand.Int(rand.Reader, new(big.Int).Lsh(big.NewInt(1), 128))
|
||||
if err != nil {
|
||||
return tls.Certificate{}, fmt.Errorf("failed to generate serial number: %w", err)
|
||||
}
|
||||
|
||||
template := x509.Certificate{
|
||||
SerialNumber: serialNumber,
|
||||
Subject: pkix.Name{CommonName: "newt-fallback"},
|
||||
NotBefore: time.Now().Add(-time.Hour),
|
||||
NotAfter: time.Now().AddDate(1, 0, 0),
|
||||
KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment,
|
||||
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
|
||||
}
|
||||
|
||||
derBytes, err := x509.CreateCertificate(rand.Reader, &template, &template, &priv.PublicKey, priv)
|
||||
if err != nil {
|
||||
return tls.Certificate{}, fmt.Errorf("failed to create certificate: %w", err)
|
||||
}
|
||||
|
||||
return tls.Certificate{
|
||||
Certificate: [][]byte{derBytes},
|
||||
PrivateKey: priv,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// getProxy returns a cached *httputil.ReverseProxy for the given target,
|
||||
// creating one on first use. Reusing the proxy preserves its http.Transport
|
||||
// connection pool, avoiding repeated TCP/TLS handshakes to the downstream.
|
||||
|
||||
@@ -167,3 +167,95 @@ func configureLinux(interfaceName string, ip net.IP, ipNet *net.IPNet) error {
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// AddSecondaryAddress adds an additional IP address (given as CIDR, e.g. "10.10.0.5/32")
|
||||
// to an already-configured interface. It also records the address in the shared
|
||||
// NetworkSettings (see AddIPv4Address) so mobile (iOS/Android) packet-tunnel providers
|
||||
// pick it up on their next settings poll - those platforms have no OS-level interface
|
||||
// to configure directly, so this is the only way they learn about the address.
|
||||
func AddSecondaryAddress(interfaceName string, addr string) error {
|
||||
ip, ipNet, err := net.ParseCIDR(addr)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid IP address: %v", err)
|
||||
}
|
||||
|
||||
mask := net.IP(ipNet.Mask).String()
|
||||
AddIPv4Address(ip.String(), mask)
|
||||
|
||||
if interfaceName == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
switch runtime.GOOS {
|
||||
case "linux":
|
||||
return configureLinux(interfaceName, ip, ipNet)
|
||||
case "darwin":
|
||||
return configureDarwin(interfaceName, ip, ipNet)
|
||||
case "windows":
|
||||
return configureWindows(interfaceName, ip, ipNet)
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
// RemoveSecondaryAddress removes an IP address (given as CIDR) previously added with
|
||||
// AddSecondaryAddress, including from the shared NetworkSettings used by mobile
|
||||
// packet-tunnel providers.
|
||||
func RemoveSecondaryAddress(interfaceName string, addr string) error {
|
||||
ip, ipNet, err := net.ParseCIDR(addr)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid IP address: %v", err)
|
||||
}
|
||||
|
||||
RemoveIPv4Address(ip.String())
|
||||
|
||||
if interfaceName == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
switch runtime.GOOS {
|
||||
case "linux":
|
||||
return removeLinuxAddress(interfaceName, ip, ipNet)
|
||||
case "darwin":
|
||||
return removeDarwinAddress(interfaceName, ip, ipNet)
|
||||
case "windows":
|
||||
return removeWindowsAddress(interfaceName, ip, ipNet)
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func removeLinuxAddress(interfaceName string, ip net.IP, ipNet *net.IPNet) error {
|
||||
link, err := netlink.LinkByName(interfaceName)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get interface %s: %v", interfaceName, err)
|
||||
}
|
||||
|
||||
addr := &netlink.Addr{
|
||||
IPNet: &net.IPNet{
|
||||
IP: ip,
|
||||
Mask: ipNet.Mask,
|
||||
},
|
||||
}
|
||||
|
||||
if err := netlink.AddrDel(link, addr); err != nil {
|
||||
return fmt.Errorf("failed to remove IP address: %v", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func removeDarwinAddress(interfaceName string, ip net.IP, ipNet *net.IPNet) error {
|
||||
prefix, _ := ipNet.Mask.Size()
|
||||
ipStr := fmt.Sprintf("%s/%d", ip.String(), prefix)
|
||||
|
||||
cmd := exec.Command("/sbin/ifconfig", interfaceName, "inet", ipStr, "-alias")
|
||||
logger.Info("Running command: %v", cmd)
|
||||
|
||||
out, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
return fmt.Errorf("ifconfig command failed: %v, output: %s", err, out)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -10,3 +10,7 @@ import (
|
||||
func configureWindows(interfaceName string, ip net.IP, ipNet *net.IPNet) error {
|
||||
return fmt.Errorf("configureWindows called on non-Windows platform")
|
||||
}
|
||||
|
||||
func removeWindowsAddress(interfaceName string, ip net.IP, ipNet *net.IPNet) error {
|
||||
return fmt.Errorf("removeWindowsAddress called on non-Windows platform")
|
||||
}
|
||||
|
||||
@@ -61,3 +61,35 @@ func configureWindows(interfaceName string, ip net.IP, ipNet *net.IPNet) error {
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func removeWindowsAddress(interfaceName string, ip net.IP, ipNet *net.IPNet) error {
|
||||
iface, err := net.InterfaceByName(interfaceName)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get interface %s: %v", interfaceName, err)
|
||||
}
|
||||
|
||||
luid, err := winipcfg.LUIDFromIndex(uint32(iface.Index))
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get LUID for interface %s: %v", interfaceName, err)
|
||||
}
|
||||
|
||||
maskBits, _ := ipNet.Mask.Size()
|
||||
|
||||
var addr netip.Addr
|
||||
if ip4 := ip.To4(); ip4 != nil {
|
||||
addr, _ = netip.AddrFromSlice(ip4)
|
||||
} else {
|
||||
addr, _ = netip.AddrFromSlice(ip)
|
||||
}
|
||||
if !addr.IsValid() {
|
||||
return fmt.Errorf("failed to convert IP address")
|
||||
}
|
||||
prefix := netip.PrefixFrom(addr, maskBits)
|
||||
|
||||
logger.Info("Removing IP address %s from interface %s", prefix.String(), interfaceName)
|
||||
if err := luid.DeleteIPAddress(prefix); err != nil {
|
||||
return fmt.Errorf("failed to remove IP address: %v", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,164 @@
|
||||
package network
|
||||
|
||||
import (
|
||||
"net"
|
||||
"regexp"
|
||||
"sort"
|
||||
"strconv"
|
||||
|
||||
"github.com/fosrl/newt/logger"
|
||||
)
|
||||
|
||||
// Interface name patterns used to rank candidate local endpoints. Interfaces
|
||||
// matching physicalInterfacePatterns are tried first, interfaces matching
|
||||
// virtualInterfacePatterns (container/VPN/hypervisor bridges and the like)
|
||||
// are tried last, and everything else falls in between.
|
||||
var (
|
||||
physicalInterfacePatterns = []*regexp.Regexp{
|
||||
regexp.MustCompile(`(?i)^eth\d+$`),
|
||||
regexp.MustCompile(`(?i)^en\d+$`),
|
||||
regexp.MustCompile(`(?i)^eno\d+$`),
|
||||
regexp.MustCompile(`(?i)^ens\d+$`),
|
||||
regexp.MustCompile(`(?i)^enp\d+s\d+`),
|
||||
regexp.MustCompile(`(?i)^wlan\d*$`),
|
||||
regexp.MustCompile(`(?i)^wlp\d+s\d+`),
|
||||
regexp.MustCompile(`(?i)^wl\d+$`),
|
||||
regexp.MustCompile(`(?i)ethernet`),
|
||||
regexp.MustCompile(`(?i)wi-?fi`),
|
||||
regexp.MustCompile(`(?i)wireless`),
|
||||
}
|
||||
|
||||
virtualInterfacePatterns = []*regexp.Regexp{
|
||||
regexp.MustCompile(`(?i)docker`),
|
||||
regexp.MustCompile(`(?i)podman`),
|
||||
regexp.MustCompile(`(?i)^veth`),
|
||||
regexp.MustCompile(`(?i)^virbr`),
|
||||
regexp.MustCompile(`(?i)vmnet`),
|
||||
regexp.MustCompile(`(?i)vboxnet`),
|
||||
regexp.MustCompile(`(?i)virtualbox`),
|
||||
regexp.MustCompile(`(?i)^vbox`),
|
||||
regexp.MustCompile(`(?i)vmware`),
|
||||
regexp.MustCompile(`(?i)hyper-?v`),
|
||||
regexp.MustCompile(`(?i)vethernet`),
|
||||
regexp.MustCompile(`(?i)npcap`),
|
||||
regexp.MustCompile(`(?i)^tun\d*$`),
|
||||
regexp.MustCompile(`(?i)^tap\d*$`),
|
||||
regexp.MustCompile(`(?i)^wg\d*$`),
|
||||
regexp.MustCompile(`(?i)^utun\d*$`),
|
||||
regexp.MustCompile(`(?i)zerotier`),
|
||||
regexp.MustCompile(`(?i)^zt`),
|
||||
regexp.MustCompile(`(?i)tailscale`),
|
||||
regexp.MustCompile(`(?i)^ppp\d*$`),
|
||||
regexp.MustCompile(`(?i)bridge`),
|
||||
regexp.MustCompile(`(?i)^br-`),
|
||||
regexp.MustCompile(`(?i)^br\d+$`),
|
||||
regexp.MustCompile(`(?i)^cni`),
|
||||
regexp.MustCompile(`(?i)flannel`),
|
||||
regexp.MustCompile(`(?i)weave`),
|
||||
regexp.MustCompile(`(?i)kube`),
|
||||
regexp.MustCompile(`(?i)isatap`),
|
||||
regexp.MustCompile(`(?i)teredo`),
|
||||
regexp.MustCompile(`(?i)bluetooth`),
|
||||
regexp.MustCompile(`(?i)^awdl\d*$`),
|
||||
regexp.MustCompile(`(?i)^llw\d*$`),
|
||||
regexp.MustCompile(`(?i)p2p`),
|
||||
}
|
||||
)
|
||||
|
||||
const (
|
||||
scorePhysical = 0
|
||||
scoreUnknown = 10
|
||||
scoreVirtual = 20
|
||||
)
|
||||
|
||||
// interfaceScore ranks an interface name by how likely it is to be a
|
||||
// real, usable network interface. Lower scores are tried first.
|
||||
func interfaceScore(name string) int {
|
||||
for _, re := range physicalInterfacePatterns {
|
||||
if re.MatchString(name) {
|
||||
return scorePhysical
|
||||
}
|
||||
}
|
||||
for _, re := range virtualInterfacePatterns {
|
||||
if re.MatchString(name) {
|
||||
return scoreVirtual
|
||||
}
|
||||
}
|
||||
return scoreUnknown
|
||||
}
|
||||
|
||||
// GetLocalEndpoints returns "ip:port" strings (bracketed for IPv6, e.g.
|
||||
// "[fe80::1]:51820") for every usable, non-loopback IP address bound to a
|
||||
// network interface on this host. The list is ordered with interfaces most
|
||||
// likely to be a genuine host network (wired/Wi-Fi) first, and interfaces
|
||||
// that are typically synthetic (Docker, VPN tunnels, hypervisor bridges,
|
||||
// etc.) last, so callers should try the results roughly in order.
|
||||
//
|
||||
// excludeInterface, if non-empty, is skipped entirely - this is normally the
|
||||
// name of our own WireGuard/TUN interface, whose address is the tunnel IP
|
||||
// and not a useful endpoint to advertise.
|
||||
//
|
||||
// If interfaces cannot be enumerated (e.g. insufficient OS permissions),
|
||||
// an info message is logged and an empty slice is returned.
|
||||
func GetLocalEndpoints(port uint16, excludeInterface string) []string {
|
||||
ifaces, err := net.Interfaces()
|
||||
if err != nil {
|
||||
logger.Info("Unable to enumerate local network interfaces, localEndpoints will not be reported: %v", err)
|
||||
return nil
|
||||
}
|
||||
|
||||
type candidate struct {
|
||||
score int
|
||||
ip string
|
||||
}
|
||||
var candidates []candidate
|
||||
|
||||
for _, iface := range ifaces {
|
||||
if excludeInterface != "" && iface.Name == excludeInterface {
|
||||
continue
|
||||
}
|
||||
if iface.Flags&net.FlagUp == 0 || iface.Flags&net.FlagLoopback != 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
addrs, err := iface.Addrs()
|
||||
if err != nil {
|
||||
logger.Debug("Unable to read addresses for interface %s: %v", iface.Name, err)
|
||||
continue
|
||||
}
|
||||
|
||||
baseScore := interfaceScore(iface.Name)
|
||||
|
||||
for _, addr := range addrs {
|
||||
var ip net.IP
|
||||
switch v := addr.(type) {
|
||||
case *net.IPNet:
|
||||
ip = v.IP
|
||||
case *net.IPAddr:
|
||||
ip = v.IP
|
||||
}
|
||||
if ip == nil || ip.IsLoopback() || ip.IsUnspecified() {
|
||||
continue
|
||||
}
|
||||
if ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() {
|
||||
// Link-local addresses (169.254.0.0/16, fe80::/10) aren't
|
||||
// routable off the local segment, so they're never a
|
||||
// reachable endpoint for a peer.
|
||||
continue
|
||||
}
|
||||
|
||||
candidates = append(candidates, candidate{score: baseScore, ip: ip.String()})
|
||||
}
|
||||
}
|
||||
|
||||
sort.SliceStable(candidates, func(i, j int) bool {
|
||||
return candidates[i].score < candidates[j].score
|
||||
})
|
||||
|
||||
portStr := strconv.Itoa(int(port))
|
||||
endpoints := make([]string, 0, len(candidates))
|
||||
for _, c := range candidates {
|
||||
endpoints = append(endpoints, net.JoinHostPort(c.ip, portStr))
|
||||
}
|
||||
return endpoints
|
||||
}
|
||||
+60
-10
@@ -11,6 +11,33 @@ import (
|
||||
"github.com/vishvananda/netlink"
|
||||
)
|
||||
|
||||
// VPNRouteMetric is the route metric/priority assigned to routes we add for
|
||||
// the tunnel, so that an overlapping local/connected route is always
|
||||
// preferred over the VPN route to the same destination rather than the two
|
||||
// silently racing based on insertion order. It needs to be higher than any
|
||||
// metric a local route is realistically going to have: on Linux, automatic
|
||||
// metrics assigned by NetworkManager (which also apply to the connected
|
||||
// subnet route, not just the default route) go up to 600 for Wi-Fi; on
|
||||
// Windows, automatic interface metrics plus route metric rarely exceed a few
|
||||
// hundred. 9999 comfortably clears both without needing to query the local
|
||||
// routing table at add-time.
|
||||
const VPNRouteMetric = 9999
|
||||
|
||||
// PreferLocalRoutes controls whether routes added by AddRoutes are given the
|
||||
// explicit high VPNRouteMetric priority, so that an overlapping local/
|
||||
// connected route always takes precedence over the VPN route to the same
|
||||
// destination. Defaults to false (routes are added with the OS default
|
||||
// metric/priority, matching behavior prior to the introduction of
|
||||
// VPNRouteMetric); callers that want local routes to win opt in by setting
|
||||
// this to true (e.g. from a config value) before routes are added.
|
||||
var PreferLocalRoutes = false
|
||||
|
||||
// DarwinAddRoute adds a route via the BSD routing table. Unlike Linux/Windows,
|
||||
// BSD's routing table has no per-route metric - preference between an
|
||||
// overlapping local route and this VPN route is instead resolved by
|
||||
// longest-prefix-match, and `route add` (as opposed to `route change`) fails
|
||||
// rather than replacing an existing route to the same destination, so a local
|
||||
// route is never displaced by one we add here.
|
||||
func DarwinAddRoute(destination string, gateway string, interfaceName string) error {
|
||||
if runtime.GOOS != "darwin" {
|
||||
return nil
|
||||
@@ -65,10 +92,16 @@ func LinuxAddRoute(destination string, gateway string, interfaceName string) err
|
||||
return fmt.Errorf("invalid destination address: %v", err)
|
||||
}
|
||||
|
||||
// Create route
|
||||
// Create route. When PreferLocalRoutes is enabled, Priority is set
|
||||
// explicitly (rather than left at the default of 0) so that this route
|
||||
// never outranks a local/connected route to the same destination - see
|
||||
// VPNRouteMetric.
|
||||
route := &netlink.Route{
|
||||
Dst: ipNet,
|
||||
}
|
||||
if PreferLocalRoutes {
|
||||
route.Priority = VPNRouteMetric
|
||||
}
|
||||
|
||||
if gateway != "" {
|
||||
// Route with specific gateway
|
||||
@@ -98,7 +131,7 @@ func LinuxAddRoute(destination string, gateway string, interfaceName string) err
|
||||
return nil
|
||||
}
|
||||
|
||||
func LinuxRemoveRoute(destination string) error {
|
||||
func LinuxRemoveRoute(destination string, interfaceName string) error {
|
||||
if runtime.GOOS != "linux" {
|
||||
return nil
|
||||
}
|
||||
@@ -109,12 +142,26 @@ func LinuxRemoveRoute(destination string) error {
|
||||
return fmt.Errorf("invalid destination address: %v", err)
|
||||
}
|
||||
|
||||
// Create route to delete
|
||||
// Create route to delete. LinkIndex and Priority are set to match the
|
||||
// route we added exactly, so this only ever deletes the route we own -
|
||||
// a local/native route to the same destination on a different
|
||||
// interface (or with a different metric) must never be touched.
|
||||
route := &netlink.Route{
|
||||
Dst: ipNet,
|
||||
}
|
||||
if PreferLocalRoutes {
|
||||
route.Priority = VPNRouteMetric
|
||||
}
|
||||
|
||||
logger.Info("Removing route to %s", destination)
|
||||
if interfaceName != "" {
|
||||
link, err := netlink.LinkByName(interfaceName)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get interface %s: %v", interfaceName, err)
|
||||
}
|
||||
route.LinkIndex = link.Attrs().Index
|
||||
}
|
||||
|
||||
logger.Info("Removing route to %s via interface %s", destination, interfaceName)
|
||||
|
||||
// Delete the route
|
||||
if err := netlink.RouteDel(route); err != nil {
|
||||
@@ -157,9 +204,9 @@ func RemoveRouteForServerIP(serverIP string, interfaceName string) error {
|
||||
return DarwinRemoveRoute(serverIP)
|
||||
}
|
||||
// else if runtime.GOOS == "windows" {
|
||||
// return WindowsRemoveRoute(serverIP)
|
||||
// return WindowsRemoveRoute(serverIP, interfaceName)
|
||||
// } else if runtime.GOOS == "linux" {
|
||||
// return LinuxRemoveRoute(serverIP)
|
||||
// return LinuxRemoveRoute(serverIP, interfaceName)
|
||||
// }
|
||||
return nil
|
||||
}
|
||||
@@ -242,8 +289,11 @@ func AddRoutes(remoteSubnets []string, interfaceName string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// removeRoutesForRemoteSubnets removes routes for each subnet in RemoteSubnets
|
||||
func RemoveRoutes(remoteSubnets []string) error {
|
||||
// removeRoutesForRemoteSubnets removes routes for each subnet in RemoteSubnets.
|
||||
// interfaceName must match the interface the routes were added on (see
|
||||
// AddRoutes) so that only the routes we own are deleted, never an unrelated
|
||||
// local/native route to the same destination on another interface.
|
||||
func RemoveRoutes(remoteSubnets []string, interfaceName string) error {
|
||||
if len(remoteSubnets) == 0 {
|
||||
return nil
|
||||
}
|
||||
@@ -267,11 +317,11 @@ func RemoveRoutes(remoteSubnets []string) error {
|
||||
logger.Error("Failed to remove Darwin route for subnet %s: %v", subnet, err)
|
||||
}
|
||||
case "windows":
|
||||
if err := WindowsRemoveRoute(subnet); err != nil {
|
||||
if err := WindowsRemoveRoute(subnet, interfaceName); err != nil {
|
||||
logger.Error("Failed to remove Windows route for subnet %s: %v", subnet, err)
|
||||
}
|
||||
case "linux":
|
||||
if err := LinuxRemoveRoute(subnet); err != nil {
|
||||
if err := LinuxRemoveRoute(subnet, interfaceName); err != nil {
|
||||
logger.Error("Failed to remove Linux route for subnet %s: %v", subnet, err)
|
||||
}
|
||||
case "android", "ios":
|
||||
|
||||
@@ -6,6 +6,6 @@ func WindowsAddRoute(destination string, gateway string, interfaceName string) e
|
||||
return nil
|
||||
}
|
||||
|
||||
func WindowsRemoveRoute(destination string) error {
|
||||
func WindowsRemoveRoute(destination string, interfaceName string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
+46
-12
@@ -84,8 +84,15 @@ func WindowsAddRoute(destination string, gateway string, interfaceName string) e
|
||||
return fmt.Errorf("either gateway or interface must be specified")
|
||||
}
|
||||
|
||||
// Add the route using winipcfg
|
||||
err = luid.AddRoute(prefix, nextHop, 1)
|
||||
// Add the route using winipcfg. When PreferLocalRoutes is enabled,
|
||||
// metric is set explicitly (rather than a low value like 1, which would
|
||||
// nearly always outrank local routes) so that an overlapping local/
|
||||
// connected route is preferred over this VPN route - see VPNRouteMetric.
|
||||
var metric uint32
|
||||
if PreferLocalRoutes {
|
||||
metric = VPNRouteMetric
|
||||
}
|
||||
err = luid.AddRoute(prefix, nextHop, metric)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to add route: %v", err)
|
||||
}
|
||||
@@ -93,7 +100,7 @@ func WindowsAddRoute(destination string, gateway string, interfaceName string) e
|
||||
return nil
|
||||
}
|
||||
|
||||
func WindowsRemoveRoute(destination string) error {
|
||||
func WindowsRemoveRoute(destination string, interfaceName string) error {
|
||||
// Parse destination CIDR
|
||||
_, ipNet, err := net.ParseCIDR(destination)
|
||||
if err != nil {
|
||||
@@ -117,8 +124,25 @@ func WindowsRemoveRoute(destination string) error {
|
||||
}
|
||||
prefix := netip.PrefixFrom(addr, maskBits)
|
||||
|
||||
// Resolve the LUID of the interface we added the route on, so we only
|
||||
// ever delete the route we own rather than any route matching the
|
||||
// destination - a local/native route to the same destination on a
|
||||
// different interface must never be touched.
|
||||
var luid winipcfg.LUID
|
||||
var haveLuid bool
|
||||
if interfaceName != "" {
|
||||
iface, err := net.InterfaceByName(interfaceName)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get interface %s: %v", interfaceName, err)
|
||||
}
|
||||
luid, err = winipcfg.LUIDFromIndex(uint32(iface.Index))
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get LUID for interface %s: %v", interfaceName, err)
|
||||
}
|
||||
haveLuid = true
|
||||
}
|
||||
|
||||
// Get all routes and find the one to delete
|
||||
// We need to get the LUID from the existing route
|
||||
var family winipcfg.AddressFamily
|
||||
if addr.Is4() {
|
||||
family = 2 // AF_INET
|
||||
@@ -131,17 +155,27 @@ func WindowsRemoveRoute(destination string) error {
|
||||
return fmt.Errorf("failed to get route table: %v", err)
|
||||
}
|
||||
|
||||
// Find and delete matching route
|
||||
// Find and delete matching route. When we know which interface we added
|
||||
// the route on, only delete the entry on that interface with the metric
|
||||
// we added it with (see PreferLocalRoutes) so we never remove an
|
||||
// unrelated local/native route to the same destination.
|
||||
var wantMetric uint32
|
||||
if PreferLocalRoutes {
|
||||
wantMetric = VPNRouteMetric
|
||||
}
|
||||
for _, route := range routes {
|
||||
routePrefix := route.DestinationPrefix.Prefix()
|
||||
if routePrefix == prefix {
|
||||
logger.Info("Removing route to %s", destination)
|
||||
err = route.Delete()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to delete route: %v", err)
|
||||
}
|
||||
return nil
|
||||
if routePrefix != prefix {
|
||||
continue
|
||||
}
|
||||
if haveLuid && (route.InterfaceLUID != luid || route.Metric != wantMetric) {
|
||||
continue
|
||||
}
|
||||
logger.Info("Removing route to %s on interface %s", destination, interfaceName)
|
||||
if err := route.Delete(); err != nil {
|
||||
return fmt.Errorf("failed to delete route: %v", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
return fmt.Errorf("route to %s not found", destination)
|
||||
|
||||
@@ -81,6 +81,45 @@ func SetIPv4Settings(addresses []string, subnetMasks []string) {
|
||||
logger.Info("Set IPv4 addresses: %v, subnet masks: %v", addresses, subnetMasks)
|
||||
}
|
||||
|
||||
// AddIPv4Address appends an additional IPv4 address/subnet mask pair to the
|
||||
// tunnel's network settings. This is how a secondary interface address gets
|
||||
// exposed to mobile (iOS/Android) packet-tunnel providers, which read the
|
||||
// full IPv4Addresses/IPv4SubnetMasks arrays (not just the first entry) and
|
||||
// re-apply them on every settings poll.
|
||||
func AddIPv4Address(address string, subnetMask string) {
|
||||
networkSettingsMutex.Lock()
|
||||
defer networkSettingsMutex.Unlock()
|
||||
|
||||
for _, a := range networkSettings.IPv4Addresses {
|
||||
if a == address {
|
||||
logger.Info("IPv4 address already exists: %s", address)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
networkSettings.IPv4Addresses = append(networkSettings.IPv4Addresses, address)
|
||||
networkSettings.IPv4SubnetMasks = append(networkSettings.IPv4SubnetMasks, subnetMask)
|
||||
incrementor++
|
||||
logger.Info("Added IPv4 address: %s/%s", address, subnetMask)
|
||||
}
|
||||
|
||||
// RemoveIPv4Address removes a previously added secondary IPv4 address.
|
||||
func RemoveIPv4Address(address string) {
|
||||
networkSettingsMutex.Lock()
|
||||
defer networkSettingsMutex.Unlock()
|
||||
|
||||
for i, a := range networkSettings.IPv4Addresses {
|
||||
if a == address {
|
||||
networkSettings.IPv4Addresses = append(networkSettings.IPv4Addresses[:i], networkSettings.IPv4Addresses[i+1:]...)
|
||||
networkSettings.IPv4SubnetMasks = append(networkSettings.IPv4SubnetMasks[:i], networkSettings.IPv4SubnetMasks[i+1:]...)
|
||||
incrementor++
|
||||
logger.Info("Removed IPv4 address: %s", address)
|
||||
return
|
||||
}
|
||||
}
|
||||
logger.Info("IPv4 address not found for removal: %s", address)
|
||||
}
|
||||
|
||||
// SetIPv4IncludedRoutes sets the included IPv4 routes
|
||||
func SetIPv4IncludedRoutes(routes []IPv4Route) {
|
||||
networkSettingsMutex.Lock()
|
||||
|
||||
@@ -144,6 +144,7 @@ func (n *Newt) handleSync(msg websocket.WSMessage) {
|
||||
|
||||
// Sync clients WireGuard peers and targets, if clients are set up
|
||||
if n.wgService != nil {
|
||||
n.wgService.SetCerts(syncData.Certs)
|
||||
n.wgService.Sync(syncData.Peers, syncData.ClientTargets)
|
||||
}
|
||||
|
||||
|
||||
+3
-125
@@ -10,13 +10,13 @@ import (
|
||||
"net/http"
|
||||
"os"
|
||||
"os/signal"
|
||||
"strings"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"github.com/fosrl/newt/authdaemon"
|
||||
"github.com/fosrl/newt/browsergateway"
|
||||
"github.com/fosrl/newt/docker"
|
||||
"github.com/fosrl/newt/exitnode"
|
||||
"github.com/fosrl/newt/healthcheck"
|
||||
"github.com/fosrl/newt/internal/state"
|
||||
"github.com/fosrl/newt/internal/telemetry"
|
||||
@@ -140,129 +140,7 @@ func (n *Newt) registerHandlers(ctx context.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
if len(exitNodes) == 1 || n.config.PreferEndpoint != "" {
|
||||
logger.Debug("Only one exit node available, using it directly: %s", exitNodes[0].Endpoint)
|
||||
|
||||
if n.config.PreferEndpoint != "" {
|
||||
for _, node := range exitNodes {
|
||||
if node.Endpoint == n.config.PreferEndpoint {
|
||||
exitNodes[0] = node
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pingResults := []ExitNodePingResult{
|
||||
{
|
||||
ExitNodeID: exitNodes[0].ID,
|
||||
LatencyMs: 0,
|
||||
Weight: exitNodes[0].Weight,
|
||||
Error: "",
|
||||
Name: exitNodes[0].Name,
|
||||
Endpoint: exitNodes[0].Endpoint,
|
||||
WasPreviouslyConnected: exitNodes[0].WasPreviouslyConnected,
|
||||
},
|
||||
}
|
||||
|
||||
chainId := generateChainId()
|
||||
n.pendingRegisterChainId = chainId
|
||||
n.stopFunc = n.client.SendMessageInterval(topicWGRegister, map[string]interface{}{
|
||||
"publicKey": n.publicKey.String(),
|
||||
"pingResults": pingResults,
|
||||
"newtVersion": n.config.Version,
|
||||
"chainId": chainId,
|
||||
}, 2*time.Second)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
type nodeResult struct {
|
||||
Node ExitNode
|
||||
Latency time.Duration
|
||||
Err error
|
||||
}
|
||||
|
||||
results := make([]nodeResult, len(exitNodes))
|
||||
const pingAttempts = 3
|
||||
for i, node := range exitNodes {
|
||||
var totalLatency time.Duration
|
||||
var lastErr error
|
||||
successes := 0
|
||||
httpClient := &http.Client{
|
||||
Timeout: 5 * time.Second,
|
||||
}
|
||||
url := node.Endpoint
|
||||
if !strings.HasPrefix(url, "http://") && !strings.HasPrefix(url, "https://") {
|
||||
url = "http://" + url
|
||||
}
|
||||
if !strings.HasSuffix(url, "/ping") {
|
||||
url = strings.TrimRight(url, "/") + "/ping"
|
||||
}
|
||||
for j := 0; j < pingAttempts; j++ {
|
||||
start := time.Now()
|
||||
resp, err := httpClient.Get(url)
|
||||
latency := time.Since(start)
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
logger.Warn("Failed to ping exit node %d (%s) attempt %d: %v", node.ID, url, j+1, err)
|
||||
continue
|
||||
}
|
||||
resp.Body.Close()
|
||||
totalLatency += latency
|
||||
successes++
|
||||
}
|
||||
var avgLatency time.Duration
|
||||
if successes > 0 {
|
||||
avgLatency = totalLatency / time.Duration(successes)
|
||||
}
|
||||
if successes == 0 {
|
||||
results[i] = nodeResult{Node: node, Latency: 0, Err: lastErr}
|
||||
} else {
|
||||
results[i] = nodeResult{Node: node, Latency: avgLatency, Err: nil}
|
||||
}
|
||||
}
|
||||
|
||||
var pingResults []ExitNodePingResult
|
||||
for _, res := range results {
|
||||
errMsg := ""
|
||||
if res.Err != nil {
|
||||
errMsg = res.Err.Error()
|
||||
}
|
||||
pingResults = append(pingResults, ExitNodePingResult{
|
||||
ExitNodeID: res.Node.ID,
|
||||
LatencyMs: res.Latency.Milliseconds(),
|
||||
Weight: res.Node.Weight,
|
||||
Error: errMsg,
|
||||
Name: res.Node.Name,
|
||||
Endpoint: res.Node.Endpoint,
|
||||
WasPreviouslyConnected: res.Node.WasPreviouslyConnected,
|
||||
})
|
||||
}
|
||||
|
||||
if n.connected {
|
||||
var filteredPingResults []ExitNodePingResult
|
||||
previouslyConnectedNodeIdx := -1
|
||||
for i, res := range pingResults {
|
||||
if res.WasPreviouslyConnected {
|
||||
previouslyConnectedNodeIdx = i
|
||||
}
|
||||
}
|
||||
goodNodeCount := 0
|
||||
for i, res := range pingResults {
|
||||
if i != previouslyConnectedNodeIdx && res.LatencyMs > 0 && res.Error == "" {
|
||||
goodNodeCount++
|
||||
}
|
||||
}
|
||||
if previouslyConnectedNodeIdx != -1 && goodNodeCount > 0 {
|
||||
for i, res := range pingResults {
|
||||
if i != previouslyConnectedNodeIdx {
|
||||
filteredPingResults = append(filteredPingResults, res)
|
||||
}
|
||||
}
|
||||
pingResults = filteredPingResults
|
||||
logger.Info("Excluding previously connected exit node from ping results due to other available nodes")
|
||||
}
|
||||
}
|
||||
pingResults := exitnode.PingExitNodes(exitNodes, n.config.PreferEndpoint, n.connected)
|
||||
|
||||
chainId := generateChainId()
|
||||
n.pendingRegisterChainId = chainId
|
||||
@@ -430,7 +308,7 @@ func (n *Newt) registerHandlers(ctx context.Context) {
|
||||
}
|
||||
|
||||
if n.config.UseNativeMainInterface {
|
||||
if err := network.RemoveRoutes(data.Subnets); err != nil {
|
||||
if err := network.RemoveRoutes(data.Subnets, n.config.NativeMainInterfaceName); err != nil {
|
||||
logger.Warn("Failed to remove routes for subnets: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
+2
-2
@@ -24,7 +24,7 @@ func (n *Newt) updateRemoteExitNodeSubnets(subnets []string) {
|
||||
}
|
||||
}
|
||||
if len(toRemove) > 0 {
|
||||
if err := network.RemoveRoutes(toRemove); err != nil {
|
||||
if err := network.RemoveRoutes(toRemove, n.config.NativeMainInterfaceName); err != nil {
|
||||
logger.Warn("Failed to remove old subnet routes: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -88,7 +88,7 @@ func (n *Newt) closeWgTunnel() {
|
||||
}
|
||||
toRemove = append(toRemove, n.activeRemoteSubnets...)
|
||||
if len(toRemove) > 0 {
|
||||
if err := network.RemoveRoutes(toRemove); err != nil {
|
||||
if err := network.RemoveRoutes(toRemove, n.config.NativeMainInterfaceName); err != nil {
|
||||
logger.Warn("Failed to remove native main tunnel routes: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
+10
-22
@@ -2,6 +2,7 @@ package newt
|
||||
|
||||
import (
|
||||
wgclients "github.com/fosrl/newt/clients"
|
||||
"github.com/fosrl/newt/exitnode"
|
||||
"github.com/fosrl/newt/healthcheck"
|
||||
)
|
||||
|
||||
@@ -35,28 +36,14 @@ type TargetData struct {
|
||||
Targets []string `json:"targets"`
|
||||
}
|
||||
|
||||
type ExitNodeData struct {
|
||||
ExitNodes []ExitNode `json:"exitNodes"`
|
||||
ChainId string `json:"chainId"`
|
||||
}
|
||||
|
||||
type ExitNode struct {
|
||||
ID int `json:"exitNodeId"`
|
||||
Name string `json:"exitNodeName"`
|
||||
Endpoint string `json:"endpoint"`
|
||||
Weight float64 `json:"weight"`
|
||||
WasPreviouslyConnected bool `json:"wasPreviouslyConnected"`
|
||||
}
|
||||
|
||||
type ExitNodePingResult struct {
|
||||
ExitNodeID int `json:"exitNodeId"`
|
||||
LatencyMs int64 `json:"latencyMs"`
|
||||
Weight float64 `json:"weight"`
|
||||
Error string `json:"error,omitempty"`
|
||||
Name string `json:"exitNodeName"`
|
||||
Endpoint string `json:"endpoint"`
|
||||
WasPreviouslyConnected bool `json:"wasPreviouslyConnected"`
|
||||
}
|
||||
// ExitNodeData, ExitNode and ExitNodePingResult are aliases for the shared
|
||||
// exit-node ping dance types in package exitnode, kept here so existing code
|
||||
// in this package can keep referring to them unqualified.
|
||||
type (
|
||||
ExitNodeData = exitnode.ExitNodeData
|
||||
ExitNode = exitnode.ExitNode
|
||||
ExitNodePingResult = exitnode.ExitNodePingResult
|
||||
)
|
||||
|
||||
type BlueprintResult struct {
|
||||
Success bool `json:"success"`
|
||||
@@ -70,5 +57,6 @@ type SyncData struct {
|
||||
RemoteExitNodeSubnets []string `json:"remoteExitNodeSubnets"`
|
||||
Peers []wgclients.Peer `json:"peers"`
|
||||
ClientTargets []wgclients.Target `json:"clientTargets"`
|
||||
Certs []wgclients.CertData `json:"certs"`
|
||||
BrowserGatewayTargets []BrowserGatewayTarget `json:"browserGatewayTargets"`
|
||||
}
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
# Maintainer: Fossorial <hello@fossorial.io>
|
||||
pkgname=newt
|
||||
pkgver=1.13.0
|
||||
pkgrel=1
|
||||
pkgdesc="Fully user space WireGuard tunnel client and TCP/UDP proxy for Pangolin"
|
||||
arch=('x86_64' 'aarch64' 'armv7h' 'armv6h' 'riscv64')
|
||||
url="https://github.com/fosrl/newt"
|
||||
license=('AGPL3')
|
||||
makedepends=('go')
|
||||
source=("$pkgname-$pkgver.tar.gz::https://github.com/fosrl/newt/archive/refs/tags/v$pkgver.tar.gz")
|
||||
sha256sums=('SKIP')
|
||||
|
||||
build() {
|
||||
cd "$pkgname-$pkgver"
|
||||
export CGO_ENABLED=0
|
||||
export GOFLAGS="-trimpath -mod=readonly -modcacherw"
|
||||
go build -ldflags "-X main.newtVersion=$pkgver -X main.newtPlatform=linux_$(go env GOARCH)" -o "$pkgname" .
|
||||
}
|
||||
|
||||
check() {
|
||||
cd "$pkgname-$pkgver"
|
||||
go test ./... || true
|
||||
}
|
||||
|
||||
package() {
|
||||
cd "$pkgname-$pkgver"
|
||||
install -Dm755 "$pkgname" "$pkgdir/usr/bin/$pkgname"
|
||||
install -Dm644 LICENSE "$pkgdir/usr/share/licenses/$pkgname/LICENSE"
|
||||
install -Dm644 README.md "$pkgdir/usr/share/doc/$pkgname/README.md"
|
||||
}
|
||||
+20
-2
@@ -28,6 +28,15 @@ import (
|
||||
"go.opentelemetry.io/otel"
|
||||
)
|
||||
|
||||
// writeDeadline bounds how long a websocket write may block before it is
|
||||
// treated as a failure. Without this, a write to a TCP connection whose
|
||||
// underlying network interface has disappeared (e.g. laptop sleep/resume,
|
||||
// Wi-Fi roam) can sit buffered in the kernel for minutes without erroring.
|
||||
// This matters even with the read-deadline/pong machinery below: if the
|
||||
// WriteJSON call in sendPing blocks, execution never reaches the
|
||||
// WriteControl ping that would otherwise trigger that read-side detection.
|
||||
const writeDeadline = 10 * time.Second
|
||||
|
||||
type Client struct {
|
||||
conn *websocket.Conn
|
||||
config *Config
|
||||
@@ -257,6 +266,9 @@ func (c *Client) SendMessage(messageType string, data interface{}) error {
|
||||
|
||||
c.writeMux.Lock()
|
||||
defer c.writeMux.Unlock()
|
||||
if err := c.conn.SetWriteDeadline(time.Now().Add(writeDeadline)); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := c.conn.WriteJSON(msg); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -277,6 +289,9 @@ func (c *Client) SendMessageNoLog(messageType string, data interface{}) error {
|
||||
|
||||
c.writeMux.Lock()
|
||||
defer c.writeMux.Unlock()
|
||||
if err := c.conn.SetWriteDeadline(time.Now().Add(writeDeadline)); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := c.conn.WriteJSON(msg); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -760,14 +775,17 @@ func (c *Client) sendPing() {
|
||||
c.writeMux.Unlock()
|
||||
return
|
||||
}
|
||||
err := c.conn.WriteJSON(pingMsg)
|
||||
err := c.conn.SetWriteDeadline(time.Now().Add(writeDeadline))
|
||||
if err == nil {
|
||||
err = c.conn.WriteJSON(pingMsg)
|
||||
}
|
||||
if err == nil {
|
||||
telemetry.IncWSMessage(c.metricsContext(), "out", "ping")
|
||||
// Protocol-level ping: a standards-compliant server replies with a PONG,
|
||||
// which refreshes the read deadline. This is what lets us notice a
|
||||
// half-open connection where writes still succeed (buffered) but the
|
||||
// peer is gone.
|
||||
_ = c.conn.WriteControl(websocket.PingMessage, nil, time.Now().Add(10*time.Second))
|
||||
_ = c.conn.WriteControl(websocket.PingMessage, nil, time.Now().Add(writeDeadline))
|
||||
}
|
||||
c.writeMux.Unlock()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user