diff --git a/pkg/beholder/batch_emitter_service.go b/pkg/beholder/batch_emitter_service.go index d2560a981e..865251b619 100644 --- a/pkg/beholder/batch_emitter_service.go +++ b/pkg/beholder/batch_emitter_service.go @@ -171,8 +171,7 @@ func (e *ChipIngressBatchEmitterService) emitInternal(ctx context.Context, body // every dropped event. chip_ingress.events_dropped (error_code) already // captures that it's happening and roughly why; the full reason isn't // needed at fleet-wide log volume. - var pubErr *batch.PublishError - if !errors.As(sendErr, &pubErr) { + if _, ok := errors.AsType[*batch.PublishError](sendErr); !ok { e.eng.Errorw("failed to emit to chip ingress", "error", sendErr, "error_code", errorCode, diff --git a/pkg/chipingress/batch/client.go b/pkg/chipingress/batch/client.go index b07d9c7999..52d32b4619 100644 --- a/pkg/chipingress/batch/client.go +++ b/pkg/chipingress/batch/client.go @@ -65,8 +65,7 @@ func ErrorCodeFor(err error) string { return "" } - var pubErr *PublishError - if errors.As(err, &pubErr) { + if pubErr, ok := errors.AsType[*PublishError](err); ok { if pubErr.Code == ErrCodeResultsMismatch { return "results_mismatch" } diff --git a/pkg/config/toml.go b/pkg/config/toml.go index f51db76365..00d65b354e 100644 --- a/pkg/config/toml.go +++ b/pkg/config/toml.go @@ -12,8 +12,7 @@ import ( func DecodeTOML(r io.Reader, v any) error { d := toml.NewDecoder(r).DisallowUnknownFields() if err := d.Decode(v); err != nil { - var strict *toml.StrictMissingError - if errors.As(err, &strict) { + if strict, ok := errors.AsType[*toml.StrictMissingError](err); ok { return errors.New(strict.String()) } return err diff --git a/pkg/loop/internal/core/services/capability/capabilities.go b/pkg/loop/internal/core/services/capability/capabilities.go index 20ec9f1485..e05d2aa289 100644 --- a/pkg/loop/internal/core/services/capability/capabilities.go +++ b/pkg/loop/internal/core/services/capability/capabilities.go @@ -226,8 +226,7 @@ func (t *triggerExecutableServer) RegisterTrigger(request *pb.TriggerRegistratio // If it's a capability error, serialize it and send it to the client for proper deserialization and handling on the client side. errorString := err.Error() - var capErr caperrors.Error - if errors.As(err, &capErr) { + if capErr, ok := errors.AsType[caperrors.Error](err); ok { errorString = capErr.SerializeToString() } msg := &pb.TriggerResponseMessage{ @@ -448,8 +447,7 @@ func (c *executableServer) Execute(reqpb *pb.CapabilityRequest, server pb.Execut var responseMessage *pb.CapabilityResponse response, err := c.impl.Execute(server.Context(), req) if err != nil { - var capabilityError caperrors.Error - if errors.As(err, &capabilityError) { + if capabilityError, ok := errors.AsType[caperrors.Error](err); ok { responseMessage = &pb.CapabilityResponse{Error: capabilityError.SerializeToString()} } else { // All other errors are treated as private visibility and are marked as such to prevent accidental or malicious diff --git a/pkg/loop/internal/core/services/capability/capabilities_test.go b/pkg/loop/internal/core/services/capability/capabilities_test.go index 4a84c25fb2..d7e6085acc 100644 --- a/pkg/loop/internal/core/services/capability/capabilities_test.go +++ b/pkg/loop/internal/core/services/capability/capabilities_test.go @@ -689,8 +689,7 @@ func Test_Capabilities(t *testing.T) { capabilities.CapabilityRequest{Config: cmap, Inputs: imap}) require.Error(t, err) - var capErr caperrors.Error - ok := errors.As(err, &capErr) + capErr, ok := errors.AsType[caperrors.Error](err) require.True(t, ok, "expected caperrors.Error, got %T: %v", err, err) require.Equal(t, caperrors.Unavailable, capErr.Code()) require.Equal(t, caperrors.VisibilityPublic, capErr.Visibility()) @@ -738,8 +737,7 @@ func Test_Capabilities(t *testing.T) { } require.Error(t, execErr) - var capErr caperrors.Error - ok := errors.As(execErr, &capErr) + capErr, ok := errors.AsType[caperrors.Error](execErr) require.True(t, ok, "expected caperrors.Error, got %T: %v", execErr, execErr) require.Equal(t, caperrors.Unavailable, capErr.Code()) require.Equal(t, caperrors.VisibilityPublic, capErr.Visibility()) @@ -776,8 +774,7 @@ func Test_Capabilities(t *testing.T) { capabilities.CapabilityRequest{Config: cmap, Inputs: imap}) require.Error(t, err) - var capErr caperrors.Error - ok := errors.As(err, &capErr) + capErr, ok := errors.AsType[caperrors.Error](err) require.True(t, ok, "expected caperrors.Error, got %T: %v", err, err) require.Equal(t, caperrors.Internal, capErr.Code()) require.Equal(t, caperrors.VisibilityPublic, capErr.Visibility()) diff --git a/pkg/loop/internal/relayer/pluginprovider/ext/median/test/median.go b/pkg/loop/internal/relayer/pluginprovider/ext/median/test/median.go index 0a5413785c..0b54779724 100644 --- a/pkg/loop/internal/relayer/pluginprovider/ext/median/test/median.go +++ b/pkg/loop/internal/relayer/pluginprovider/ext/median/test/median.go @@ -142,8 +142,7 @@ func (s staticMedianFactoryServer) NewMedianFactory(ctx context.Context, provide err = s.gasPriceSubunitsDataSource.Evaluate(ctx, gasPriceSubunitsDataSource) if err != nil { - var compareError *CompareError - isCompareError := errors.As(err, &compareError) + compareError, isCompareError := errors.AsType[*CompareError](err) // allow 0 as valid data source value with the same staticMedianFactoryServer (because it is only defined once as a global var for all tests) if !isCompareError || !compareError.GotZero() { return nil, fmt.Errorf("NewMedianFactory: gasPriceSubunitsDataSource does not equal a static gas price subunits data source implementation: %w", err) diff --git a/pkg/loop/internal/relayer/test/relayer.go b/pkg/loop/internal/relayer/test/relayer.go index 7e320f235d..2688e77a82 100644 --- a/pkg/loop/internal/relayer/test/relayer.go +++ b/pkg/loop/internal/relayer/test/relayer.go @@ -422,8 +422,21 @@ func (s staticRelayer) Replay(ctx context.Context, fromBlock string, args map[st } func (s staticRelayer) AssertEqual(_ context.Context, t *testing.T, relayer looptypes.Relayer) { + s.assertEqual(t, relayer, true) +} + +// assertEqual checks that the relayer exposes all expected functionality. When parallel is true, +// subtests run concurrently; when false, they run sequentially so callers can ensure the underlying +// service has recovered between assertions. +func (s staticRelayer) assertEqual(t *testing.T, relayer looptypes.Relayer, parallel bool) { + maybeParallel := func(t *testing.T) { + if parallel { + t.Parallel() + } + } + t.Run("ContractReader", func(t *testing.T) { - //t.Parallel() + // ContractReader is intentionally not parallel so it is created before the other subtests. ctx := t.Context() contractReader, err := relayer.NewContractReader(ctx, []byte("test")) require.NoError(t, err) @@ -431,7 +444,7 @@ func (s staticRelayer) AssertEqual(_ context.Context, t *testing.T, relayer loop }) t.Run("ConfigProvider", func(t *testing.T) { - t.Parallel() + maybeParallel(t) ctx := t.Context() configProvider, err := relayer.NewConfigProvider(ctx, RelayArgs) require.NoError(t, err) @@ -441,7 +454,7 @@ func (s staticRelayer) AssertEqual(_ context.Context, t *testing.T, relayer loop }) t.Run("MedianProvider", func(t *testing.T) { - t.Parallel() + maybeParallel(t) ctx := t.Context() ra := newRelayArgsWithProviderType(types.Median) p, err := relayer.NewPluginProvider(ctx, ra, PluginArgs) @@ -451,25 +464,25 @@ func (s staticRelayer) AssertEqual(_ context.Context, t *testing.T, relayer loop servicetest.Run(t, provider) t.Run("ReportingPluginProvider", func(t *testing.T) { - t.Parallel() + maybeParallel(t) s.medianProvider.AssertEqual(ctx, t, provider) }) }) t.Run("PluginProvider", func(t *testing.T) { - t.Parallel() + maybeParallel(t) ra := newRelayArgsWithProviderType(types.GenericPlugin) provider, err := relayer.NewPluginProvider(t.Context(), ra, PluginArgs) require.NoError(t, err) servicetest.Run(t, provider) t.Run("ReportingPluginProvider", func(t *testing.T) { - t.Parallel() + maybeParallel(t) s.agnosticProvider.AssertEqual(t.Context(), t, provider) }) }) t.Run("GetChainStatus", func(t *testing.T) { - t.Parallel() + maybeParallel(t) ctx := t.Context() gotChain, err := relayer.GetChainStatus(ctx) require.NoError(t, err) @@ -477,7 +490,7 @@ func (s staticRelayer) AssertEqual(_ context.Context, t *testing.T, relayer loop }) t.Run("GetChainInfo", func(t *testing.T) { - t.Parallel() + maybeParallel(t) ctx := t.Context() chainInfoReply, err := relayer.GetChainInfo(ctx) require.NoError(t, err) @@ -485,7 +498,7 @@ func (s staticRelayer) AssertEqual(_ context.Context, t *testing.T, relayer loop }) t.Run("ListNodeStatuses", func(t *testing.T) { - t.Parallel() + maybeParallel(t) ctx := t.Context() gotNodes, gotNextToken, gotCount, err := relayer.ListNodeStatuses(ctx, s.nodeRequest.pageSize, s.nodeRequest.pageToken) require.NoError(t, err) @@ -495,14 +508,14 @@ func (s staticRelayer) AssertEqual(_ context.Context, t *testing.T, relayer loop }) t.Run("Transact", func(t *testing.T) { - t.Parallel() + maybeParallel(t) ctx := t.Context() err := relayer.Transact(ctx, s.transactionRequest.from, s.transactionRequest.to, s.transactionRequest.amount, s.transactionRequest.balanceCheck) require.NoError(t, err) }) t.Run("Replay", func(t *testing.T) { - t.Parallel() + maybeParallel(t) ctx := t.Context() err := relayer.Replay(ctx, s.replayRequest.fromBlock, s.replayRequest.args) require.NoError(t, err) @@ -565,6 +578,13 @@ func Run(t *testing.T, relayer looptypes.Relayer) { expectedRelayer.AssertEqual(ctx, t, relayer) } +// RunSequential is like Run but executes the subtests sequentially. This is useful when the relayer +// is expected to recover from plugin failures between assertions, so each assertion has a chance to +// wait for the service to relaunch before the next one runs. +func RunSequential(t *testing.T, relayer looptypes.Relayer) { + newStaticRelayer(logger.Test(t), false).assertEqual(t, relayer, false) +} + func RunFuzzPluginRelayer(f *testing.F, relayerFunc func(*testing.T) looptypes.PluginRelayer) { var ( account = "testaccount" diff --git a/pkg/loop/relayer_service_test.go b/pkg/loop/relayer_service_test.go index 001c7f6191..d438005c8b 100644 --- a/pkg/loop/relayer_service_test.go +++ b/pkg/loop/relayer_service_test.go @@ -75,7 +75,6 @@ func TestRelayerService(t *testing.T) { } func TestRelayerService_recovery(t *testing.T) { - t.Parallel() var limit atomic.Int32 relayer := loop.NewRelayerService(logger.Test(t), loop.GRPCOpts{}, func() *exec.Cmd { return HelperProcessCommand{ @@ -85,7 +84,7 @@ func TestRelayerService_recovery(t *testing.T) { }, test.ConfigTOML, keystoretest.Keystore, keystoretest.Keystore, nil) servicetest.Run(t, relayer) - relayertest.Run(t, relayer) + relayertest.RunSequential(t, relayer) if hp := relayer.HealthReport(); len(hp) == 2 { servicetest.AssertHealthReportNames(t, hp, relayerServiceNames[:2]...) diff --git a/pkg/monitoring/schema_registry.go b/pkg/monitoring/schema_registry.go index a7856fbfb7..00cffe76d2 100644 --- a/pkg/monitoring/schema_registry.go +++ b/pkg/monitoring/schema_registry.go @@ -68,8 +68,7 @@ func isNotFoundErr(err error) bool { if strings.HasPrefix(err.Error(), "Subject not found") { // for mock schema registry return true } - var srErr srclient.Error - if errors.As(err, &srErr) && srErr.Code == 40401 { // for the actual schema registry api. + if srErr, ok := errors.AsType[srclient.Error](err); ok && srErr.Code == 40401 { // for the actual schema registry api. return true } return false diff --git a/pkg/settings/limits/errors.go b/pkg/settings/limits/errors.go index b0ab536731..0b9c531175 100644 --- a/pkg/settings/limits/errors.go +++ b/pkg/settings/limits/errors.go @@ -12,10 +12,9 @@ import ( ) // LimitError is implemented by errors returned when a limit is exceeded. -// Use [errors.As] to identify limit errors, for example: +// Use [errors.AsType] to identify limit errors, for example: // -// var limitErr LimitError -// if errors.As(err, &limitErr) { ... } +// if limitErr, ok := errors.AsType[LimitError](err); ok { ... } type LimitError interface { error limitError() diff --git a/pkg/settings/limits/errors_test.go b/pkg/settings/limits/errors_test.go index 589b9f776a..739678a436 100644 --- a/pkg/settings/limits/errors_test.go +++ b/pkg/settings/limits/errors_test.go @@ -31,15 +31,18 @@ func TestLimitError_As(t *testing.T) { for _, err := range cases { t.Run(err.Error(), func(t *testing.T) { t.Parallel() - var limitErr LimitError - require.True(t, errors.As(err, &limitErr)) - require.True(t, errors.As(fmt.Errorf("wrapped: %w", err), &limitErr)) + limitErr, ok := errors.AsType[LimitError](err) + require.True(t, ok) + _ = limitErr + _, ok = errors.AsType[LimitError](fmt.Errorf("wrapped: %w", err)) + require.True(t, ok) }) } - var limitErr LimitError - require.False(t, errors.As(errors.New("other"), &limitErr)) - require.False(t, errors.As(ErrQueueEmpty, &limitErr)) + _, ok := errors.AsType[LimitError](errors.New("other")) + require.False(t, ok) + _, ok = errors.AsType[LimitError](ErrQueueEmpty) + require.False(t, ok) } func TestErrorRateLimited(t *testing.T) { diff --git a/pkg/workflows/host/execution_restrictions_test.go b/pkg/workflows/host/execution_restrictions_test.go index 125c3ace06..4679d02312 100644 --- a/pkg/workflows/host/execution_restrictions_test.go +++ b/pkg/workflows/host/execution_restrictions_test.go @@ -88,8 +88,8 @@ func TestRequirementSelectingModule_CallCapWithRestrictions(t *testing.T) { inner := mocks.NewMockExecutionHelper(t) // no expectations: inner must not be called h := host.NewRestrictedExecutionHelper(inner, restrictions) _, err := h.CallCapability(t.Context(), &sdk.CapabilityRequest{Id: "blocked@1.0.0", Method: "Bar"}) - var capErr caperrors.Error - require.True(t, errors.As(err, &capErr)) + capErr, ok := errors.AsType[caperrors.Error](err) + require.True(t, ok) assert.Contains(t, capErr.Error(), "denied by user pre-hook restrictions") assert.Equal(t, caperrors.LimitExceeded, capErr.Code()) }) @@ -440,8 +440,8 @@ func TestRequirementSelectingModule_ConfidentialHTTPWithRestrictions(t *testing. req := confidentialHTTPRequest(t, "confhttp@1.0.0", "Call", &confidentialhttp.SecretIdentifier{Key: "blocked-secret", Namespace: "ns"}) _, err := h.CallCapability(t.Context(), req) - var capErr caperrors.Error - require.True(t, errors.As(err, &capErr)) + capErr, ok := errors.AsType[caperrors.Error](err) + require.True(t, ok) assert.Contains(t, capErr.Error(), "denied by user pre-hook restrictions") assert.Equal(t, caperrors.LimitExceeded, capErr.Code()) }) diff --git a/pkg/workflows/wasm/host/execution.go b/pkg/workflows/wasm/host/execution.go index 53caaf3774..a24bf4a9c7 100644 --- a/pkg/workflows/wasm/host/execution.go +++ b/pkg/workflows/wasm/host/execution.go @@ -62,8 +62,7 @@ func (e *execution[T]) callCapAsync(ctx context.Context, req *sdkpb.CapabilityRe if err != nil { errString := err.Error() - var caperror caperrors.Error - if errors.As(err, &caperror) { + if caperror, ok := errors.AsType[caperrors.Error](err); ok { errString = caperror.SerializeToString() } resp = &sdkpb.CapabilityResponse{