mirror of
https://github.com/fosrl/olm.git
synced 2026-10-04 11:49:07 +02:00
Exclude the websocket routes
This commit is contained in:
+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