mirror of
https://github.com/fosrl/olm.git
synced 2026-10-02 02:39:10 +02:00
Exclude the websocket routes
This commit is contained in:
@@ -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
|
||||
|
||||
+61
-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,63 @@ 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
|
||||
}
|
||||
|
||||
// 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()
|
||||
}
|
||||
|
||||
+31
-8
@@ -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()
|
||||
@@ -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
|
||||
|
||||
+19
-32
@@ -105,10 +105,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,13 +115,13 @@ 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
|
||||
@@ -323,7 +320,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 +345,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
|
||||
// resolving and excluding every currently-tracked peer's active endpoint and
|
||||
// every extra bypass endpoint (including the control plane - see
|
||||
// gatewayExtraEndpoints) (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)
|
||||
}
|
||||
|
||||
func (pm *PeerManager) activateGatewayLocked() error {
|
||||
for _, peer := range pm.peers {
|
||||
if ip, ok := pm.resolveActiveEndpointIPLocked(peer); ok {
|
||||
pm.excludeEndpointLocked(ip)
|
||||
@@ -381,8 +372,7 @@ func (pm *PeerManager) activateGatewayLocked(controlEndpointHost string) error {
|
||||
}
|
||||
|
||||
// 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
|
||||
// default-route-equivalent, then every remaining bypass route, and resets gateway exclusion state. Must be called with
|
||||
// pm.mu held.
|
||||
func (pm *PeerManager) deactivateGatewayLocked() {
|
||||
if err := network.RemoveGatewayDefaultRoute(pm.interfaceName); err != nil {
|
||||
@@ -395,7 +385,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 +439,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 +466,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 +637,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