Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 9 additions & 0 deletions client/inner_client.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down
81 changes: 52 additions & 29 deletions client/resource_manager_client.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -352,47 +358,64 @@ 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 {
case <-dispatcherCtx.Done():
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))
}
Expand Down
180 changes: 179 additions & 1 deletion client/resource_manager_client_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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)
}
Expand Down Expand Up @@ -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) {
Expand Down Expand Up @@ -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())),
}
Expand All @@ -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{})
Expand Down Expand Up @@ -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())
}
Loading