diff --git a/internal/controller/worker/grpc_handler.go b/internal/controller/worker/grpc_handler.go new file mode 100644 index 0000000..62b2313 --- /dev/null +++ b/internal/controller/worker/grpc_handler.go @@ -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 +} diff --git a/internal/controller/worker/grpc_handler_test.go b/internal/controller/worker/grpc_handler_test.go new file mode 100644 index 0000000..a6f4e69 --- /dev/null +++ b/internal/controller/worker/grpc_handler_test.go @@ -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() } +}