proxy-pool/internal/controller/worker/grpc_handler_test.go
youfak 6b6fb54075
Some checks are pending
ci / proto (push) Waiting to run
ci / test (ubuntu-latest) (push) Waiting to run
ci / test (windows-latest) (push) Waiting to run
ci / race (push) Waiting to run
ci / integration (push) Waiting to run
feat: publish initial worker snapshots
2026-07-31 13:22:45 +08:00

146 lines
6.2 KiB
Go

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)
}
}
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)
}
}
type grpcServiceStub struct {
registration Registration
registerErr error
acknowledgement SnapshotAcknowledgement
acknowledgeErr error
report workerruntime.Report
reportErr error
issued workerruntime.SnapshotReference
}
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) IssueSnapshot(_ context.Context, reference workerruntime.SnapshotReference) error {
stub.issued = reference
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
}
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() }
}