Skip to content

Commit 376bcff

Browse files
pingsutwCopilot
andauthored
fix(connector): reuse persistent connections for ListConnectors polling (#7705)
Signed-off-by: Kevin Su <pingsutw@apache.org> Signed-off-by: Kevin Su <pingsutw@gmail.com> Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com>
1 parent aa60b1b commit 376bcff

3 files changed

Lines changed: 132 additions & 26 deletions

File tree

flyteplugins/go/tasks/plugins/webapi/connector/client.go

Lines changed: 79 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -4,13 +4,16 @@ import (
44
"context"
55
"crypto/x509"
66
"strings"
7+
"sync"
8+
"time"
79

810
"golang.org/x/exp/maps"
911
"google.golang.org/grpc"
1012
"google.golang.org/grpc/codes"
1113
"google.golang.org/grpc/credentials"
1214
"google.golang.org/grpc/credentials/insecure"
1315
"google.golang.org/grpc/grpclog"
16+
"google.golang.org/grpc/keepalive"
1417
"google.golang.org/grpc/status"
1518

1619
"github.com/flyteorg/flyte/v2/flytestdlib/config"
@@ -32,10 +35,66 @@ type Connector struct {
3235

3336
// ClientSet contains the clients exposed to communicate with various connector services.
3437
type ClientSet struct {
38+
mu sync.RWMutex
3539
asyncConnectorClients map[string]connector.AsyncConnectorServiceClient // map[endpoint] => AsyncConnectorServiceClient
3640
connectorMetadataClients map[string]connector.ConnectorMetadataServiceClient // map[endpoint] => ConnectorMetadataServiceClient
3741
}
3842

43+
// getOrDialAsyncClient returns the cached AsyncConnectorService client for the
44+
// endpoint, dialing (and caching) a persistent connection on first use.
45+
func (cs *ClientSet) getOrDialAsyncClient(ctx context.Context, deployment *Deployment) (connector.AsyncConnectorServiceClient, error) {
46+
cs.mu.Lock()
47+
defer cs.mu.Unlock()
48+
49+
if cs.asyncConnectorClients == nil {
50+
cs.asyncConnectorClients = make(map[string]connector.AsyncConnectorServiceClient)
51+
}
52+
if cs.connectorMetadataClients == nil {
53+
cs.connectorMetadataClients = make(map[string]connector.ConnectorMetadataServiceClient)
54+
}
55+
56+
if client, ok := cs.asyncConnectorClients[deployment.Endpoint]; ok {
57+
return client, nil
58+
}
59+
60+
conn, err := getGrpcConnection(ctx, deployment)
61+
if err != nil {
62+
return nil, err
63+
}
64+
65+
asyncClient := connector.NewAsyncConnectorServiceClient(conn)
66+
cs.asyncConnectorClients[deployment.Endpoint] = asyncClient
67+
if _, ok := cs.connectorMetadataClients[deployment.Endpoint]; !ok {
68+
cs.connectorMetadataClients[deployment.Endpoint] = connector.NewConnectorMetadataServiceClient(conn)
69+
}
70+
return asyncClient, nil
71+
}
72+
73+
// getOrDialMetadataClient returns the cached ConnectorMetadataService client for
74+
// the endpoint, dialing (and caching) a persistent connection on first use.
75+
func (cs *ClientSet) getOrDialMetadataClient(ctx context.Context, deployment *Deployment) (connector.ConnectorMetadataServiceClient, error) {
76+
cs.mu.Lock()
77+
defer cs.mu.Unlock()
78+
if client, ok := cs.connectorMetadataClients[deployment.Endpoint]; ok {
79+
return client, nil
80+
}
81+
conn, err := getGrpcConnection(ctx, deployment)
82+
if err != nil {
83+
return nil, err
84+
}
85+
client := connector.NewConnectorMetadataServiceClient(conn)
86+
cs.connectorMetadataClients[deployment.Endpoint] = client
87+
return client, nil
88+
}
89+
90+
// metadataClient returns the cached ConnectorMetadataService client for the endpoint.
91+
func (cs *ClientSet) metadataClient(endpoint string) (connector.ConnectorMetadataServiceClient, bool) {
92+
cs.mu.RLock()
93+
defer cs.mu.RUnlock()
94+
client, ok := cs.connectorMetadataClients[endpoint]
95+
return client, ok
96+
}
97+
3998
func getGrpcConnection(ctx context.Context, connector *Deployment) (*grpc.ClientConn, error) {
4099
var opts []grpc.DialOption
41100

@@ -59,6 +118,13 @@ func getGrpcConnection(ctx context.Context, connector *Deployment) (*grpc.Client
59118
opts = append(opts,
60119
grpc.WithChainUnaryInterceptor(clientMetrics.UnaryClientInterceptor()),
61120
grpc.WithChainStreamInterceptor(clientMetrics.StreamClientInterceptor()),
121+
// Keepalive lets gRPC notice a dead/half-open transport within seconds
122+
// (via HTTP/2 PINGs during active RPCs) instead of relying on the OS TCP
123+
// timeout.
124+
grpc.WithKeepaliveParams(keepalive.ClientParameters{
125+
Time: 30 * time.Second,
126+
Timeout: 10 * time.Second,
127+
}),
62128
)
63129

64130
var err error
@@ -102,7 +168,7 @@ func updateRegistry(
102168
isConnectorApp bool,
103169
) {
104170
for connectorID, connectorDeployment := range connectorDeployments {
105-
client, ok := cs.connectorMetadataClients[connectorDeployment.Endpoint]
171+
client, ok := cs.metadataClient(connectorDeployment.Endpoint)
106172
if !ok {
107173
logger.Warningf(ctx, "Connector client not found in the clientSet for the endpoint: %v", connectorDeployment.Endpoint)
108174
continue
@@ -199,28 +265,29 @@ func getConnectorRegistry(ctx context.Context, cs *ClientSet) Registry {
199265
return newConnectorRegistry
200266
}
201267

202-
func getConnectorClientSets(ctx context.Context) *ClientSet {
203-
clientSet := &ClientSet{
204-
asyncConnectorClients: make(map[string]connector.AsyncConnectorServiceClient),
205-
connectorMetadataClients: make(map[string]connector.ConnectorMetadataServiceClient),
206-
}
207-
268+
// allConnectorDeployments returns every configured connector endpoint: the
269+
// default connector, explicit deployments, and connector apps.
270+
func allConnectorDeployments(cfg *Config) []*Deployment {
208271
var connectorDeployments []*Deployment
209-
cfg := GetConfig()
210-
211272
if len(cfg.DefaultConnector.Endpoint) != 0 {
212273
connectorDeployments = append(connectorDeployments, &cfg.DefaultConnector)
213274
}
214-
215275
for _, deployment := range cfg.ConnectorDeployments {
216276
connectorDeployments = append(connectorDeployments, deployment)
217277
}
218-
219278
for _, deployment := range cfg.ConnectorApps {
220279
connectorDeployments = append(connectorDeployments, deployment)
221280
}
281+
return connectorDeployments
282+
}
283+
284+
func getConnectorClientSets(ctx context.Context) *ClientSet {
285+
clientSet := &ClientSet{
286+
asyncConnectorClients: make(map[string]connector.AsyncConnectorServiceClient),
287+
connectorMetadataClients: make(map[string]connector.ConnectorMetadataServiceClient),
288+
}
222289

223-
for _, connectorDeployment := range connectorDeployments {
290+
for _, connectorDeployment := range allConnectorDeployments(GetConfig()) {
224291
if _, ok := clientSet.connectorMetadataClients[connectorDeployment.Endpoint]; ok {
225292
logger.Infof(ctx, "Connector client already initialized for [%v]", connectorDeployment.Endpoint)
226293
continue

flyteplugins/go/tasks/plugins/webapi/connector/client_test.go

Lines changed: 46 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,10 @@ import (
55
"testing"
66

77
"github.com/stretchr/testify/assert"
8+
"google.golang.org/grpc"
9+
"google.golang.org/grpc/credentials/insecure"
10+
11+
connectorpb "github.com/flyteorg/flyte/v2/gen/go/flyteidl2/connector"
812
)
913

1014
func TestInitializeClients(t *testing.T) {
@@ -26,3 +30,45 @@ func TestInitializeClients(t *testing.T) {
2630
_, ok = cs.asyncConnectorClients["x"]
2731
assert.True(t, ok)
2832
}
33+
34+
func TestAllConnectorDeployments(t *testing.T) {
35+
cfg := defaultConfig
36+
cfg.DefaultConnector = Deployment{Endpoint: "default"}
37+
cfg.ConnectorDeployments = map[string]*Deployment{"x": {Endpoint: "x"}}
38+
cfg.ConnectorApps = map[string]*Deployment{"app": {Endpoint: "app"}}
39+
40+
endpoints := make([]string, 0)
41+
for _, d := range allConnectorDeployments(&cfg) {
42+
endpoints = append(endpoints, d.Endpoint)
43+
}
44+
assert.ElementsMatch(t, []string{"default", "x", "app"}, endpoints)
45+
}
46+
47+
func TestGetOrDialMetadataClientReusesConnection(t *testing.T) {
48+
// Cancelable context so getGrpcConnection's close goroutine is cleaned up
49+
// when the test finishes.
50+
ctx, cancel := context.WithCancel(context.Background())
51+
defer cancel()
52+
53+
// A pre-existing cached client for endpoint "ep".
54+
conn, err := grpc.NewClient("passthrough:///ep", grpc.WithTransportCredentials(insecure.NewCredentials()))
55+
assert.NoError(t, err)
56+
t.Cleanup(func() { _ = conn.Close() })
57+
existing := connectorpb.NewConnectorMetadataServiceClient(conn)
58+
59+
cs := &ClientSet{
60+
asyncConnectorClients: map[string]connectorpb.AsyncConnectorServiceClient{},
61+
connectorMetadataClients: map[string]connectorpb.ConnectorMetadataServiceClient{"ep": existing},
62+
}
63+
64+
// Already cached -> must reuse the same client, not re-dial/replace it.
65+
got, err := cs.getOrDialMetadataClient(ctx, &Deployment{Endpoint: "ep"})
66+
assert.NoError(t, err)
67+
assert.True(t, existing == got, "cached connection should be reused, not re-dialed")
68+
69+
// New endpoint -> dialed and cached. The passthrough scheme avoids DNS.
70+
_, err = cs.getOrDialMetadataClient(ctx, &Deployment{Endpoint: "passthrough:///new", Insecure: true})
71+
assert.NoError(t, err)
72+
_, ok := cs.metadataClient("passthrough:///new")
73+
assert.True(t, ok, "new endpoint should be dialed and cached")
74+
}

flyteplugins/go/tasks/plugins/webapi/connector/plugin.go

Lines changed: 7 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -374,24 +374,17 @@ func (p *Plugin) Status(ctx context.Context, taskCtx webapi.StatusContext) (phas
374374
}
375375

376376
func (p *Plugin) getAsyncConnectorClient(ctx context.Context, connector *Deployment) (connectorPb.AsyncConnectorServiceClient, error) {
377-
client, ok := p.cs.asyncConnectorClients[connector.Endpoint]
378-
if !ok {
379-
conn, err := getGrpcConnection(ctx, connector)
380-
if err != nil {
381-
return nil, err
382-
}
383-
client = connectorPb.NewAsyncConnectorServiceClient(conn)
384-
p.cs.asyncConnectorClients[connector.Endpoint] = client
385-
}
386-
return client, nil
377+
return p.cs.getOrDialAsyncClient(ctx, connector)
387378
}
388379

389380
func (p *Plugin) watchConnectors(ctx context.Context, connectorService *ConnectorService) {
390381
go wait.Until(func() {
391-
childCtx, cancel := context.WithCancel(ctx)
392-
defer cancel()
393-
clientSet := getConnectorClientSets(childCtx)
394-
connectorRegistry := getConnectorRegistry(childCtx, clientSet)
382+
for _, deployment := range allConnectorDeployments(GetConfig()) {
383+
if _, err := p.cs.getOrDialMetadataClient(ctx, deployment); err != nil {
384+
logger.Errorf(ctx, "failed to connect to connector [%v]: %v", deployment.Endpoint, err)
385+
}
386+
}
387+
connectorRegistry := getConnectorRegistry(ctx, p.cs)
395388
p.setRegistry(connectorRegistry)
396389
connectorService.SetSupportedTaskType(connectorRegistry.getSupportedTaskTypes())
397390
}, p.cfg.PollInterval.Duration, ctx.Done())

0 commit comments

Comments
 (0)