Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 1 addition & 2 deletions pkg/beholder/batch_emitter_service.go
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
3 changes: 1 addition & 2 deletions pkg/chipingress/batch/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
}
Expand Down
3 changes: 1 addition & 2 deletions pkg/config/toml.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
6 changes: 2 additions & 4 deletions pkg/loop/internal/core/services/capability/capabilities.go
Original file line number Diff line number Diff line change
Expand Up @@ -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{
Expand Down Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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())
Expand Down Expand Up @@ -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())
Expand Down Expand Up @@ -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())
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
42 changes: 31 additions & 11 deletions pkg/loop/internal/relayer/test/relayer.go
Original file line number Diff line number Diff line change
Expand Up @@ -422,16 +422,29 @@ 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)
}
Comment on lines 424 to +426

// 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)
servicetest.Run(t, contractReader)
})

t.Run("ConfigProvider", func(t *testing.T) {
t.Parallel()
maybeParallel(t)
ctx := t.Context()
configProvider, err := relayer.NewConfigProvider(ctx, RelayArgs)
require.NoError(t, err)
Expand All @@ -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)
Expand All @@ -451,41 +464,41 @@ 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)
assert.Equal(t, s.chainStatus, gotChain)
})

t.Run("GetChainInfo", func(t *testing.T) {
t.Parallel()
maybeParallel(t)
ctx := t.Context()
chainInfoReply, err := relayer.GetChainInfo(ctx)
require.NoError(t, err)
assert.Equal(t, s.chainInfo, chainInfoReply)
})

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)
Expand All @@ -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)
Expand Down Expand Up @@ -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"
Expand Down
3 changes: 1 addition & 2 deletions pkg/loop/relayer_service_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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{
Expand All @@ -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)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

cc @pkcll


if hp := relayer.HealthReport(); len(hp) == 2 {
servicetest.AssertHealthReportNames(t, hp, relayerServiceNames[:2]...)
Expand Down
3 changes: 1 addition & 2 deletions pkg/monitoring/schema_registry.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
5 changes: 2 additions & 3 deletions pkg/settings/limits/errors.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
15 changes: 9 additions & 6 deletions pkg/settings/limits/errors_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down
8 changes: 4 additions & 4 deletions pkg/workflows/host/execution_restrictions_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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())
})
Expand Down Expand Up @@ -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())
})
Expand Down
3 changes: 1 addition & 2 deletions pkg/workflows/wasm/host/execution.go
Original file line number Diff line number Diff line change
Expand Up @@ -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{
Expand Down
Loading