diff --git a/fs/rc/internal.go b/fs/rc/internal.go index 36d2fafd5..c1bbe7819 100644 --- a/fs/rc/internal.go +++ b/fs/rc/internal.go @@ -305,15 +305,14 @@ func init() { // Terminates app func rcQuit(ctx context.Context, in Params) (out Params, err error) { - code, err := in.GetInt64("exitCode") + exitCode, err := in.GetInt("exitCode") if IsErrParamInvalid(err) { return nil, err } if IsErrParamNotFound(err) { - code = 0 + exitCode = 0 } - exitCode := int(code) go func(exitCode int) { time.Sleep(time.Millisecond * 1500) @@ -353,11 +352,11 @@ Results: } func rcSetMutexProfileFraction(ctx context.Context, in Params) (out Params, err error) { - rate, err := in.GetInt64("rate") + rate, err := in.GetInt("rate") if err != nil { return nil, err } - previousRate := runtime.SetMutexProfileFraction(int(rate)) + previousRate := runtime.SetMutexProfileFraction(rate) out = make(Params) out["previousRate"] = previousRate return out, nil @@ -388,11 +387,11 @@ Parameters: } func rcSetBlockProfileRate(ctx context.Context, in Params) (out Params, err error) { - rate, err := in.GetInt64("rate") + rate, err := in.GetInt("rate") if err != nil { return nil, err } - runtime.SetBlockProfileRate(int(rate)) + runtime.SetBlockProfileRate(rate) return nil, nil } @@ -471,11 +470,11 @@ Parameters: } func rcSetGCPercent(ctx context.Context, in Params) (out Params, err error) { - gcPercent, err := in.GetInt64("gc-percent") + gcPercent, err := in.GetInt("gc-percent") if err != nil { return nil, err } - oldGCPercent := debug.SetGCPercent(int(gcPercent)) + oldGCPercent := debug.SetGCPercent(gcPercent) out = Params{ "existing-gc-percent": oldGCPercent, } diff --git a/fs/rc/jobs/job.go b/fs/rc/jobs/job.go index b4ef2acca..0390c3ad7 100644 --- a/fs/rc/jobs/job.go +++ b/fs/rc/jobs/job.go @@ -689,10 +689,10 @@ func rcBatch(ctx context.Context, in rc.Params) (out rc.Params, err error) { } // Read concurrency - concurrency, err := in.GetInt64("concurrency") + concurrency, err := in.GetInt("concurrency") if rc.IsErrParamNotFound(err) { ci := fs.GetConfig(ctx) - concurrency = int64(ci.Transfers) + concurrency = ci.Transfers } else if err != nil { return nil, err } @@ -702,7 +702,7 @@ func rcBatch(ctx context.Context, in rc.Params) (out rc.Params, err error) { out["results"] = results g, gCtx := errgroup.WithContext(ctx) - g.SetLimit(int(concurrency)) + g.SetLimit(concurrency) for i, inputAny := range inputs { input, ok := inputAny.(map[string]any) if !ok { diff --git a/fs/rc/params.go b/fs/rc/params.go index 402be3c7b..994d543ca 100644 --- a/fs/rc/params.go +++ b/fs/rc/params.go @@ -182,6 +182,22 @@ func (p Params) GetInt64(key string) (int64, error) { return 0, ErrParamInvalid{fmt.Errorf("expecting int64 value for key %q (was %T)", key, value)} } +// GetInt gets an int parameter from the input +// +// If the parameter isn't found then error will be of type +// ErrParamNotFound and the returned value will be 0. If the value +// doesn't fit in an int then the error will be of type ErrParamInvalid. +func (p Params) GetInt(key string) (int, error) { + i, err := p.GetInt64(key) + if err != nil { + return 0, err + } + if i > math.MaxInt || i < math.MinInt { + return 0, ErrParamInvalid{fmt.Errorf("key %q (%v) overflows int", key, i)} + } + return int(i), nil +} + // GetFloat64 gets a float64 parameter from the input // // If the parameter isn't found then error will be of type diff --git a/fs/rc/params_test.go b/fs/rc/params_test.go index 6714a5a3f..b6ad6f20d 100644 --- a/fs/rc/params_test.go +++ b/fs/rc/params_test.go @@ -158,6 +158,41 @@ func TestParamsGetInt64(t *testing.T) { assert.Equal(t, true, IsErrParamInvalid(e3), e3.Error()) } +func TestParamsGetInt(t *testing.T) { + in := Params{ + "int": "123", + "bad": "123x", + "notInt": []string{"a", "b"}, + "maxInt": int64(math.MaxInt), + "minInt": int64(math.MinInt), + "overflow": "9223372036854775808", + } + v, err := in.GetInt("int") + require.NoError(t, err) + assert.Equal(t, 123, v) + v, err = in.GetInt("maxInt") + require.NoError(t, err) + assert.Equal(t, math.MaxInt, v) + v, err = in.GetInt("minInt") + require.NoError(t, err) + assert.Equal(t, math.MinInt, v) + for _, key := range []string{"bad", "notInt", "overflow"} { + v, err = in.GetInt(key) + assert.True(t, IsErrParamInvalid(err), key) + assert.Equal(t, 0, v) + } + v, err = in.GetInt("notFound") + assert.Equal(t, ErrParamNotFound("notFound"), err) + assert.Equal(t, 0, v) + if math.MaxInt == math.MaxInt32 { + in["big"] = int64(math.MaxInt32) + 1 + v, err = in.GetInt("big") + assert.True(t, IsErrParamInvalid(err)) + assert.Contains(t, err.Error(), "overflows int") + assert.Equal(t, 0, v) + } +} + func TestParamsGetFloat64(t *testing.T) { for _, test := range []struct { value any