feat: map worker grpc requests
This commit is contained in:
parent
05e1758c00
commit
5a678dc66f
136
internal/controller/worker/grpc_handler.go
Normal file
136
internal/controller/worker/grpc_handler.go
Normal file
@ -0,0 +1,136 @@
|
||||
package worker
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
|
||||
controlplanev1 "proxy-pool/gen/controlplane/v1"
|
||||
"proxy-pool/internal/domain/workerruntime"
|
||||
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
"google.golang.org/protobuf/types/known/durationpb"
|
||||
"google.golang.org/protobuf/types/known/emptypb"
|
||||
)
|
||||
|
||||
type IdentityAuthorizer interface {
|
||||
Authorize(context.Context, string) error
|
||||
}
|
||||
|
||||
type GRPCHandler struct {
|
||||
controlplanev1.UnimplementedWorkerControlPlaneServer
|
||||
service Service
|
||||
identity IdentityAuthorizer
|
||||
}
|
||||
|
||||
func NewGRPCHandler(service Service, identity IdentityAuthorizer) *GRPCHandler {
|
||||
return &GRPCHandler{service: service, identity: identity}
|
||||
}
|
||||
|
||||
func (handler *GRPCHandler) RegisterWorker(ctx context.Context, request *controlplanev1.RegisterWorkerRequest) (*controlplanev1.RegisterWorkerResponse, error) {
|
||||
if request == nil || handler == nil || handler.service == nil || handler.identity == nil {
|
||||
return nil, grpcError(ErrInvalidCommand)
|
||||
}
|
||||
if err := handler.authorize(ctx, request.GetWorkerId()); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
registration, err := handler.service.Register(ctx, RegisterCommand{
|
||||
WorkerID: request.GetWorkerId(), InstanceID: request.GetInstanceId(), Zone: request.GetZone(),
|
||||
ProtocolVersion: request.GetSupportedProtocolVersion(), Labels: cloneLabels(request.GetLabels()),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, grpcError(err)
|
||||
}
|
||||
return &controlplanev1.RegisterWorkerResponse{
|
||||
WorkerId: registration.WorkerID, SessionId: registration.SessionID, OwnershipEpoch: registration.OwnershipEpoch,
|
||||
HeartbeatInterval: durationpb.New(registration.HeartbeatInterval), MaxStaleAge: durationpb.New(registration.MaxStaleAge),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (handler *GRPCHandler) AcknowledgeSnapshot(ctx context.Context, request *controlplanev1.AcknowledgeSnapshotRequest) (*emptypb.Empty, error) {
|
||||
if request == nil || handler == nil || handler.service == nil || handler.identity == nil {
|
||||
return nil, grpcError(ErrInvalidCommand)
|
||||
}
|
||||
if err := handler.authorize(ctx, request.GetWorkerId()); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
err := handler.service.Acknowledge(ctx, SnapshotAcknowledgement{
|
||||
WorkerID: request.GetWorkerId(), SessionID: request.GetSessionId(), Version: request.GetVersion(),
|
||||
OwnershipEpoch: request.GetOwnershipEpoch(), Checksum: append([]byte(nil), request.GetChecksum()...),
|
||||
Applied: request.GetApplied(), ErrorCode: request.GetErrorCode(), ErrorMessage: request.GetErrorMessage(),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, grpcError(err)
|
||||
}
|
||||
return &emptypb.Empty{}, nil
|
||||
}
|
||||
|
||||
func (handler *GRPCHandler) ReportRuntime(ctx context.Context, request *controlplanev1.ReportRuntimeRequest) (*controlplanev1.ReportRuntimeResponse, error) {
|
||||
if request == nil || handler == nil || handler.service == nil || handler.identity == nil || request.GetObservedAt() == nil || request.GetObservedAt().CheckValid() != nil {
|
||||
return nil, grpcError(ErrInvalidCommand)
|
||||
}
|
||||
if err := handler.authorize(ctx, request.GetWorkerId()); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
counters := make([]workerruntime.Counter, len(request.GetCounters()))
|
||||
for index, counter := range request.GetCounters() {
|
||||
if counter == nil {
|
||||
return nil, grpcError(ErrInvalidCommand)
|
||||
}
|
||||
counters[index] = workerruntime.Counter{
|
||||
ProxyID: counter.GetProxyId(), Active: int64(counter.GetActive()), Reserved: int64(counter.GetReserved()), Draining: counter.GetDraining(),
|
||||
}
|
||||
}
|
||||
decision, err := handler.service.ReportRuntime(ctx, workerruntime.Report{
|
||||
WorkerID: request.GetWorkerId(), SessionID: request.GetSessionId(), Sequence: request.GetReportSequence(),
|
||||
SnapshotVersion: request.GetSnapshotVersion(), OwnershipEpoch: request.GetOwnershipEpoch(),
|
||||
ObservedAt: request.GetObservedAt().AsTime(), Counters: counters,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, grpcError(err)
|
||||
}
|
||||
return &controlplanev1.ReportRuntimeResponse{
|
||||
AcceptedOwnershipEpoch: decision.AcceptedOwnershipEpoch, RequireFullSnapshot: decision.RequireFullSnapshot,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (handler *GRPCHandler) authorize(ctx context.Context, workerID string) error {
|
||||
if err := handler.identity.Authorize(ctx, workerID); err != nil {
|
||||
return status.Error(codes.PermissionDenied, "worker identity is not authorized")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func grpcError(err error) error {
|
||||
switch {
|
||||
case errors.Is(err, context.Canceled):
|
||||
return status.Error(codes.Canceled, "worker control request canceled")
|
||||
case errors.Is(err, context.DeadlineExceeded):
|
||||
return status.Error(codes.DeadlineExceeded, "worker control request deadline exceeded")
|
||||
case errors.Is(err, ErrInvalidCommand), errors.Is(err, workerruntime.ErrInvalidReport),
|
||||
errors.Is(err, workerruntime.ErrInvalidAcknowledgement), errors.Is(err, workerruntime.ErrInvalidSnapshotReference):
|
||||
return status.Error(codes.InvalidArgument, "invalid worker control request")
|
||||
case errors.Is(err, ErrProtocolVersion):
|
||||
return status.Error(codes.FailedPrecondition, "unsupported worker protocol version")
|
||||
case errors.Is(err, workerruntime.ErrStaleSession):
|
||||
return status.Error(codes.FailedPrecondition, "worker session is stale")
|
||||
case errors.Is(err, workerruntime.ErrSnapshotMismatch):
|
||||
return status.Error(codes.FailedPrecondition, "worker snapshot does not match issued snapshot")
|
||||
case errors.Is(err, workerruntime.ErrStaleAcknowledgement):
|
||||
return status.Error(codes.Aborted, "worker snapshot acknowledgement is stale")
|
||||
case errors.Is(err, workerruntime.ErrStaleReport):
|
||||
return status.Error(codes.Aborted, "worker runtime sequence is stale")
|
||||
case errors.Is(err, workerruntime.ErrConflictingReport):
|
||||
return status.Error(codes.AlreadyExists, "worker runtime sequence conflicts")
|
||||
default:
|
||||
return status.Error(codes.Unavailable, "worker control plane unavailable")
|
||||
}
|
||||
}
|
||||
|
||||
func cloneLabels(labels map[string]string) map[string]string {
|
||||
result := make(map[string]string, len(labels))
|
||||
for key, value := range labels {
|
||||
result[key] = value
|
||||
}
|
||||
return result
|
||||
}
|
||||
103
internal/controller/worker/grpc_handler_test.go
Normal file
103
internal/controller/worker/grpc_handler_test.go
Normal file
@ -0,0 +1,103 @@
|
||||
package worker
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
controlplanev1 "proxy-pool/gen/controlplane/v1"
|
||||
"proxy-pool/internal/domain/workerruntime"
|
||||
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
"google.golang.org/grpc/test/bufconn"
|
||||
"google.golang.org/protobuf/types/known/timestamppb"
|
||||
)
|
||||
|
||||
func TestGRPCHandlerMapsWorkerRequests(t *testing.T) {
|
||||
service := &grpcServiceStub{registration: Registration{WorkerID: "worker-a", SessionID: "session-a", OwnershipEpoch: 9, HeartbeatInterval: time.Second, MaxStaleAge: 2 * time.Second}}
|
||||
client, cleanup := grpcWorkerClient(t, service, allowIdentity{})
|
||||
defer cleanup()
|
||||
registered, err := client.RegisterWorker(context.Background(), &controlplanev1.RegisterWorkerRequest{
|
||||
WorkerId: "worker-a", InstanceId: "instance-a", Zone: "zone-a", SupportedProtocolVersion: 1,
|
||||
})
|
||||
if err != nil || registered.GetSessionId() != "session-a" || registered.GetHeartbeatInterval().AsDuration() != time.Second {
|
||||
t.Fatalf("RegisterWorker() = %+v, %v", registered, err)
|
||||
}
|
||||
if err := service.acknowledgeErr; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err = client.AcknowledgeSnapshot(context.Background(), &controlplanev1.AcknowledgeSnapshotRequest{
|
||||
WorkerId: "worker-a", SessionId: "session-a", Version: 7, OwnershipEpoch: 9, Checksum: make([]byte, 32),
|
||||
})
|
||||
if err != nil || service.acknowledgement.Version != 7 {
|
||||
t.Fatalf("AcknowledgeSnapshot() error = %v; command=%+v", err, service.acknowledgement)
|
||||
}
|
||||
response, err := client.ReportRuntime(context.Background(), &controlplanev1.ReportRuntimeRequest{
|
||||
WorkerId: "worker-a", SessionId: "session-a", SnapshotVersion: 7, OwnershipEpoch: 9, ReportSequence: 1,
|
||||
ObservedAt: timestamppb.New(time.Now()), Counters: []*controlplanev1.ProxyRuntime{{ProxyId: "proxy-a", Active: 2, Reserved: 1}},
|
||||
})
|
||||
if err != nil || response.GetAcceptedOwnershipEpoch() != 9 || service.report.Counters[0].Active != 2 {
|
||||
t.Fatalf("ReportRuntime() = %+v, %v; report=%+v", response, err, service.report)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGRPCHandlerMapsErrorsAndLeavesStreamsUnimplemented(t *testing.T) {
|
||||
service := &grpcServiceStub{registerErr: ErrProtocolVersion}
|
||||
client, cleanup := grpcWorkerClient(t, service, allowIdentity{})
|
||||
defer cleanup()
|
||||
_, err := client.RegisterWorker(context.Background(), &controlplanev1.RegisterWorkerRequest{WorkerId: "worker-a"})
|
||||
if status.Code(err) != codes.FailedPrecondition {
|
||||
t.Fatalf("RegisterWorker() code = %s, want FailedPrecondition", status.Code(err))
|
||||
}
|
||||
stream, err := client.WatchSnapshots(context.Background(), &controlplanev1.WatchSnapshotsRequest{})
|
||||
_, streamErr := stream.Recv()
|
||||
if err != nil || status.Code(streamErr) != codes.Unimplemented {
|
||||
t.Fatalf("WatchSnapshots() = %v, %v", err, streamErr)
|
||||
}
|
||||
outcomes, err := client.ReportOutcomes(context.Background())
|
||||
_, outcomesErr := outcomes.CloseAndRecv()
|
||||
if err != nil || status.Code(outcomesErr) != codes.Unimplemented {
|
||||
t.Fatalf("ReportOutcomes() = %v, %v", err, outcomesErr)
|
||||
}
|
||||
}
|
||||
|
||||
type grpcServiceStub struct {
|
||||
registration Registration
|
||||
registerErr error
|
||||
acknowledgement SnapshotAcknowledgement
|
||||
acknowledgeErr error
|
||||
report workerruntime.Report
|
||||
reportErr error
|
||||
}
|
||||
|
||||
func (stub *grpcServiceStub) Register(context.Context, RegisterCommand) (Registration, error) {
|
||||
return stub.registration, stub.registerErr
|
||||
}
|
||||
func (stub *grpcServiceStub) Acknowledge(_ context.Context, acknowledgement SnapshotAcknowledgement) error {
|
||||
stub.acknowledgement = acknowledgement
|
||||
return stub.acknowledgeErr
|
||||
}
|
||||
func (stub *grpcServiceStub) ReportRuntime(_ context.Context, report workerruntime.Report) (RuntimeDecision, error) {
|
||||
stub.report = report
|
||||
return RuntimeDecision{AcceptedOwnershipEpoch: 9}, stub.reportErr
|
||||
}
|
||||
|
||||
type allowIdentity struct{}
|
||||
|
||||
func (allowIdentity) Authorize(context.Context, string) error { return nil }
|
||||
|
||||
func grpcWorkerClient(t *testing.T, service Service, identity IdentityAuthorizer) (controlplanev1.WorkerControlPlaneClient, func()) {
|
||||
t.Helper()
|
||||
listener := bufconn.Listen(1 << 20)
|
||||
server := grpc.NewServer()
|
||||
controlplanev1.RegisterWorkerControlPlaneServer(server, NewGRPCHandler(service, identity))
|
||||
go func() { _ = server.Serve(listener) }()
|
||||
connection, err := grpc.NewClient("passthrough:///bufnet", grpc.WithContextDialer(func(context.Context, string) (net.Conn, error) { return listener.Dial() }), grpc.WithInsecure())
|
||||
if err != nil {
|
||||
t.Fatalf("grpc.NewClient(): %v", err)
|
||||
}
|
||||
return controlplanev1.NewWorkerControlPlaneClient(connection), func() { _ = connection.Close(); server.Stop(); _ = listener.Close() }
|
||||
}
|
||||
Loading…
Reference in New Issue
Block a user