364 lines
15 KiB
Go
364 lines
15 KiB
Go
package worker
|
|
|
|
import (
|
|
"context"
|
|
"io"
|
|
"net"
|
|
"testing"
|
|
"time"
|
|
|
|
controlplanev1 "proxy-pool/gen/controlplane/v1"
|
|
"proxy-pool/internal/domain/outcome"
|
|
"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/proto"
|
|
"google.golang.org/protobuf/types/known/durationpb"
|
|
"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 TestGRPCHandlerMapsErrorsAndLeavesSnapshotStreamUnimplemented(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)
|
|
}
|
|
}
|
|
|
|
func TestGRPCHandlerReportsOutcomeStream(t *testing.T) {
|
|
service := &grpcServiceStub{}
|
|
client, cleanup := grpcWorkerClient(t, service, allowIdentity{})
|
|
defer cleanup()
|
|
stream, err := client.ReportOutcomes(context.Background())
|
|
if err != nil {
|
|
t.Fatalf("ReportOutcomes() = %v", err)
|
|
}
|
|
if err := stream.Send(&controlplanev1.OutcomeBatch{
|
|
WorkerId: "worker-a", SessionId: "session-a", Sequence: 3,
|
|
Outcomes: []*controlplanev1.ProxyOutcome{{
|
|
ProxyId: "proxy-a", RoutingName: "route-a", Stage: controlplanev1.OutcomeStage_OUTCOME_STAGE_DIAL,
|
|
Success: false, ErrorClass: "timeout", Latency: durationpb.New(25 * time.Millisecond), ObservedAt: timestamppb.New(time.Now()),
|
|
}},
|
|
}); err != nil {
|
|
t.Fatalf("Send() = %v", err)
|
|
}
|
|
response, err := stream.CloseAndRecv()
|
|
if err != nil || response.GetAcceptedThroughSequence() != 3 || service.outcome.Sequence != 3 ||
|
|
service.outcome.Events[0].ErrorClass != outcome.ErrorClassTimeout {
|
|
t.Fatalf("CloseAndRecv() = %+v, %v; outcome = %+v", response, err, service.outcome)
|
|
}
|
|
}
|
|
|
|
func TestGRPCHandlerFencesOutcomesThroughWorkerService(t *testing.T) {
|
|
now := time.Date(2026, 7, 31, 12, 0, 0, 0, time.UTC)
|
|
store, err := workerruntime.NewMemoryStore(func() time.Time { return now })
|
|
if err != nil {
|
|
t.Fatalf("NewMemoryStore() = %v", err)
|
|
}
|
|
service, err := NewService(store, Options{
|
|
ProtocolVersion: 1, HeartbeatInterval: time.Second, SessionTTL: 3 * time.Second,
|
|
MaxStaleAge: time.Second, MaxRuntimeCounters: 4,
|
|
SessionID: func() (string, error) { return "session-a", nil },
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("NewService() = %v", err)
|
|
}
|
|
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" {
|
|
t.Fatalf("RegisterWorker() = %+v, %v", registered, err)
|
|
}
|
|
batch := &controlplanev1.OutcomeBatch{
|
|
WorkerId: "worker-a", SessionId: "session-a", Sequence: 1,
|
|
Outcomes: []*controlplanev1.ProxyOutcome{{
|
|
ProxyId: "proxy-a", Stage: controlplanev1.OutcomeStage_OUTCOME_STAGE_DIAL, Success: true,
|
|
Latency: durationpb.New(time.Millisecond), ObservedAt: timestamppb.New(now),
|
|
}},
|
|
}
|
|
if accepted, err := reportOutcomeBatch(context.Background(), client, batch); err != nil || accepted != 1 {
|
|
t.Fatalf("ReportOutcomes(first) = %d, %v", accepted, err)
|
|
}
|
|
if accepted, err := reportOutcomeBatch(context.Background(), client, batch); err != nil || accepted != 1 {
|
|
t.Fatalf("ReportOutcomes(replay) = %d, %v", accepted, err)
|
|
}
|
|
conflicting := proto.Clone(batch).(*controlplanev1.OutcomeBatch)
|
|
conflicting.Outcomes[0].Success = false
|
|
conflicting.Outcomes[0].ErrorClass = "dial"
|
|
if _, err := reportOutcomeBatch(context.Background(), client, conflicting); status.Code(err) != codes.AlreadyExists {
|
|
t.Fatalf("ReportOutcomes(conflict) code = %s, want AlreadyExists; error=%v", status.Code(err), err)
|
|
}
|
|
}
|
|
|
|
func reportOutcomeBatch(ctx context.Context, client controlplanev1.WorkerControlPlaneClient, batch *controlplanev1.OutcomeBatch) (uint64, error) {
|
|
stream, err := client.ReportOutcomes(ctx)
|
|
if err != nil {
|
|
return 0, err
|
|
}
|
|
if err := stream.Send(batch); err != nil {
|
|
return 0, err
|
|
}
|
|
response, err := stream.CloseAndRecv()
|
|
if err != nil {
|
|
return 0, err
|
|
}
|
|
return response.GetAcceptedThroughSequence(), nil
|
|
}
|
|
|
|
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 TestGRPCHandlerBindsDrainBarriersAfterIssuingExcludedSnapshot(t *testing.T) {
|
|
checksum := make([]byte, 32)
|
|
checksum[0] = 1
|
|
service := &drainBindingServiceStub{grpcServiceStub: grpcServiceStub{}}
|
|
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)),
|
|
Proxies: []*controlplanev1.OwnedProxy{{Id: "proxy-present"}},
|
|
}}})
|
|
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(); err != nil {
|
|
t.Fatalf("Recv(): %v", err)
|
|
}
|
|
if !service.boundAfterIssue || service.boundReference.Version != 3 || len(service.presentProxyIDs) != 1 || service.presentProxyIDs[0] != "proxy-present" {
|
|
t.Fatalf("BindDrainBarriers() = afterIssue:%t reference:%+v present:%v", service.boundAfterIssue, service.boundReference, service.presentProxyIDs)
|
|
}
|
|
}
|
|
|
|
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
|
|
outcome outcome.Batch
|
|
outcomeAccepted uint64
|
|
outcomeErr error
|
|
issued workerruntime.SnapshotReference
|
|
issuedSessionID string
|
|
validateErr error
|
|
issueErr error
|
|
}
|
|
|
|
type drainBindingServiceStub struct {
|
|
grpcServiceStub
|
|
boundAfterIssue bool
|
|
boundReference workerruntime.SnapshotReference
|
|
presentProxyIDs []string
|
|
}
|
|
|
|
func (stub *drainBindingServiceStub) BindDrainBarriers(_ context.Context, _ string, _ string, reference workerruntime.SnapshotReference, proxyIDs []string) error {
|
|
stub.boundAfterIssue = stub.issued.Version == reference.Version
|
|
stub.boundReference = reference
|
|
stub.presentProxyIDs = append([]string(nil), proxyIDs...)
|
|
return nil
|
|
}
|
|
|
|
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
|
|
}
|
|
func (stub *grpcServiceStub) ReportOutcomes(_ context.Context, batch outcome.Batch) (uint64, error) {
|
|
stub.outcome = batch
|
|
if stub.outcomeAccepted == 0 {
|
|
stub.outcomeAccepted = batch.Sequence
|
|
}
|
|
return stub.outcomeAccepted, stub.outcomeErr
|
|
}
|
|
|
|
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() }
|
|
}
|