Exclude the websocket routes

This commit is contained in:
Owen
2026-10-01 12:27:24 -04:00
parent 815f3e506c
commit 98fe56ed13
6 changed files with 221 additions and 60 deletions
+71 -6
View File
@@ -3,12 +3,14 @@ package websocket
import (
"bytes"
"compress/gzip"
"context"
"crypto/tls"
"crypto/x509"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"net/http"
"net/url"
"os"
@@ -19,6 +21,7 @@ import (
"software.sslmate.com/src/go-pkcs12"
"github.com/fosrl/newt/logger"
"github.com/fosrl/newt/util"
"github.com/gorilla/websocket"
)
@@ -96,6 +99,8 @@ type Client struct {
onConnect func() error
onTokenUpdate func(token string, exitNodes []ExitNode)
onAuthError func(statusCode int, message string) // Callback for auth errors
onDialTargets func(ips []string) // Callback with the server IPs about to be dialed, see dialContext
publicDNS func() []string // Provides the DNS servers used to resolve the server host, see dialContext
writeMux sync.Mutex
clientType string // Type of client (e.g., "newt", "olm")
tlsConfig TLSConfig
@@ -156,6 +161,16 @@ func WithPingDataProvider(fn func() map[string]any) ClientOption {
}
}
// WithPublicDNSProvider sets a callback returning the DNS servers ("ip:port")
// used to resolve the server host when dialing (see dialContext). Called on
// every dial, so it may return a value that changes over time. If unset, or
// if every server fails, the platform resolver is used.
func WithPublicDNSProvider(fn func() []string) ClientOption {
return func(c *Client) {
c.publicDNS = fn
}
}
func (c *Client) OnConnect(callback func() error) {
c.onConnect = callback
}
@@ -168,6 +183,53 @@ func (c *Client) OnAuthError(callback func(statusCode int, message string)) {
c.onAuthError = callback
}
// OnDialTargets sets a callback invoked synchronously with every IP address
// the server host resolved to, immediately before each dial (token fetch and
// websocket). The callback runs before any packet is sent to those addresses,
// so it can install routes for them - see dialContext.
func (c *Client) OnDialTargets(callback func(ips []string)) {
c.onDialTargets = callback
}
// dialContext is the dial function for every connection to the server (token
// fetch and websocket, including through a proxy). It resolves the host
// itself rather than leaving it to net.Dialer, reports the resulting IPs via
// onDialTargets, and then dials only those IPs - so the caller knows exactly
// which addresses the connection can use and can keep them off the gateway
// (full-tunnel) default route before the first packet is sent. Each dial
// resolves again, so DNS changes are picked up on the next (re)connect; the
// address of a connection that is already established never changes.
func (c *Client) dialContext(ctx context.Context, network, addr string) (net.Conn, error) {
host, port, err := net.SplitHostPort(addr)
if err != nil {
return nil, err
}
var publicDNS []string
if c.publicDNS != nil {
publicDNS = c.publicDNS()
}
ips, err := util.ResolveDomainAllUpstream(host, publicDNS)
if err != nil {
return nil, fmt.Errorf("failed to resolve %s: %w", host, err)
}
if c.onDialTargets != nil {
c.onDialTargets(ips)
}
dialer := net.Dialer{Timeout: 10 * time.Second}
var lastErr error
for _, ip := range ips {
conn, err := dialer.DialContext(ctx, network, net.JoinHostPort(ip, port))
if err == nil {
return conn, nil
}
lastErr = err
}
return nil, lastErr
}
// NewClient creates a new websocket client
func NewClient(ID, secret, userToken, orgId, endpoint string, pingInterval time.Duration, opts ...ClientOption) (*Client, error) {
config := &Config{
@@ -483,12 +545,12 @@ func (c *Client) getToken() (string, []ExitNode, string, error) {
logger.Debug("websocket: Requesting token from %s with body: %s", req.URL.String(), string(jsonData))
// Make the request
client := &http.Client{}
transport := http.DefaultTransport.(*http.Transport).Clone()
transport.DialContext = c.dialContext
if tlsConfig != nil {
client.Transport = &http.Transport{
TLSClientConfig: tlsConfig,
}
transport.TLSClientConfig = tlsConfig
}
client := &http.Client{Transport: transport}
resp, err := client.Do(req)
if err != nil {
return "", nil, "", fmt.Errorf("failed to request new token: %w", err)
@@ -610,8 +672,11 @@ func (c *Client) establishConnection() error {
}
u.RawQuery = q.Encode()
// Connect to WebSocket
dialer := websocket.DefaultDialer
// Connect to WebSocket. Copy DefaultDialer rather than mutating the
// shared global.
dialerCopy := *websocket.DefaultDialer
dialer := &dialerCopy
dialer.NetDialContext = c.dialContext
// Use new TLS configuration method
if c.tlsConfig.ClientCertFile != "" || c.tlsConfig.ClientKeyFile != "" || len(c.tlsConfig.CAFiles) > 0 || c.tlsConfig.PKCS12File != "" {
+38
View File
@@ -1,6 +1,8 @@
package websocket
import (
"context"
"net"
"net/http"
"net/http/httptest"
"strings"
@@ -111,3 +113,39 @@ func TestReconnectClearsCurrentConnection(t *testing.T) {
t.Fatalf("expected conn to be cleared after reconnect(a) when a was still current, got %v", got)
}
}
// TestDialContextReportsTargetsBeforeDialing checks that dialContext reports
// the resolved server IPs via onDialTargets before connecting, and then
// connects to one of exactly those IPs - the contract olm relies on to install
// gateway bypass routes for the control connection.
func TestDialContextReportsTargetsBeforeDialing(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {}))
t.Cleanup(srv.Close)
_, port, err := net.SplitHostPort(strings.TrimPrefix(srv.URL, "http://"))
if err != nil {
t.Fatal(err)
}
c := newTestClient()
var reported []string
c.OnDialTargets(func(ips []string) {
if len(reported) != 0 {
t.Errorf("onDialTargets called more than once")
}
reported = append([]string(nil), ips...)
})
conn, err := c.dialContext(context.Background(), "tcp", net.JoinHostPort("127.0.0.1", port))
if err != nil {
t.Fatalf("dialContext: %v", err)
}
defer conn.Close()
if len(reported) != 1 || reported[0] != "127.0.0.1" {
t.Fatalf("onDialTargets got %v, want [127.0.0.1]", reported)
}
remoteHost, _, _ := net.SplitHostPort(conn.RemoteAddr().String())
if remoteHost != reported[0] {
t.Fatalf("dialed %s, but reported %v", remoteHost, reported)
}
}