Skip to content
155 changes: 106 additions & 49 deletions internal/service/coordinator/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -136,10 +136,11 @@ type pinnedStateCoordinator struct {
lastUsed time.Time
}

// client holds the gRPC connection and clients for the coordinator service.
// it should be closed and removed when no longer needed or when the coordinator
// is unhealthy.
// client holds the gRPC connection and clients for one coordinator endpoint.
// The connection is closed when the endpoint is replaced or during cleanup.
type client struct {
address string
startedAt time.Time
conn *grpc.ClientConn
client coordinatorv1.CoordinatorServiceClient
healthClient grpc_health_v1.HealthClient
Expand Down Expand Up @@ -302,22 +303,25 @@ func (cli *clientImpl) attemptCall(ctx context.Context, members []exec.HostInfo,
// Try each coordinator in order (round-robin style)
var lastErr error
for _, member := range members {
if err := ctx.Err(); err != nil {
return err
}

// Get or create client for this coordinator
client, err := cli.getOrCreateClient(member)
client, err := cli.getOrCreateDiscoveredClient(member)
if err != nil {
logger.Warn(ctx, "Failed to connect to coordinator",
slog.String("coordinator-id", member.ID),
tag.Host(member.Host),
tag.Port(member.Port),
tag.Error(err))
cli.removeClient(member) // Remove failed client
cli.recordFailure(err)
lastErr = err
continue
}

// Check if the coordinator is healthy
if err := cli.isHealthy(ctx, member); err != nil {
if err := cli.isHealthy(ctx, client); err != nil {
logger.Warn(ctx, "Failed to check coordinator health",
slog.String("coordinator-id", member.ID),
tag.Host(member.Host),
Expand All @@ -341,7 +345,6 @@ func (cli *clientImpl) attemptCall(ctx context.Context, members []exec.HostInfo,
if errors.Is(err, backoff.ErrPermanent) {
return err
}

lastErr = err
} else {
// Success - record and return immediately
Expand All @@ -365,7 +368,6 @@ func (cli *clientImpl) callPinnedStateCoordinator(ctx context.Context, routingKe
return callback(ctx, member, client)
})
if shouldRefreshPinnedStateCoordinator(err) {
cli.removeClient(member)
cli.refreshPinnedStateCoordinator(ctx, routingKey, member)
}
return err
Expand Down Expand Up @@ -394,6 +396,9 @@ func (cli *clientImpl) pinnedStateCoordinator(ctx context.Context, routingKey st
if err != nil {
return exec.HostInfo{}, err
}
if _, err := cli.getOrCreateDiscoveredClient(member); err != nil {
return exec.HostInfo{}, err
}

cli.stateCoordinatorMu.Lock()
defer cli.stateCoordinatorMu.Unlock()
Expand Down Expand Up @@ -481,6 +486,16 @@ func (cli *clientImpl) refreshPinnedStateCoordinator(ctx context.Context, routin
break
}
}
if replacement != nil {
if _, err := cli.getOrCreateDiscoveredClient(*replacement); err != nil {
logger.Debug(ctx, "Failed to refresh pinned state coordinator client",
slog.String("coordinator-id", replacement.ID),
tag.Host(replacement.Host),
tag.Port(replacement.Port),
tag.Error(err))
return
}
}

cli.stateCoordinatorMu.Lock()
defer cli.stateCoordinatorMu.Unlock()
Expand All @@ -504,9 +519,6 @@ func (cli *clientImpl) callMember(ctx context.Context, member exec.HostInfo, cal
}
if err := callback(ctx, client); err != nil {
cli.recordFailure(err)
if st, ok := status.FromError(err); ok && st.Code() == codes.Unavailable {
cli.removeClient(member)
}
return err
}
cli.recordSuccess(ctx)
Expand All @@ -523,13 +535,7 @@ func (cli *clientImpl) callMemberWithTimeout(ctx context.Context, member exec.Ho
return cli.callMember(callCtx, member, callback)
}

func (cli *clientImpl) isHealthy(ctx context.Context, member exec.HostInfo) error {
// Get or create client for this coordinator
client, err := cli.getOrCreateClient(member)
if err != nil {
return fmt.Errorf("failed to get coordinator client: %w", err)
}

func (cli *clientImpl) isHealthy(ctx context.Context, client *client) error {
// Check health
req := &grpc_health_v1.HealthCheckRequest{
Service: "", // Check overall server health
Expand All @@ -547,13 +553,26 @@ func (cli *clientImpl) isHealthy(ctx context.Context, member exec.HostInfo) erro
return nil
}

// getOrCreateClient gets an existing client or creates a new one for the given member
// getOrCreateClient gets a client for the member without changing the cached address.
func (cli *clientImpl) getOrCreateClient(member exec.HostInfo) (*client, error) {
return cli.getOrCreateClientWithAddressRefresh(member, false)
}

// getOrCreateDiscoveredClient treats the discovered address as authoritative.
func (cli *clientImpl) getOrCreateDiscoveredClient(member exec.HostInfo) (*client, error) {
return cli.getOrCreateClientWithAddressRefresh(member, true)
}

func (cli *clientImpl) getOrCreateClientWithAddressRefresh(member exec.HostInfo, refreshAddress bool) (*client, error) {
key := coordinatorMemberKey(member)
address := coordinatorAddress(member)

// Try to get existing client with read lock
cli.clientsMu.RLock()
if c, exists := cli.clients[key]; exists {
if c, exists := cli.clients[key]; exists &&
(!refreshAddress ||
isOlderCoordinatorIncarnation(member.StartedAt, c.startedAt) ||
(c.address == address && !member.StartedAt.After(c.startedAt))) {
cli.clientsMu.RUnlock()
return c, nil
}
Expand All @@ -565,20 +584,39 @@ func (cli *clientImpl) getOrCreateClient(member exec.HostInfo) (*client, error)

// Double-check after acquiring write lock
if c, exists := cli.clients[key]; exists {
return c, nil
if !refreshAddress || isOlderCoordinatorIncarnation(member.StartedAt, c.startedAt) {
return c, nil
}
if c.address == address {
if member.StartedAt.After(c.startedAt) {
c.startedAt = member.StartedAt
}
return c, nil
}
}

// Create new client
c, err := cli.createClient(member)
if err != nil {
return nil, err
}
if refreshAddress {
c.startedAt = member.StartedAt
}

if stale, exists := cli.clients[key]; exists {
_ = stale.conn.Close()
}

// Cache it
cli.clients[key] = c
return c, nil
}

func isOlderCoordinatorIncarnation(candidate, current time.Time) bool {
return !candidate.IsZero() && !current.IsZero() && candidate.Before(current)
}

// createClient creates a new gRPC client for the given coordinator
func (cli *clientImpl) createClient(member exec.HostInfo) (*client, error) {
// Get dial options based on TLS configuration
Expand All @@ -588,7 +626,7 @@ func (cli *clientImpl) createClient(member exec.HostInfo) (*client, error) {
}

// Construct address from host and port
address := net.JoinHostPort(member.Host, strconv.Itoa(member.Port))
address := coordinatorAddress(member)

// Create gRPC connection
conn, err := grpc.NewClient(address, dialOpts...)
Expand All @@ -597,25 +635,13 @@ func (cli *clientImpl) createClient(member exec.HostInfo) (*client, error) {
}

return &client{
address: address,
conn: conn,
client: coordinatorv1.NewCoordinatorServiceClient(conn),
healthClient: grpc_health_v1.NewHealthClient(conn),
}, nil
}

// removeClient removes a client from the cache
func (cli *clientImpl) removeClient(member exec.HostInfo) {
key := coordinatorMemberKey(member)

cli.clientsMu.Lock()
defer cli.clientsMu.Unlock()

if c, exists := cli.clients[key]; exists {
_ = c.conn.Close()
delete(cli.clients, key)
}
}

// Cleanup cleans up all connections
func (cli *clientImpl) Cleanup(ctx context.Context) error {
cli.clientsMu.Lock()
Expand Down Expand Up @@ -694,7 +720,7 @@ func (cli *clientImpl) GetWorkers(ctx context.Context) ([]*coordinatorv1.WorkerI

for _, member := range members {
// Get or create client for this member
c, err := cli.getOrCreateClient(member)
c, err := cli.getOrCreateDiscoveredClient(member)
if err != nil {
logger.Warn(ctx, "Failed to connect to coordinator",
tag.ID(member.ID),
Expand All @@ -714,11 +740,6 @@ func (cli *clientImpl) GetWorkers(ctx context.Context) ([]*coordinatorv1.WorkerI
tag.Port(member.Port),
tag.Error(err))
lastErr = err

// If this is a connection error, remove the client from cache
if st, ok := status.FromError(err); ok && st.Code() == codes.Unavailable {
cli.removeClient(member)
}
continue
}
successfulReads = true
Expand Down Expand Up @@ -796,6 +817,10 @@ func coordinatorMemberKey(member exec.HostInfo) string {
return fmt.Sprintf("%s:%d", member.Host, member.Port)
}

func coordinatorAddress(member exec.HostInfo) string {
return net.JoinHostPort(member.Host, strconv.Itoa(member.Port))
}

func selectAuthoritativeWorker(current, candidate *coordinatorv1.WorkerInfo) *coordinatorv1.WorkerInfo {
if candidate == nil {
return current
Expand All @@ -811,21 +836,53 @@ func selectAuthoritativeWorker(current, candidate *coordinatorv1.WorkerInfo) *co

// Heartbeat sends a heartbeat to coordinators and returns the response
func (cli *clientImpl) Heartbeat(ctx context.Context, req *coordinatorv1.HeartbeatRequest) (*coordinatorv1.HeartbeatResponse, error) {
if cli.config.HeartbeatTimeout > 0 {
var cancel context.CancelFunc
ctx, cancel = context.WithTimeout(ctx, cli.config.HeartbeatTimeout)
defer cancel()
}

members, err := cli.getCoordinatorMembers(ctx)
if err != nil {
return nil, err
}

var resp *coordinatorv1.HeartbeatResponse
err = cli.attemptCall(ctx, members, func(ctx context.Context, _ exec.HostInfo, client *client) error {
var callErr error
resp, callErr = client.client.Heartbeat(ctx, req)
call := func(ctx context.Context, _ exec.HostInfo, client *client) error {
callResp, callErr := client.client.Heartbeat(ctx, req)
if callErr != nil {
return fmt.Errorf("heartbeat failed: %w", callErr)
}
resp = callResp
return nil
}

deadline, hasDeadline := ctx.Deadline()
if !hasDeadline || len(members) == 1 {
err = cli.attemptCall(ctx, members, call)
return resp, err
}

rand.Shuffle(len(members), func(i, j int) {
members[i], members[j] = members[j], members[i]
})
return resp, err

var lastErr error
for i := range members {
if err := ctx.Err(); err != nil {
return nil, err
}

remainingAttempts := len(members) - i
attemptTimeout := time.Until(deadline) / time.Duration(remainingAttempts)
attemptCtx, cancel := context.WithTimeout(ctx, attemptTimeout)
lastErr = cli.attemptCall(attemptCtx, members[i:i+1], call)
cancel()
if lastErr == nil {
return resp, nil
}
}
return nil, lastErr
}

func (cli *clientImpl) AckTaskClaimTo(ctx context.Context, owner exec.HostInfo, req *coordinatorv1.AckTaskClaimRequest) (*coordinatorv1.AckTaskClaimResponse, error) {
Expand Down Expand Up @@ -937,13 +994,13 @@ func openStreamWithFailover[T any](

var lastErr error
for _, member := range members {
memberClient, err := cli.getOrCreateClient(member)
memberClient, err := cli.getOrCreateDiscoveredClient(member)
if err != nil {
cli.recordFailure(err)
lastErr = err
continue
}
if err := cli.isHealthy(ctx, member); err != nil {
if err := cli.isHealthy(ctx, memberClient); err != nil {
cli.recordFailure(err)
lastErr = err
continue
Expand Down Expand Up @@ -1100,12 +1157,12 @@ func (cli *clientImpl) PutWorkspaceBundle(ctx context.Context, desc workspacebun
var satisfied int
errs := make([]error, 0)
for _, member := range members {
memberClient, err := cli.getOrCreateClient(member)
memberClient, err := cli.getOrCreateDiscoveredClient(member)
if err != nil {
errs = append(errs, fmt.Errorf("coordinator %q: %w", member.ID, err))
continue
}
if err := cli.isHealthy(ctx, member); err != nil {
if err := cli.isHealthy(ctx, memberClient); err != nil {
errs = append(errs, fmt.Errorf("coordinator %q is unhealthy: %w", member.ID, err))
continue
}
Expand Down
9 changes: 1 addition & 8 deletions internal/service/coordinator/client_internal_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -111,12 +111,5 @@ func TestClientCacheUsesDerivedKeyForEmptyCoordinatorIDs(t *testing.T) {
require.Len(t, cli.clients, 2)
require.Contains(t, cli.clients, coordinatorMemberKey(member1))
require.Contains(t, cli.clients, coordinatorMemberKey(member2))

cli.removeClient(member1)
require.Len(t, cli.clients, 1)
require.NotContains(t, cli.clients, coordinatorMemberKey(member1))
require.Contains(t, cli.clients, coordinatorMemberKey(member2))

cli.removeClient(member2)
require.Empty(t, cli.clients)
require.NoError(t, cli.Cleanup(t.Context()))
}
2 changes: 2 additions & 0 deletions internal/service/coordinator/client_state_internal_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -54,9 +54,11 @@ func TestPinnedStateCoordinatorUsesDeterministicOwnerWithoutHealthFailover(t *te
cli := &clientImpl{
config: DefaultConfig(),
registry: staticStateRegistry{members: members},
clients: make(map[string]*client),
stateCoordinators: make(map[string]pinnedStateCoordinator),
state: &Metrics{IsConnected: true},
}
t.Cleanup(func() { require.NoError(t, cli.Cleanup(t.Context())) })

pinned, err := cli.pinnedStateCoordinator(context.Background(), key)
require.NoError(t, err)
Expand Down
Loading