// Package wire enforces the bounded, exact JSON protocol shared by CLI and packages. package wire import ( "bytes" "encoding/json" "errors" "io" "reflect" "time" "unicode/utf8" ) // Decode refuses duplicate keys, case aliases, unknown/missing fields, // nulls and trailing values. encoding/json alone accepts several of these. func Decode(in io.Reader, target any, limit int64) error { if limit <= 0 || limit > 16<<20 { return errors.New("invalid input limit") } typ := reflect.TypeOf(target) if typ == nil || typ.Kind() != reflect.Pointer || reflect.ValueOf(target).IsNil() { return errors.New("expected nonnil decode target") } data, err := io.ReadAll(io.LimitReader(in, limit+1)) if err != nil || int64(len(data)) > limit || !utf8.Valid(data) { return errors.New("invalid input") } decoder := json.NewDecoder(bytes.NewReader(data)) if err := uniqueValue(decoder, 0); err != nil { return err } if _, err := decoder.Token(); err != io.EOF { return errors.New("trailing input") } if err := exactFields(data, typ.Elem()); err != nil { return err } return json.Unmarshal(data, target) } func uniqueValue(d *json.Decoder, depth int) error { if depth > 32 { return errors.New("input nesting exceeds limit") } token, err := d.Token() if err != nil { return err } delim, composite := token.(json.Delim) if !composite { return nil } switch delim { case '{': seen := make(map[string]bool) for d.More() { token, err := d.Token() if err != nil { return err } key, ok := token.(string) if !ok || seen[key] { return errors.New("duplicate or invalid field") } seen[key] = true if err := uniqueValue(d, depth+1); err != nil { return err } } case '[': for d.More() { if err := uniqueValue(d, depth+1); err != nil { return err } } default: return errors.New("unexpected delimiter") } _, err = d.Token() return err } func exactFields(data []byte, typ reflect.Type) error { if bytes.Equal(bytes.TrimSpace(data), []byte("null")) { return errors.New("null is not permitted") } if typ == reflect.TypeOf(time.Time{}) { var text string if err := json.Unmarshal(data, &text); err != nil { return err } parsed, err := time.Parse(time.RFC3339, text) if err != nil || text != parsed.UTC().Format("2006-01-02T15:04:05Z") { return errors.New("expected canonical UTC timestamp") } return nil } if typ.Kind() == reflect.Slice { var items []json.RawMessage if err := json.Unmarshal(data, &items); err != nil { return err } for _, item := range items { if err := exactFields(item, typ.Elem()); err != nil { return err } } return nil } if typ.Kind() == reflect.Map { var entries map[string]json.RawMessage if err := json.Unmarshal(data, &entries); err != nil { return err } for _, value := range entries { if err := exactFields(value, typ.Elem()); err != nil { return err } } return nil } if typ.Kind() != reflect.Struct { return nil } var fields map[string]json.RawMessage if err := json.Unmarshal(data, &fields); err != nil { return err } if len(fields) != typ.NumField() { return errors.New("unexpected field set") } for n := 0; n < typ.NumField(); n++ { field := typ.Field(n) value, exists := fields[field.Tag.Get("json")] if !exists { return errors.New("missing field") } if err := exactFields(value, field.Type); err != nil { return err } } return nil }