package worker import ( "context" "errors" "net" "testing" "time" controlplanev1 "proxy-pool/gen/controlplane/v1" "proxy-pool/internal/config" "google.golang.org/grpc" "google.golang.org/grpc/credentials/insecure" ) func TestNewServerRejectsInvalidOptions(t *testing.T) { controlPlane := validServerControlPlane() service := &grpcServiceStub{} if _, err := NewServer(controlPlane, nil, DefaultServerOptions()); !errors.Is(err, ErrInvalidServer) { t.Fatalf("NewServer(nil service) error = %v, want ErrInvalidServer", err) } if _, err := NewServer(controlPlane, service, ServerOptions{ShutdownTimeout: -time.Second}); !errors.Is(err, ErrInvalidServer) { t.Fatalf("NewServer(negative shutdown timeout) error = %v, want ErrInvalidServer", err) } controlPlane.Enabled = false if _, err := NewServer(controlPlane, service, DefaultServerOptions()); !errors.Is(err, ErrInvalidServer) { t.Fatalf("NewServer(disabled) error = %v, want ErrInvalidServer", err) } } func TestServerServesAndStopsOnContextCancellation(t *testing.T) { listener, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("net.Listen(): %v", err) } controlPlane := validServerControlPlane() controlPlane.Listen = listener.Addr().String() service := &grpcServiceStub{registration: Registration{WorkerID: "worker-a", SessionID: "session-a", OwnershipEpoch: 5, HeartbeatInterval: time.Second, MaxStaleAge: 2 * time.Second}} server, err := NewServer(controlPlane, service, ServerOptions{ShutdownTimeout: time.Second}) if err != nil { t.Fatalf("NewServer(): %v", err) } ctx, cancel := context.WithCancel(context.Background()) result := make(chan error, 1) go func() { result <- server.Serve(ctx, listener) }() connection, err := grpc.NewClient(listener.Addr().String(), grpc.WithTransportCredentials(insecure.NewCredentials())) if err != nil { t.Fatalf("grpc.NewClient(): %v", err) } client := controlplanev1.NewWorkerControlPlaneClient(connection) requestCtx, requestCancel := context.WithTimeout(context.Background(), 3*time.Second) defer requestCancel() response, err := client.RegisterWorker(requestCtx, &controlplanev1.RegisterWorkerRequest{WorkerId: "worker-a", SupportedProtocolVersion: 1}) if err != nil || response.GetSessionId() != "session-a" { t.Fatalf("RegisterWorker() = %+v, %v", response, err) } if err := connection.Close(); err != nil { t.Fatalf("connection.Close(): %v", err) } cancel() select { case err := <-result: if err != nil { t.Fatalf("Serve() error = %v", err) } case <-time.After(3 * time.Second): t.Fatal("Serve() did not stop after context cancellation") } } func TestServerRegistersCheckerServiceOnTheExistingControlPlaneListener(t *testing.T) { listener, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("net.Listen(): %v", err) } controlPlane := validServerControlPlane() controlPlane.Listen = listener.Addr().String() checker := &checkerServiceStub{} server, err := NewServer(controlPlane, &grpcServiceStub{}, ServerOptions{ShutdownTimeout: time.Second, Checker: checker}) if err != nil { t.Fatalf("NewServer(): %v", err) } ctx, cancel := context.WithCancel(context.Background()) result := make(chan error, 1) go func() { result <- server.Serve(ctx, listener) }() connection, err := grpc.NewClient(listener.Addr().String(), grpc.WithTransportCredentials(insecure.NewCredentials())) if err != nil { t.Fatalf("grpc.NewClient(): %v", err) } response, err := controlplanev1.NewCheckerControlPlaneClient(connection).ReportObservations(context.Background(), &controlplanev1.ObservationBatch{ CheckerId: "checker-a", }) if err != nil || response.GetAccepted() != 1 || checker.checkerID != "checker-a" { t.Fatalf("ReportObservations() = (%+v, %v); checker=%q", response, err, checker.checkerID) } if err := connection.Close(); err != nil { t.Fatalf("connection.Close(): %v", err) } cancel() select { case err := <-result: if err != nil { t.Fatalf("Serve() error = %v", err) } case <-time.After(3 * time.Second): t.Fatal("Serve() did not stop after context cancellation") } } type checkerServiceStub struct { controlplanev1.UnimplementedCheckerControlPlaneServer checkerID string } func (stub *checkerServiceStub) ReportObservations(_ context.Context, request *controlplanev1.ObservationBatch) (*controlplanev1.ReportObservationsResponse, error) { stub.checkerID = request.GetCheckerId() return &controlplanev1.ReportObservationsResponse{Accepted: 1}, nil } func validServerControlPlane() config.ControlPlane { return config.ControlPlane{ Enabled: true, Listen: "127.0.0.1:8443", ProtocolVersion: 1, HeartbeatInterval: config.Duration(10 * time.Second), SessionTTL: config.Duration(30 * time.Second), MaxStaleAge: config.Duration(10 * time.Second), MaxMessageBytes: 1 << 20, MaxRuntimeCounters: 100, MaxConcurrentStreams: 10, TLS: config.ControlPlaneTLS{Mode: "disabled"}, } }