Skip to content

Commit dbd9788

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 f1b5c57 commit dbd9788

2 files changed

Lines changed: 105 additions & 81 deletions

File tree

opensearchtransport/address_resolver_internal_test.go

Lines changed: 75 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -389,7 +389,7 @@ func TestAddressResolver(t *testing.T) {
389389
ctx, cancel := context.WithCancel(t.Context())
390390
cancel()
391391
return ctx, func(_ context.Context, _ NodeInfo) (*url.URL, error) {
392-
return nil, nil //nolint:nilnil // testing (nil, nil) protocol case
392+
return nil, nil
393393
}
394394
},
395395
wantErr: context.Canceled,
@@ -399,12 +399,13 @@ func TestAddressResolver(t *testing.T) {
399399
setup: func(t *testing.T) (context.Context, AddressResolverFunc) {
400400
t.Helper()
401401
ctx, cancel := context.WithCancel(t.Context())
402+
t.Cleanup(cancel)
402403
var calls atomic.Int32
403404
return ctx, func(_ context.Context, _ NodeInfo) (*url.URL, error) {
404405
if calls.Add(1) == 1 {
405406
cancel()
406407
}
407-
return nil, nil //nolint:nilnil // testing (nil, nil) protocol case
408+
return nil, nil
408409
}
409410
},
410411
maxResolvers: 1,
@@ -433,7 +434,7 @@ func TestAddressResolver(t *testing.T) {
433434
return
434435
}
435436
require.NoError(t, err)
436-
require.Equal(t, tt.wantNodes, len(nodes))
437+
require.Len(t, nodes, tt.wantNodes)
437438
})
438439
}
439440

@@ -635,42 +636,83 @@ func TestAddressResolverRunner(t *testing.T) {
635636
"resolve param should be nil when AddressResolver is not set")
636637
})
637638

638-
t.Run("runner context cancellation propagated", func(t *testing.T) {
639-
t.Parallel()
640-
641-
ctx, cancel := context.WithCancel(t.Context())
642-
cancel()
643-
644-
tp, err := New(Config{
645-
URLs: []*url.URL{testSeedURL(t)},
646-
Transport: newResolverTestTransport(t, nodesJSON),
647-
HealthCheck: NoOpHealthCheck,
648-
AddressResolverRunner: func(ctx context.Context, _ []NodeInfo, _ AddressResolverFunc) ([]ResolvedAddress, error) {
639+
runnerEdgeCases := []struct {
640+
name string
641+
makeCtx func(t *testing.T) context.Context
642+
runner AddressResolverRunnerFunc
643+
wantErr error
644+
wantNodes int
645+
}{
646+
{
647+
name: "context cancellation propagated",
648+
makeCtx: func(t *testing.T) context.Context {
649+
t.Helper()
650+
ctx, cancel := context.WithCancel(t.Context())
651+
cancel()
652+
return ctx
653+
},
654+
runner: func(ctx context.Context, _ []NodeInfo, _ AddressResolverFunc) ([]ResolvedAddress, error) {
649655
return nil, ctx.Err()
650656
},
651-
})
652-
require.NoError(t, err)
657+
wantErr: context.Canceled,
658+
},
659+
{
660+
name: "drops all returns ErrAllResolversFailed",
661+
runner: func(_ context.Context, _ []NodeInfo, _ AddressResolverFunc) ([]ResolvedAddress, error) {
662+
return nil, nil
663+
},
664+
wantErr: ErrAllResolversFailed,
665+
},
666+
{
667+
name: "nil URLs are skipped",
668+
runner: func(_ context.Context, nodes []NodeInfo, _ AddressResolverFunc) ([]ResolvedAddress, error) {
669+
out := make([]ResolvedAddress, len(nodes))
670+
for i, n := range nodes {
671+
out[i] = ResolvedAddress{Node: n, URL: nil}
672+
}
673+
return out, nil
674+
},
675+
wantErr: ErrAllResolversFailed,
676+
},
677+
{
678+
name: "unknown node ID is ignored",
679+
runner: func(_ context.Context, nodes []NodeInfo, _ AddressResolverFunc) ([]ResolvedAddress, error) {
680+
bogus := NodeInfo{ID: "bogus-id", Name: "ghost"}
681+
return []ResolvedAddress{
682+
{Node: nodes[0], URL: nodes[0].URL},
683+
{Node: bogus, URL: nodes[0].URL},
684+
}, nil
685+
},
686+
wantNodes: 1,
687+
},
688+
}
653689

654-
_, err = tp.getNodesInfo(ctx)
655-
require.ErrorIs(t, err, context.Canceled)
656-
})
690+
for _, tt := range runnerEdgeCases {
691+
t.Run(tt.name, func(t *testing.T) {
692+
t.Parallel()
657693

658-
t.Run("runner drops all returns ErrAllResolversFailed", func(t *testing.T) {
659-
t.Parallel()
694+
ctx := t.Context()
695+
if tt.makeCtx != nil {
696+
ctx = tt.makeCtx(t)
697+
}
660698

661-
tp, err := New(Config{
662-
URLs: []*url.URL{testSeedURL(t)},
663-
Transport: newResolverTestTransport(t, nodesJSON),
664-
HealthCheck: NoOpHealthCheck,
665-
AddressResolverRunner: func(_ context.Context, _ []NodeInfo, _ AddressResolverFunc) ([]ResolvedAddress, error) {
666-
return nil, nil
667-
},
668-
})
669-
require.NoError(t, err)
699+
tp, err := New(Config{
700+
URLs: []*url.URL{testSeedURL(t)},
701+
Transport: newResolverTestTransport(t, nodesJSON),
702+
HealthCheck: NoOpHealthCheck,
703+
AddressResolverRunner: tt.runner,
704+
})
705+
require.NoError(t, err)
670706

671-
_, err = tp.getNodesInfo(t.Context())
672-
require.ErrorIs(t, err, ErrAllResolversFailed)
673-
})
707+
nodes, err := tp.getNodesInfo(ctx)
708+
if tt.wantErr != nil {
709+
require.ErrorIs(t, err, tt.wantErr)
710+
return
711+
}
712+
require.NoError(t, err)
713+
require.Len(t, nodes, tt.wantNodes)
714+
})
715+
}
674716

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

opensearchtransport/discovery.go

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

13251325
node := p.node
1326-
u := p.defaultURL
1327-
1328-
if resolved != nil && resolved.String() != u.String() {
1329-
node.rewritten = true
1330-
1331-
if c.metrics != nil {
1332-
c.metrics.addressResolverRewrites.Add(1)
1333-
}
1334-
1335-
if obs != nil {
1336-
obs.OnAddressRewrite(AddressRewriteEvent{
1337-
ID: node.ID,
1338-
Name: node.Name,
1339-
Roles: node.Roles,
1340-
OriginalURL: u.String(),
1341-
RewrittenURL: resolved.String(),
1342-
Timestamp: time.Now().UTC(),
1343-
})
1344-
}
1345-
1346-
u = resolved
1347-
}
1348-
1349-
node.url = u
1326+
node.url = c.applyRewrite(&node, p.defaultURL, resolved, obs)
13501327
results[i] = resolvedNode{node: node}
13511328
}(i, p)
13521329
}
@@ -1379,6 +1356,34 @@ func (c *Client) resolveDiscoveredNodes(ctx context.Context, pending []discovery
13791356
return out, nil
13801357
}
13811358

