diff --git a/expr_json.go b/expr_json.go index a81b5f4bd..91cd8ed4e 100644 --- a/expr_json.go +++ b/expr_json.go @@ -21,7 +21,9 @@ import ( "bytes" "encoding/hex" "encoding/json" + "errors" "fmt" + "io" "math" "sort" "strings" @@ -488,8 +490,21 @@ func decodeExpr(raw json.RawMessage, schema *Schema, caseSensitive bool) (Boolea // without inspecting bytes by hand. func firstToken(raw json.RawMessage) (json.Token, error) { dec := json.NewDecoder(bytes.NewReader(raw)) + tok, err := dec.Token() + if err != nil { + return nil, err + } + if _, ok := tok.(bool); ok { + if _, err := dec.Token(); !errors.Is(err, io.EOF) { + if err == nil { + return nil, errors.New("trailing data after boolean expression") + } + + return nil, err + } + } - return dec.Token() + return tok, nil } func decodePredicate(op Operation, node exprNode, schema *Schema, caseSensitive bool) (BooleanExpression, error) { diff --git a/expr_json_test.go b/expr_json_test.go index b5841a4b8..1b9de31d6 100644 --- a/expr_json_test.go +++ b/expr_json_test.go @@ -393,6 +393,20 @@ func TestUnmarshalExpressionErrors(t *testing.T) { } } +func TestUnmarshalBooleanExpressionRejectsTrailingData(t *testing.T) { + for _, input := range []string{ + `true false`, + `false null`, + `true{}`, + `false garbage`, + } { + t.Run(input, func(t *testing.T) { + _, err := iceberg.ParseExpr([]byte(input), nil) + require.ErrorIs(t, err, iceberg.ErrInvalidArgument) + }) + } +} + // TestExpressionTransformTermRoundTrip covers a residual filter whose term is a // partition transform, e.g. a Java server's bucket[100](id) <= 50. func TestExpressionTransformTermRoundTrip(t *testing.T) {