diff --git a/Makefile b/Makefile index 6192ad8..420f306 100644 --- a/Makefile +++ b/Makefile @@ -23,7 +23,7 @@ generate: root build .root/src/$(PKG)/tests/omitempty.go .root/bin/easyjson -all .root/src/$(PKG)/tests/data.go - .root/bin/easyjson -snake_case .root/src/$(PKG)/tests/snake.go + .root/bin/easyjson -snake_case .root/src/$(PKG)/tests/snake.go .root/bin/easyjson -omit_empty .root/src/$(PKG)/tests/omitempty.go .root/bin/easyjson -build_tags=use_easyjson .root/src/$(PKG)/benchmark/data.go diff --git a/gen/decoder.go b/gen/decoder.go index a662329..a1291fc 100644 --- a/gen/decoder.go +++ b/gen/decoder.go @@ -170,7 +170,40 @@ func (g *Generator) genStructFieldDecoder(t reflect.Type, f reflect.StructField) } fmt.Fprintf(g.out, " case %q:\n", jsonName) - return g.genTypeDecoder(f.Type, "out."+f.Name, tags, 3) + if err := g.genTypeDecoder(f.Type, "out."+f.Name, tags, 3); err != nil { + return err + } + + if tags.required { + fmt.Fprintf(g.out, "%sSet = true\n", f.Name) + } + + return nil +} + +func (g *Generator) genRequiredFieldSet(t reflect.Type, f reflect.StructField) { + tags := parseFieldTags(f) + + if !tags.required { + return + } + + fmt.Fprintf(g.out, "var %sSet bool\n", f.Name) +} + +func (g *Generator) genRequiredFieldCheck(t reflect.Type, f reflect.StructField) { + jsonName := g.namer.GetJSONFieldName(t, f) + tags := parseFieldTags(f) + + if !tags.required { + return + } + + g.imports["fmt"] = "fmt" + + fmt.Fprintf(g.out, "if !%sSet {\n", f.Name) + fmt.Fprintf(g.out, " in.AddError(fmt.Errorf(\"key '%s' is required\"))\n", jsonName) + fmt.Fprintf(g.out, "}\n") } func mergeStructFields(fields1, fields2 []reflect.StructField) (fields []reflect.StructField) { @@ -250,6 +283,15 @@ func (g *Generator) genStructDecoder(t reflect.Type) error { fmt.Fprintln(g.out, " out."+f.Name+" = new("+g.getType(f.Type.Elem())+")") } + fs, err := getStructFields(t) + if err != nil { + return fmt.Errorf("cannot generate decoder for %v: %v", t, err) + } + + for _, f := range fs { + g.genRequiredFieldSet(t, f) + } + fmt.Fprintln(g.out, " in.Delim('{')") fmt.Fprintln(g.out, " for !in.IsDelim('}') {") fmt.Fprintln(g.out, " key := in.UnsafeString()") @@ -259,13 +301,8 @@ func (g *Generator) genStructDecoder(t reflect.Type) error { fmt.Fprintln(g.out, " in.WantComma()") fmt.Fprintln(g.out, " continue") fmt.Fprintln(g.out, " }") + fmt.Fprintln(g.out, " switch key {") - - fs, err := getStructFields(t) - if err != nil { - return fmt.Errorf("cannot generate decoder for %v: %v", t, err) - } - for _, f := range fs { if err := g.genStructFieldDecoder(t, f); err != nil { return err @@ -278,6 +315,11 @@ func (g *Generator) genStructDecoder(t reflect.Type) error { fmt.Fprintln(g.out, " in.WantComma()") fmt.Fprintln(g.out, " }") fmt.Fprintln(g.out, " in.Delim('}')") + + for _, f := range fs { + g.genRequiredFieldCheck(t, f) + } + fmt.Fprintln(g.out, "}") return nil diff --git a/gen/encoder.go b/gen/encoder.go index dfe21f0..c263c21 100644 --- a/gen/encoder.go +++ b/gen/encoder.go @@ -52,6 +52,7 @@ type fieldTags struct { omitEmpty bool noOmitEmpty bool asString bool + required bool } // parseFieldTags parses the json field tag into a structure. @@ -70,6 +71,8 @@ func parseFieldTags(f reflect.StructField) fieldTags { ret.noOmitEmpty = true case s == "string": ret.asString = true + case s == "required": + ret.required = true } } diff --git a/tests/data.go b/tests/data.go index 73ced30..aa8963f 100644 --- a/tests/data.go +++ b/tests/data.go @@ -435,3 +435,8 @@ var mapsString = `{` + `"NilMap":null,` + `"CustomMap":{"c":"d"}` + `}` + +type RequiredOptionalStruct struct { + FirstName string `json:"first_name,required"` + Lastname string `json:"last_name"` +} diff --git a/tests/required_test.go b/tests/required_test.go new file mode 100644 index 0000000..8b03be6 --- /dev/null +++ b/tests/required_test.go @@ -0,0 +1,28 @@ +package tests + +import ( + "testing" + "fmt" +) + +func TestRequiredField(t *testing.T) { + cases := []struct{ json, errorMessage string }{ + {`{"first_name":"Foo", "last_name": "Bar"}`, ""}, + {`{"last_name":"Bar"}`, "key 'first_name' is required"}, + {"{}", "key 'first_name' is required"}, + } + + for _, tc := range cases { + var v RequiredOptionalStruct + err := v.UnmarshalJSON([]byte(tc.json)) + if tc.errorMessage == "" { + if err != nil { + t.Errorf("%s. UnmarshallJSON didn`t expect error: %v", tc.json, err) + } + } else { + if fmt.Sprintf("%v", err) != tc.errorMessage { + t.Errorf("%s. UnmarshallJSON expected error: %v. got: %v", tc.json, tc.errorMessage, err) + } + } + } +}