Skip to content

Commit 5ace75c

Browse files
committed
Extract applyRewrite helper; update context-cancellation test
Deduplicate the rewrite-detection, metrics-increment, and observer- notification block from resolveDiscoveredNodes and runAddressResolverRunner into (*Client).applyRewrite. Signed-off-by: Sean Chittenden <sean.chittenden@crowdstrike.com>
1 parent bbf2628 commit 5ace75c

2 files changed

Lines changed: 136 additions & 81 deletions

File tree

opensearchtransport/address_resolver_internal_test.go

Lines changed: 87 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -177,6 +177,18 @@ func TestAddressResolver(t *testing.T) {
177177
wantPort: map[string]string{"rewrite-me": "9201", "keep-me": "9200"},
178178
metrics: wantMetrics{calls: 3, rewrites: 1, errors: 1},
179179
},
180+
{
181+
name: "malformed URL keeps default and counts an error",
182+
inAddrs: map[string]string{"alpha": "10.0.0.1:9200", "beta": "10.0.0.2:9200"},
183+
resolver: func(_ context.Context, _ NodeInfo) (*url.URL, error) {
184+
// Empty Scheme/Host -- the kind of URL a resolver might
185+
// return from a buggy parse. The client must reject it
186+
// rather than enrolling a malformed connection.
187+
return &url.URL{}, nil
188+
},
189+
wantPort: map[string]string{"alpha": "9200", "beta": "9200"},
190+
metrics: wantMetrics{calls: 2, rewrites: 0, errors: 2},
191+
},
180192
}
181193

182194
for _, tt := range protocolTests {
@@ -389,7 +401,7 @@ func TestAddressResolver(t *testing.T) {
389401
ctx, cancel := context.WithCancel(t.Context())
390402
cancel()
391403
return ctx, func(_ context.Context, _ NodeInfo) (*url.URL, error) {
392-
return nil, nil //nolint:nilnil // testing (nil, nil) protocol case
404+
return nil, nil
393405
}
394406
},
395407
wantErr: context.Canceled,
@@ -399,12 +411,13 @@ func TestAddressResolver(t *testing.T) {
399411
setup: func(t *testing.T) (context.Context, AddressResolverFunc) {
400412
t.Helper()
401413
ctx, cancel := context.WithCancel(t.Context())
414+
t.Cleanup(cancel)
402415
var calls atomic.Int32
403416
return ctx, func(_ context.Context, _ NodeInfo) (*url.URL, error) {
404417
if calls.Add(1) == 1 {
405418
cancel()
406419
}
407-
return nil, nil //nolint:nilnil // testing (nil, nil) protocol case
420+
return nil, nil
408421
}
409422
},
410423
maxResolvers: 1,
@@ -433,7 +446,7 @@ func TestAddressResolver(t *testing.T) {
433446
return
434447
}
435448
require.NoError(t, err)
436-
require.Equal(t, tt.wantNodes, len(nodes))
449+
require.Len(t, nodes, tt.wantNodes)
437450
})
438451
}
439452

@@ -767,42 +780,83 @@ func TestAddressResolverRunner(t *testing.T) {
767780
"resolve param should be nil when AddressResolver is not set")
768781
})
769782

