package worker import ( "context" "io" "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) } } func TestGRPCHandlerStreamsIssuedFullSnapshots(t *testing.T) { checksum := make([]byte, 32) checksum[0] = 1 full := &controlplanev1.WorkerSnapshot{ Version: 3, OwnershipEpoch: 9, Checksum: checksum, GeneratedAt: timestamppb.New(time.Now()), ValidUntil: timestamppb.New(time.Now().Add(time.Minute)), } service := &grpcServiceStub{} client, cleanup := grpcWorkerClient(t, service, allowIdentity{}, snapshotSourceStub{snapshots: []*controlplanev1.WorkerSnapshot{full}}) defer cleanup() stream, err := client.WatchSnapshots(context.Background(), &controlplanev1.WatchSnapshotsRequest{WorkerId: "worker-a", SessionId: "session-a"}) if err != nil { t.Fatalf("WatchSnapshots(): %v", err) } received, err := stream.Recv() if err != nil || received.GetFull().GetVersion() != 3 { t.Fatalf("Recv() = %+v, %v", received, err) } if service.issued.WorkerID != "worker-a" || service.issued.Version != 3 || service.issued.Checksum[0] != 1 { t.Fatalf("issued snapshot = %+v", service.issued) } if service.issuedSessionID != "session-a" { t.Fatalf("issued session = %q, want session-a", service.issuedSessionID) } } func TestGRPCHandlerClosesSnapshotStreamAtValidityDeadline(t *testing.T) { checksum := make([]byte, 32) checksum[0] = 1 full := &controlplanev1.WorkerSnapshot{ Version: 3, OwnershipEpoch: 9, Checksum: checksum, GeneratedAt: timestamppb.New(time.Now()), ValidUntil: timestamppb.New(time.Now().Add(100 * time.Millisecond)), } service := &grpcServiceStub{} client, cleanup := grpcWorkerClient(t, service, allowIdentity{}, holdingSnapshotSource{snapshot: full}) defer cleanup() ctx, cancel := context.WithTimeout(context.Background(), time.Second) defer cancel() stream, err := client.WatchSnapshots(ctx, &controlplanev1.WatchSnapshotsRequest{WorkerId: "worker-a", SessionId: "session-a"}) if err != nil { t.Fatalf("WatchSnapshots(): %v", err) } if received, err := stream.Recv(); err != nil || received.GetFull().GetVersion() != 3 { t.Fatalf("first Recv() = %+v, %v", received, err) } if _, err := stream.Recv(); err != io.EOF { t.Fatalf("Recv(after validity deadline) error = %v, want EOF", err) } } func TestGRPCHandlerRejectsExpiredSnapshot(t *testing.T) { checksum := make([]byte, 32) checksum[0] = 1 full := &controlplanev1.WorkerSnapshot{ Version: 3, OwnershipEpoch: 9, Checksum: checksum, GeneratedAt: timestamppb.New(time.Now().Add(-time.Minute)), ValidUntil: timestamppb.New(time.Now().Add(-time.Second)), } service := &grpcServiceStub{} client, cleanup := grpcWorkerClient(t, service, allowIdentity{}, snapshotSourceStub{snapshots: []*controlplanev1.WorkerSnapshot{full}}) defer cleanup() stream, err := client.WatchSnapshots(context.Background(), &controlplanev1.WatchSnapshotsRequest{WorkerId: "worker-a", SessionId: "session-a"}) if err != nil { t.Fatalf("WatchSnapshots(): %v", err) } if _, err := stream.Recv(); status.Code(err) != codes.InvalidArgument { t.Fatalf("Recv() code = %s, want InvalidArgument; error=%v", status.Code(err), err) } if service.issued.Version != 0 { t.Fatalf("expired snapshot was issued: %+v", service.issued) } } func TestGRPCHandlerDoesNotDeliverSnapshotWhenSessionBecomesStale(t *testing.T) { checksum := make([]byte, 32) checksum[0] = 1 service := &grpcServiceStub{issueErr: workerruntime.ErrStaleSession} client, cleanup := grpcWorkerClient(t, service, allowIdentity{}, snapshotSourceStub{snapshots: []*controlplanev1.WorkerSnapshot{{ Version: 3, OwnershipEpoch: 9, Checksum: checksum, GeneratedAt: timestamppb.New(time.Now()), ValidUntil: timestamppb.New(time.Now().Add(time.Minute)), }}}) defer cleanup() stream, err := client.WatchSnapshots(context.Background(), &controlplanev1.WatchSnapshotsRequest{WorkerId: "worker-a", SessionId: "session-a"}) if err != nil { t.Fatalf("WatchSnapshots(): %v", err) } _, err = stream.Recv() if status.Code(err) != codes.FailedPrecondition { t.Fatalf("Recv() error = %v, want FailedPrecondition", err) } if service.issued.Version != 0 || service.issuedSessionID != "" { t.Fatalf("stale session issued snapshot = %+v for %q", service.issued, service.issuedSessionID) } } type grpcServiceStub struct { registration Registration registerErr error acknowledgement SnapshotAcknowledgement acknowledgeErr error report workerruntime.Report reportErr error issued workerruntime.SnapshotReference issuedSessionID string validateErr error issueErr error } func (stub *grpcServiceStub) Register(context.Context, RegisterCommand) (Registration, error) { return stub.registration, stub.registerErr } func (stub *grpcServiceStub) CurrentOwnershipEpoch(context.Context) (uint64, error) { return 9, nil } func (stub *grpcServiceStub) ValidateSession(context.Context, string, string) error { return stub.validateErr } func (stub *grpcServiceStub) IssueSnapshot(_ context.Context, sessionID string, reference workerruntime.SnapshotReference) error { if stub.issueErr != nil { return stub.issueErr } stub.issued = reference stub.issuedSessionID = sessionID return nil } 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 } type snapshotSourceStub struct { snapshots []*controlplanev1.WorkerSnapshot } type holdingSnapshotSource struct { snapshot *controlplanev1.WorkerSnapshot } func (source holdingSnapshotSource) Watch(_ context.Context, _ SnapshotWatchRequest) (<-chan *controlplanev1.WorkerSnapshot, error) { updates := make(chan *controlplanev1.WorkerSnapshot, 1) updates <- source.snapshot return updates, nil } func (source snapshotSourceStub) Watch(_ context.Context, _ SnapshotWatchRequest) (<-chan *controlplanev1.WorkerSnapshot, error) { updates := make(chan *controlplanev1.WorkerSnapshot, len(source.snapshots)) for _, snapshot := range source.snapshots { updates <- snapshot } close(updates) return updates, nil } func grpcWorkerClient(t *testing.T, service Service, identity IdentityAuthorizer, snapshots ...SnapshotSource) (controlplanev1.WorkerControlPlaneClient, func()) { t.Helper() listener := bufconn.Listen(1 << 20) server := grpc.NewServer() controlplanev1.RegisterWorkerControlPlaneServer(server, NewGRPCHandler(service, identity, snapshots...)) 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() } }