1359+
// applyRewrite checks whether resolved differs from defaultURL and, when it
1360+
// does, marks the node as rewritten, increments the rewrite metric, and fires
1361+
// the OnAddressRewrite observer event. Returns the URL to use for the node.
1362+
func (c *Client) applyRewrite(node *nodeInfo, defaultURL, resolved *url.URL, obs ConnectionObserver) *url.URL {
1363+
if resolved == nil || resolved.String() == defaultURL.String() {
1364+
return defaultURL
1365+
}
1366+
1367+
node.rewritten = true
1368+
1369+
if c.metrics != nil {
1370+
c.metrics.addressResolverRewrites.Add(1)
1371+
}
1372+
1373+
if obs != nil {
1374+
obs.OnAddressRewrite(AddressRewriteEvent{
1375+
ID: node.ID,
1376+
Name: node.Name,
1377+
Roles: node.Roles,
1378+
OriginalURL: defaultURL.String(),
1379+
RewrittenURL: resolved.String(),
1380+
Timestamp: time.Now().UTC(),
1381+
})
1382+
}
1383+
1384+
return resolved
1385+
}
1386+
13821387
// newInstrumentedResolver wraps an AddressResolverFunc with metrics
13831388
// instrumentation. Each invocation increments addressResolverCalls, and
13841389
// non-nil errors increment addressResolverErrors and emit a debug log.
@@ -1449,30 +1454,7 @@ func (c *Client) runAddressResolverRunner(ctx context.Context, pending []discove
14491454
}
14501455

14511456
node := p.node
1452-
u := p.defaultURL
1453-
1454-
if ra.URL.String() != u.String() {
1455-
node.rewritten = true
1456-
1457-
if c.metrics != nil {
1458-
c.metrics.addressResolverRewrites.Add(1)
1459-
}
1460-
1461-
if obs != nil {
1462-
obs.OnAddressRewrite(AddressRewriteEvent{
1463-
ID: node.ID,
1464-
Name: node.Name,
1465-
Roles: node.Roles,
1466-
OriginalURL: u.String(),
1467-
RewrittenURL: ra.URL.String(),
1468-
Timestamp: time.Now().UTC(),
1469-
})
1470-
}
1471-
1472-
u = ra.URL
1473-
}
1474-
1475-
node.url = u
1457+
node.url = c.applyRewrite(&node, p.defaultURL, ra.URL, obs)
14761458
out = append(out, node)
14771459
}
14781460

0 commit comments

Comments
 (0)