770-
t.Run("runner context cancellation propagated", func(t *testing.T) {
771-
t.Parallel()
772-
773-
ctx, cancel := context.WithCancel(t.Context())
774-
cancel()
775-
776-
tp, err := New(Config{
777-
URLs: []*url.URL{testSeedURL(t)},
778-
Transport: newResolverTestTransport(t, nodesJSON),
779-
HealthCheck: NoOpHealthCheck,
780-
AddressResolverRunner: func(ctx context.Context, _ []NodeInfo, _ AddressResolverFunc) ([]ResolvedAddress, error) {
783+
runnerEdgeCases := []struct {
784+
name string
785+
makeCtx func(t *testing.T) context.Context
786+
runner AddressResolverRunnerFunc
787+
wantErr error
788+
wantNodes int
789+
}{
790+
{
791+
name: "context cancellation propagated",
792+
makeCtx: func(t *testing.T) context.Context {
793+
t.Helper()
794+
ctx, cancel := context.WithCancel(t.Context())
795+
cancel()
796+
return ctx
797+
},
798+
runner: func(ctx context.Context, _ []NodeInfo, _ AddressResolverFunc) ([]ResolvedAddress, error) {
781799
return nil, ctx.Err()
782800
},
783-
})
784-
require.NoError(t, err)
801+
wantErr: context.Canceled,
802+
},
803+
{
804+
name: "drops all returns ErrAllResolversFailed",
805+
runner: func(_ context.Context, _ []NodeInfo, _ AddressResolverFunc) ([]ResolvedAddress, error) {
806+
return nil, nil
807+
},
808+
wantErr: ErrAllResolversFailed,
809+
},
810+
{
811+
name: "nil URLs are skipped",
812+
runner: func(_ context.Context, nodes []NodeInfo, _ AddressResolverFunc) ([]ResolvedAddress, error) {
813+
out := make([]ResolvedAddress, len(nodes))
814+
for i, n := range nodes {
815+
out[i] = ResolvedAddress{Node: n, URL: nil}
816+
}
817+
return out, nil
818+
},
819+
wantErr: ErrAllResolversFailed,
820+
},
821+
{
822+
name: "unknown node ID is ignored",
823+
runner: func(_ context.Context, nodes []NodeInfo, _ AddressResolverFunc) ([]ResolvedAddress, error) {
824+
bogus := NodeInfo{ID: "bogus-id", Name: "ghost"}
825+
return []ResolvedAddress{
826+
{Node: nodes[0], URL: nodes[0].URL},
827+
{Node: bogus, URL: nodes[0].URL},
828+
}, nil
829+
},
830+
wantNodes: 1,
831+
},
832+
}
785833

786-
_, err = tp.getNodesInfo(ctx)
787-
require.ErrorIs(t, err, context.Canceled)
788-
})
834+
for _, tt := range runnerEdgeCases {
835+
t.Run(tt.name, func(t *testing.T) {
836+
t.Parallel()
789837

790-
t.Run("runner drops all returns ErrAllResolversFailed", func(t *testing.T) {
791-
t.Parallel()
838+
ctx := t.Context()
839+
if tt.makeCtx != nil {
840+
ctx = tt.makeCtx(t)
841+
}
792842

793-
tp, err := New(Config{
794-
URLs: []*url.URL{testSeedURL(t)},
795-
Transport: newResolverTestTransport(t, nodesJSON),
796-
HealthCheck: NoOpHealthCheck,
797-
AddressResolverRunner: func(_ context.Context, _ []NodeInfo, _ AddressResolverFunc) ([]ResolvedAddress, error) {
798-
return nil, nil
799-
},
800-
})
801-
require.NoError(t, err)
843+
tp, err := New(Config{
844+
URLs: []*url.URL{testSeedURL(t)},
845+
Transport: newResolverTestTransport(t, nodesJSON),
846+
HealthCheck: NoOpHealthCheck,
847+
AddressResolverRunner: tt.runner,
848+
})
849+
require.NoError(t, err)
802850

803-
_, err = tp.getNodesInfo(t.Context())
804-
require.ErrorIs(t, err, ErrAllResolversFailed)
805-
})
851+
nodes, err := tp.getNodesInfo(ctx)
852+
if tt.wantErr != nil {
853+
require.ErrorIs(t, err, tt.wantErr)
854+
return
855+
}
856+
require.NoError(t, err)
857+
require.Len(t, nodes, tt.wantNodes)
858+
})
859+
}
806860

807861
t.Run("runner error propagated", func(t *testing.T) {
808862
t.Parallel()

opensearchtransport/discovery.go

Lines changed: 49 additions & 48 deletions
Original file line numberDiff line numberDiff line change
@@ -1333,30 +1333,7 @@ func (c *Client) resolveDiscoveredNodes(ctx context.Context, pending []discovery
13331333
}
13341334

13351335
node := p.node
1336-
u := p.defaultURL
1337-
1338-
if resolved != nil && resolved.String() != u.String() {
1339-
node.rewritten = true
1340-
1341-
if c.metrics != nil {
1342-
c.metrics.addressResolverRewrites.Add(1)
1343-
}
1344-
1345-
if obs != nil {
1346-
obs.OnAddressRewrite(AddressRewriteEvent{
1347-
ID: node.ID,
1348-
Name: node.Name,
1349-
Roles: node.Roles,
1350-
OriginalURL: u.String(),
1351-
RewrittenURL: resolved.String(),
1352-
Timestamp: time.Now().UTC(),
1353-
})
1354-
}
1355-
1356-
u = resolved
1357-
}
1358-
1359-
node.url = u
1336+
node.url = c.applyRewrite(&node, p.defaultURL, resolved, obs)
13601337
results[i] = resolvedNode{node: node}
13611338
}(i, p)
13621339
}
@@ -1389,6 +1366,53 @@ func (c *Client) resolveDiscoveredNodes(ctx context.Context, pending []discovery
13891366
return out, nil
13901367
}
13911368

1369+
// applyRewrite checks whether resolved differs from defaultURL and, when it
1370+
// does, marks the node as rewritten, increments the rewrite metric, and fires
1371+
// the OnAddressRewrite observer event. Returns the URL to use for the node.
1372+
//
1373+
// A non-nil resolved URL with an empty Scheme or Host is treated as a
1374+
// resolver error: the default URL is kept, the error counter is incremented,
1375+
// and a debug log is emitted. Validating here surfaces the misuse at
1376+
// resolver time rather than letting a malformed URL flow into the connection
1377+
// pool and fail later as a confusing HTTP error.
1378+
func (c *Client) applyRewrite(node *nodeInfo, defaultURL, resolved *url.URL, obs ConnectionObserver) *url.URL {
1379+
if resolved == nil || resolved.String() == defaultURL.String() {
1380+
return defaultURL
1381+
}
1382+
1383+
if resolved.Scheme == "" || resolved.Host == "" {
1384+
if c.metrics != nil {
1385+
c.metrics.addressResolverErrors.Add(1)
1386+
}
1387+
if dl := loadDebugLogger(); dl != nil {
1388+
dl.Logf("AddressResolver returned malformed URL for node %q (scheme=%q host=%q); keeping default %q\n",
1389+
node.Name, resolved.Scheme, resolved.Host, defaultURL)
1390+
}
1391+
return defaultURL
1392+
}
1393+
1394+
node.rewritten = true
1395+
1396+
if c.metrics != nil {
1397+
c.metrics.addressResolverRewrites.Add(1)
1398+
}
1399+
1400+
if obs != nil {
1401+
// Clone Roles so an observer that retains the event sees a stable
1402+
// copy -- matches ConnectionEvent's c.Roles.toSlice() behavior.
1403+
obs.OnAddressRewrite(AddressRewriteEvent{
1404+
ID: node.ID,
1405+
Name: node.Name,
1406+
Roles: slices.Clone(node.Roles),
1407+
OriginalURL: defaultURL.String(),
1408+
RewrittenURL: resolved.String(),
1409+
Timestamp: time.Now().UTC(),
1410+
})
1411+
}
1412+
1413+
return resolved
1414+
}
1415+
13921416
// newInstrumentedResolver wraps an AddressResolverFunc with metrics
13931417
// instrumentation. Each invocation increments addressResolverCalls, and
13941418
// non-nil errors increment addressResolverErrors and emit a debug log.
@@ -1474,30 +1498,7 @@ func (c *Client) runAddressResolverRunner(ctx context.Context, pending []discove
14741498
seen[ra.Node.ID] = struct{}{}
14751499

14761500
node := p.node
1477-
u := p.defaultURL
1478-
1479-
if ra.URL.String() != u.String() {
1480-
node.rewritten = true
1481-
1482-
if c.metrics != nil {
1483-
c.metrics.addressResolverRewrites.Add(1)
1484-
}
1485-
1486-
if obs != nil {
1487-
obs.OnAddressRewrite(AddressRewriteEvent{
1488-
ID: node.ID,
1489-
Name: node.Name,
1490-
Roles: node.Roles,
1491-
OriginalURL: u.String(),
1492-
RewrittenURL: ra.URL.String(),
1493-
Timestamp: time.Now().UTC(),
1494-
})
1495-
}
1496-
1497-
u = ra.URL
1498-
}
1499-
1500-
node.url = u
1501+
node.url = c.applyRewrite(&node, p.defaultURL, ra.URL, obs)
15011502
out = append(out, node)
15021503
}
15031504

0 commit comments

Comments
 (0)