diff --git a/pkg/chain/chain.go b/pkg/chain/chain.go index b04ce3b..e9fadc2 100644 --- a/pkg/chain/chain.go +++ b/pkg/chain/chain.go @@ -3,6 +3,7 @@ package chain import ( "errors" "fmt" + "strconv" "strings" ) @@ -75,3 +76,20 @@ func ChainFromInt(chainID uint64) (Chain, error) { } return "", fmt.Errorf("unknown chain ID: %d", chainID) } + +// ParseChainIDParam parses a chain_id query param. Empty means 0 ("all chains"). +// A literal 0 is rejected unless allowZero is true. +func ParseChainIDParam(raw string, allowZero bool) (uint64, error) { + s := strings.TrimSpace(raw) + if s == "" { + return 0, nil + } + id, err := strconv.ParseUint(s, 10, 64) + if err != nil { + return 0, fmt.Errorf("invalid chain_id %q: %w", raw, err) + } + if id == 0 && !allowZero { + return 0, errors.New("chain_id must be a non-zero uint64") + } + return id, nil +} diff --git a/pkg/chain/chain_test.go b/pkg/chain/chain_test.go index 3f39321..91e2b3a 100644 --- a/pkg/chain/chain_test.go +++ b/pkg/chain/chain_test.go @@ -70,3 +70,42 @@ func TestGenesisTime(t *testing.T) { _, ok := chain.GenesisTime("unknown") require.False(t, ok) } + +func TestParseChainIDParam(t *testing.T) { + cases := []struct { + raw string + allowZero bool + want uint64 + wantErr bool + }{ + {"", false, 0, false}, + {" ", false, 0, false}, + {"1", false, 1, false}, + {" 1 ", false, 1, false}, + {"18446744073709551615", false, 18446744073709551615, false}, + {"0", false, 0, true}, + {"00", false, 0, true}, + {"-1", false, 0, true}, + {"+1", false, 0, true}, + {"abc", false, 0, true}, + {"1.0", false, 0, true}, + {"18446744073709551616", false, 0, true}, + {"", true, 0, false}, + {"0", true, 0, false}, + {"00", true, 0, false}, + {"1", true, 1, false}, + {"-1", true, 0, true}, + {"abc", true, 0, true}, + {"18446744073709551616", true, 0, true}, + } + for _, c := range cases { + id, err := chain.ParseChainIDParam(c.raw, c.allowZero) + if c.wantErr { + require.Error(t, err) + require.Equal(t, uint64(0), id) + continue + } + require.NoError(t, err) + require.Equal(t, c.want, id) + } +}