diff --git a/client/inner_client.go b/client/inner_client.go index 3c7d2f5f92..773b62de09 100644 --- a/client/inner_client.go +++ b/client/inner_client.go @@ -55,6 +55,8 @@ type innerClient struct { // For internal usage. updateTokenConnectionCh chan struct{} + tokenConnectionMu sync.Mutex + tokenConnectionCancel context.CancelFunc ctx context.Context cancel context.CancelFunc @@ -231,6 +233,13 @@ func (c *innerClient) getResourceManagerDiscovery() *sd.ResourceManagerDiscovery } func (c *innerClient) scheduleUpdateTokenConnection(string) error { + // Interrupt an in-flight request so the dispatcher can reconnect to the + // newly discovered endpoint instead of waiting on the stale stream. + c.tokenConnectionMu.Lock() + defer c.tokenConnectionMu.Unlock() + if c.tokenConnectionCancel != nil { + c.tokenConnectionCancel() + } select { case c.updateTokenConnectionCh <- struct{}{}: default: diff --git a/client/resource_manager_client.go b/client/resource_manager_client.go index 58e7f87d29..56c6551818 100644 --- a/client/resource_manager_client.go +++ b/client/resource_manager_client.go @@ -312,6 +312,12 @@ type tokenDispatcher struct { tokenBatchController *tokenBatchController } +func (c *innerClient) setTokenConnectionCancel(cancel context.CancelFunc) { + c.tokenConnectionMu.Lock() + c.tokenConnectionCancel = cancel + c.tokenConnectionMu.Unlock() +} + type resourceManagerConnectionContext struct { stream rmpb.ResourceManager_AcquireTokenBucketsClient ctx context.Context @@ -352,7 +358,10 @@ func (c *innerClient) handleResourceTokenDispatcher(dispatcherCtx context.Contex ) if err = c.tryResourceManagerConnect(dispatcherCtx, &connection); err != nil { log.Warn("[resource_manager] get token stream error", zap.Error(err)) + } else { + c.setTokenConnectionCancel(connection.cancel) } +tokenRequestLoop: for { // Fetch the request from the channel. select { @@ -360,39 +369,53 @@ func (c *innerClient) handleResourceTokenDispatcher(dispatcherCtx context.Contex return case firstRequest = <-tbc.tokenRequestCh: } - // Try to get a stream connection. - stream, streamCtx = connection.stream, connection.ctx - select { - case <-c.updateTokenConnectionCh: - toReconnect = true - default: - toReconnect = stream == nil - } - // If the stream is nil or the leader has changed, try to reconnect. - if toReconnect { - connection.reset() - if err := c.tryResourceManagerConnect(dispatcherCtx, &connection); err != nil { - log.Error("[resource_manager] try to connect token leader failed", errs.ZapError(err)) - } - log.Info("[resource_manager] token leader may change, try to reconnect the stream") + for { + // Try to get a stream connection. stream, streamCtx = connection.stream, connection.ctx - } - // If the stream is still nil, return an error. - if stream == nil { - firstRequest.done <- errors.Errorf("failed to get the stream connection") - c.serviceDiscovery.ScheduleCheckMemberChanged() - connection.reset() - continue - } - select { - case <-streamCtx.Done(): - connection.reset() - log.Info("[resource_manager] token stream is canceled") - continue - default: + select { + case <-c.updateTokenConnectionCh: + toReconnect = true + default: + toReconnect = stream == nil + } + // If the stream is nil or the leader has changed, try to reconnect. + if toReconnect { + c.setTokenConnectionCancel(nil) + connection.reset() + if err := c.tryResourceManagerConnect(dispatcherCtx, &connection); err != nil { + log.Error("[resource_manager] try to connect token leader failed", errs.ZapError(err)) + } else { + c.setTokenConnectionCancel(connection.cancel) + } + log.Info("[resource_manager] token leader may change, try to reconnect the stream") + stream, streamCtx = connection.stream, connection.ctx + if stream != nil { + continue + } + } + // If the stream is still nil, return an error. + if stream == nil { + firstRequest.done <- errors.Errorf("failed to get the stream connection") + c.serviceDiscovery.ScheduleCheckMemberChanged() + connection.reset() + continue tokenRequestLoop + } + select { + case <-streamCtx.Done(): + c.setTokenConnectionCancel(nil) + connection.reset() + log.Info("[resource_manager] token stream is canceled") + if dispatcherCtx.Err() != nil { + return + } + continue + default: + } + break } if err = c.processTokenRequests(stream, firstRequest); err != nil { c.serviceDiscovery.ScheduleCheckMemberChanged() + c.setTokenConnectionCancel(nil) connection.reset() log.Info("[resource_manager] token request error", zap.Error(err)) } diff --git a/client/resource_manager_client_test.go b/client/resource_manager_client_test.go index 8ddcd32a8f..98683fb062 100644 --- a/client/resource_manager_client_test.go +++ b/client/resource_manager_client_test.go @@ -44,6 +44,9 @@ type testServiceDiscovery struct { servingURL string keyspaceID uint32 clientConns sync.Map + + getOrCreateHookMu sync.RWMutex + getOrCreateHook func() } func newTestServiceDiscovery(servingURL string, conn *grpc.ClientConn) *testServiceDiscovery { @@ -75,12 +78,23 @@ func (*testServiceDiscovery) GetServiceClient() sd.ServiceClient func (*testServiceDiscovery) GetServiceClientByKind(sd.APIKind) sd.ServiceClient { return nil } func (*testServiceDiscovery) GetAllServiceClients() []sd.ServiceClient { return nil } func (t *testServiceDiscovery) GetOrCreateGRPCConn(url string) (*grpc.ClientConn, error) { + t.getOrCreateHookMu.RLock() + hook := t.getOrCreateHook + t.getOrCreateHookMu.RUnlock() + if hook != nil { + hook() + } conn, ok := t.clientConns.Load(url) if !ok { return nil, errors.New("unexpected URL") } return conn.(*grpc.ClientConn), nil } +func (t *testServiceDiscovery) setGetOrCreateHook(hook func()) { + t.getOrCreateHookMu.Lock() + t.getOrCreateHook = hook + t.getOrCreateHookMu.Unlock() +} func (t *testServiceDiscovery) RemoveClientConn(url string) { t.clientConns.Delete(url) } @@ -134,6 +148,9 @@ type testRMServer struct { deleteCount atomic.Int32 tokenCount atomic.Int32 getErr error + + blockTokenResponse bool + tokenRequestReceived chan struct{} } func (s *testRMServer) ListResourceGroups(context.Context, *rmpb.ListResourceGroupsRequest) (*rmpb.ListResourceGroupsResponse, error) { @@ -182,6 +199,16 @@ func (s *testRMServer) AcquireTokenBuckets(stream rmpb.ResourceManager_AcquireTo return err } s.tokenCount.Add(1) + if s.tokenRequestReceived != nil { + select { + case s.tokenRequestReceived <- struct{}{}: + default: + } + } + if s.blockTokenResponse { + <-stream.Context().Done() + return stream.Context().Err() + } resp := &rmpb.TokenBucketsResponse{ Responses: make([]*rmpb.TokenBucketResponse, 0, len(req.GetRequests())), } @@ -197,13 +224,16 @@ func (s *testRMServer) AcquireTokenBuckets(stream rmpb.ResourceManager_AcquireTo } } -func startTestRMServer(t *testing.T, id string) (string, *testRMServer, func()) { +func startTestRMServer(t *testing.T, id string, opts ...func(*testRMServer)) (string, *testRMServer, func()) { t.Helper() listener, err := net.Listen("tcp", "127.0.0.1:0") require.NoError(t, err) server := grpc.NewServer() rmServer := &testRMServer{id: id} + for _, opt := range opts { + opt(rmServer) + } rmpb.RegisterResourceManagerServer(server, rmServer) done := make(chan struct{}) @@ -408,3 +438,151 @@ func TestTryResourceManagerConnectUsesRMForTokenAndFallbackToPD(t *testing.T) { require.EqualValues(t, 1, pdServer.tokenCount.Load()) }) } + +func TestTokenDispatcherReconnectsWhenRMEndpointChanges(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + pdRequestReceived := make(chan struct{}, 1) + pdAddr, pdServer, pdCleanup := startTestRMServer(t, "pd", func(server *testRMServer) { + server.blockTokenResponse = true + server.tokenRequestReceived = pdRequestReceived + }) + t.Cleanup(pdCleanup) + rmAddr, rmServer, rmCleanup := startTestRMServer(t, "rm") + t.Cleanup(rmCleanup) + + inner := newInnerClientForRMRouteTest(t, ctx, pdAddr) + inner.createTokenDispatcher() + t.Cleanup(func() { + inner.tokenDispatcher.dispatcherCancel() + inner.wg.Wait() + }) + cli := &client{inner: inner} + + firstRequestDone := make(chan error, 1) + go func() { + _, err := cli.AcquireTokenBuckets(ctx, &rmpb.TokenBucketsRequest{ + Requests: []*rmpb.TokenBucketRequest{{ResourceGroupName: "test-group"}}, + }) + firstRequestDone <- err + }() + select { + case <-pdRequestReceived: + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for the token request to reach PD") + } + + discovery := newTestResourceManagerDiscovery(t, ctx, rmAddr) + t.Cleanup(discovery.Close) + inner.Lock() + inner.resourceManagerDiscovery = discovery + inner.Unlock() + require.NoError(t, inner.scheduleUpdateTokenConnection("")) + + select { + case err := <-firstRequestDone: + require.Error(t, err) + require.Equal(t, codes.Canceled, status.Code(err)) + case <-time.After(3 * time.Second): + t.Fatal("stale token request was not canceled after the RM endpoint changed") + } + + requestCtx, requestCancel := context.WithTimeout(ctx, 3*time.Second) + defer requestCancel() + _, err := cli.AcquireTokenBuckets(requestCtx, &rmpb.TokenBucketsRequest{ + Requests: []*rmpb.TokenBucketRequest{{ResourceGroupName: "test-group"}}, + }) + require.NoError(t, err) + require.EqualValues(t, 1, pdServer.tokenCount.Load()) + require.EqualValues(t, 1, rmServer.tokenCount.Load()) +} + +func TestTokenDispatcherRechecksEndpointUpdatesAfterReconnect(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + pdRequestReceived := make(chan struct{}, 1) + pdAddr, pdServer, pdCleanup := startTestRMServer(t, "pd", func(server *testRMServer) { + server.blockTokenResponse = true + server.tokenRequestReceived = pdRequestReceived + }) + t.Cleanup(pdCleanup) + rmAddr, rmServer, rmCleanup := startTestRMServer(t, "rm") + t.Cleanup(rmCleanup) + discovery := newTestResourceManagerDiscovery(t, ctx, rmAddr) + t.Cleanup(discovery.Close) + + inner := newInnerClientForRMRouteTest(t, ctx, pdAddr) + inner.createTokenDispatcher() + t.Cleanup(func() { + inner.tokenDispatcher.dispatcherCancel() + inner.wg.Wait() + }) + cli := &client{inner: inner} + + firstRequestDone := make(chan error, 1) + go func() { + _, err := cli.AcquireTokenBuckets(ctx, &rmpb.TokenBucketsRequest{ + Requests: []*rmpb.TokenBucketRequest{{ResourceGroupName: "first-request"}}, + }) + firstRequestDone <- err + }() + select { + case <-pdRequestReceived: + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for the first token request to reach PD") + } + + reconnectStarted := make(chan struct{}) + allowReconnect := make(chan struct{}) + var hookOnce, releaseOnce sync.Once + releaseReconnect := func() { + releaseOnce.Do(func() { + close(allowReconnect) + }) + } + t.Cleanup(releaseReconnect) + testDiscovery := inner.serviceDiscovery.(*testServiceDiscovery) + testDiscovery.setGetOrCreateHook(func() { + hookOnce.Do(func() { + close(reconnectStarted) + <-allowReconnect + }) + }) + require.NoError(t, inner.scheduleUpdateTokenConnection("")) + select { + case err := <-firstRequestDone: + require.Error(t, err) + require.Equal(t, codes.Canceled, status.Code(err)) + case <-time.After(3 * time.Second): + t.Fatal("stale token request was not canceled") + } + + secondRequestDone := make(chan error, 1) + go func() { + _, err := cli.AcquireTokenBuckets(ctx, &rmpb.TokenBucketsRequest{ + Requests: []*rmpb.TokenBucketRequest{{ResourceGroupName: "second-request"}}, + }) + secondRequestDone <- err + }() + select { + case <-reconnectStarted: + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for the token dispatcher to start reconnecting") + } + inner.Lock() + inner.resourceManagerDiscovery = discovery + inner.Unlock() + require.NoError(t, inner.scheduleUpdateTokenConnection("")) + releaseReconnect() + + select { + case err := <-secondRequestDone: + require.NoError(t, err) + case <-time.After(3 * time.Second): + t.Fatal("token request was not preserved across reconnects") + } + require.EqualValues(t, 1, pdServer.tokenCount.Load()) + require.EqualValues(t, 1, rmServer.tokenCount.Load()) +}