diff --git a/Makefile b/Makefile index ce30363..57427dd 100644 --- a/Makefile +++ b/Makefile @@ -20,7 +20,9 @@ generate: build ./tests/reference_to_pointer.go \ ./tests/html.go \ ./tests/unknown_fields.go \ - ./tests/type_declaration.go + ./tests/type_declaration.go \ + ./tests/members_escaped.go \ + ./tests/members_unescaped.go \ bin/easyjson -all ./tests/data.go bin/easyjson -all ./tests/nothing.go @@ -28,7 +30,7 @@ generate: build bin/easyjson -all ./tests/html.go bin/easyjson -snake_case ./tests/snake.go bin/easyjson -omit_empty ./tests/omitempty.go - bin/easyjson -build_tags=use_easyjson ./benchmark/data.go + bin/easyjson -build_tags=use_easyjson -disable_members_unescape ./benchmark/data.go bin/easyjson ./tests/nested_easy.go bin/easyjson ./tests/named_type.go bin/easyjson ./tests/custom_map_key_type.go @@ -38,6 +40,8 @@ generate: build bin/easyjson -disallow_unknown_fields ./tests/disallow_unknown.go bin/easyjson ./tests/unknown_fields.go bin/easyjson ./tests/type_declaration.go + bin/easyjson ./tests/members_escaped.go + bin/easyjson -disable_members_unescape ./tests/members_unescaped.go test: generate go test \ diff --git a/README.md b/README.md index c88848d..92d9f1f 100644 --- a/README.md +++ b/README.md @@ -57,6 +57,8 @@ Usage of easyjson: only generate stubs for marshaler/unmarshaler funcs -disallow_unknown_fields return error if some unknown field in json appeared + -disable_members_unescape + disable unescaping of \uXXXX string sequences in member names ``` Using `-all` will generate marshalers/unmarshalers for all Go structs in the diff --git a/bootstrap/bootstrap.go b/bootstrap/bootstrap.go index 9d1c2e3..7e984d9 100644 --- a/bootstrap/bootstrap.go +++ b/bootstrap/bootstrap.go @@ -23,18 +23,19 @@ type Generator struct { PkgPath, PkgName string Types []string - NoStdMarshalers bool - SnakeCase bool - LowerCamelCase bool - OmitEmpty bool - DisallowUnknownFields bool + NoStdMarshalers bool + SnakeCase bool + LowerCamelCase bool + OmitEmpty bool + DisallowUnknownFields bool + SkipMemberNameUnescaping bool OutName string BuildTags string - StubsOnly bool - LeaveTemps bool - NoFormat bool + StubsOnly bool + LeaveTemps bool + NoFormat bool SimpleBytes bool } @@ -129,6 +130,9 @@ func (g *Generator) writeMain() (path string, err error) { if g.SimpleBytes { fmt.Fprintln(f, " g.SimpleBytes()") } + if g.SkipMemberNameUnescaping { + fmt.Fprintln(f, " g.SkipMemberNameUnescaping()") + } sort.Strings(g.Types) for _, v := range g.Types { diff --git a/easyjson/main.go b/easyjson/main.go index 26e95cb..7e935fb 100644 --- a/easyjson/main.go +++ b/easyjson/main.go @@ -29,6 +29,7 @@ var noformat = flag.Bool("noformat", false, "do not run 'gofmt -w' on output fil var specifiedName = flag.String("output_filename", "", "specify the filename of the output") var processPkg = flag.Bool("pkg", false, "process the whole package instead of just the given file") var disallowUnknownFields = flag.Bool("disallow_unknown_fields", false, "return error if any unknown field in json appeared") +var skipMemberNameUnescaping = flag.Bool("disable_members_unescape", false, "don't perform unescaping of member names to improve performance") func generate(fname string) (err error) { fInfo, err := os.Stat(fname) @@ -62,20 +63,21 @@ func generate(fname string) (err error) { } g := bootstrap.Generator{ - BuildTags: trimmedBuildTags, - PkgPath: p.PkgPath, - PkgName: p.PkgName, - Types: p.StructNames, - SnakeCase: *snakeCase, - LowerCamelCase: *lowerCamelCase, - NoStdMarshalers: *noStdMarshalers, - DisallowUnknownFields: *disallowUnknownFields, - OmitEmpty: *omitEmpty, - LeaveTemps: *leaveTemps, - OutName: outName, - StubsOnly: *stubs, - NoFormat: *noformat, - SimpleBytes: *simpleBytes, + BuildTags: trimmedBuildTags, + PkgPath: p.PkgPath, + PkgName: p.PkgName, + Types: p.StructNames, + SnakeCase: *snakeCase, + LowerCamelCase: *lowerCamelCase, + NoStdMarshalers: *noStdMarshalers, + DisallowUnknownFields: *disallowUnknownFields, + SkipMemberNameUnescaping: *skipMemberNameUnescaping, + OmitEmpty: *omitEmpty, + LeaveTemps: *leaveTemps, + OutName: outName, + StubsOnly: *stubs, + NoFormat: *noformat, + SimpleBytes: *simpleBytes, } if err := g.Run(); err != nil { diff --git a/gen/decoder.go b/gen/decoder.go index 2a17cc7..2127d0d 100644 --- a/gen/decoder.go +++ b/gen/decoder.go @@ -486,7 +486,7 @@ func (g *Generator) genStructDecoder(t reflect.Type) error { fmt.Fprintln(g.out, " in.Delim('{')") fmt.Fprintln(g.out, " for !in.IsDelim('}') {") - fmt.Fprintln(g.out, " key := in.UnsafeString()") + fmt.Fprintf(g.out, " key := in.UnsafeFieldName(%v)\n", g.skipMemberNameUnescaping) fmt.Fprintln(g.out, " in.WantColon()") fmt.Fprintln(g.out, " if in.IsNull() {") fmt.Fprintln(g.out, " in.Skip()") diff --git a/gen/generator.go b/gen/generator.go index 344e516..79f4d6f 100644 --- a/gen/generator.go +++ b/gen/generator.go @@ -33,11 +33,12 @@ type Generator struct { varCounter int - noStdMarshalers bool - omitEmpty bool - disallowUnknownFields bool - fieldNamer FieldNamer - simpleBytes bool + noStdMarshalers bool + omitEmpty bool + disallowUnknownFields bool + fieldNamer FieldNamer + simpleBytes bool + skipMemberNameUnescaping bool // package path to local alias map for tracking imports imports map[string]string @@ -117,6 +118,11 @@ func (g *Generator) DisallowUnknownFields() { g.disallowUnknownFields = true } +// SkipMemberNameUnescaping instructs to skip member names unescaping to improve performance +func (g *Generator) SkipMemberNameUnescaping() { + g.skipMemberNameUnescaping = true +} + // OmitEmpty triggers `json=",omitempty"` behaviour by default. func (g *Generator) OmitEmpty() { g.omitEmpty = true diff --git a/jlexer/lexer.go b/jlexer/lexer.go index ddd376b..d6a7f4d 100644 --- a/jlexer/lexer.go +++ b/jlexer/lexer.go @@ -5,6 +5,7 @@ package jlexer import ( + "bytes" "encoding/base64" "encoding/json" "errors" @@ -32,9 +33,10 @@ const ( type token struct { kind tokenKind // Type of a token. - boolValue bool // Value if a boolean literal token. - byteValue []byte // Raw value of a token. - delimValue byte + boolValue bool // Value if a boolean literal token. + byteValueCloned bool // true if byteValue was allocated and does not refer to original json body + byteValue []byte // Raw value of a token. + delimValue byte } // Lexer is a JSON lexer: it iterates over JSON tokens in a byte slice. @@ -240,23 +242,55 @@ func (r *Lexer) fetchNumber() { // findStringLen tries to scan into the string literal for ending quote char to determine required size. // The size will be exact if no escapes are present and may be inexact if there are escaped chars. -func findStringLen(data []byte) (isValid, hasEscapes bool, length int) { - delta := 0 - - for i := 0; i < len(data); i++ { - switch data[i] { - case '\\': - i++ - delta++ - if i < len(data) && data[i] == 'u' { - delta++ - } - case '"': - return true, (delta > 0), (i - delta) +func findStringLen(data []byte) (isValid bool, length int) { + for { + idx := bytes.IndexByte(data, '"') + if idx == -1 { + return false, len(data) } + if idx == 0 || (idx > 0 && data[idx-1] != '\\') { + return true, length + idx + } + length += idx + 1 + data = data[idx+1:] + } +} + +// unescapeStringToken performs unescaping of string token. +// if no escaping is needed, original string is returned, otherwise - a new one allocated +func (r *Lexer) unescapeStringToken() (err error) { + data := r.token.byteValue + var unescapedData []byte + + for { + i := bytes.IndexByte(data, '\\') + if i == -1 { + break + } + + escapedRune, escapedBytes, err := decodeEscape(data[i:]) + if err != nil { + r.errParse(err.Error()) + return err + } + + if unescapedData == nil { + unescapedData = make([]byte, 0, len(r.token.byteValue)) + } + + var d [4]byte + s := utf8.EncodeRune(d[:], escapedRune) + unescapedData = append(unescapedData, data[:i]...) + unescapedData = append(unescapedData, d[:s]...) + + data = data[i+escapedBytes:] } - return false, false, len(data) + if unescapedData != nil { + r.token.byteValue = append(unescapedData, data...) + r.token.byteValueCloned = true + } + return } // getu4 decodes \uXXXX from the beginning of s, returning the hex value, @@ -286,36 +320,30 @@ func getu4(s []byte) rune { return val } -// processEscape processes a single escape sequence and returns number of bytes processed. -func (r *Lexer) processEscape(data []byte) (int, error) { +// decodeEscape processes a single escape sequence and returns number of bytes processed. +func decodeEscape(data []byte) (decoded rune, bytesProcessed int, err error) { if len(data) < 2 { - return 0, fmt.Errorf("syntax error at %v", string(data)) + return 0, 0, fmt.Errorf("syntax error at %v", string(data)) } c := data[1] switch c { case '"', '/', '\\': - r.token.byteValue = append(r.token.byteValue, c) - return 2, nil + return rune(c), 2, nil case 'b': - r.token.byteValue = append(r.token.byteValue, '\b') - return 2, nil + return '\b', 2, nil case 'f': - r.token.byteValue = append(r.token.byteValue, '\f') - return 2, nil + return '\f', 2, nil case 'n': - r.token.byteValue = append(r.token.byteValue, '\n') - return 2, nil + return '\n', 2, nil case 'r': - r.token.byteValue = append(r.token.byteValue, '\r') - return 2, nil + return '\r', 2, nil case 't': - r.token.byteValue = append(r.token.byteValue, '\t') - return 2, nil + return '\t', 2, nil case 'u': rr := getu4(data) if rr < 0 { - return 0, errors.New("syntax error") + return 0, 0, errors.New("syntax error") } read := 6 @@ -328,13 +356,10 @@ func (r *Lexer) processEscape(data []byte) (int, error) { rr = unicode.ReplacementChar } } - var d [4]byte - s := utf8.EncodeRune(d[:], rr) - r.token.byteValue = append(r.token.byteValue, d[:s]...) - return read, nil + return rr, read, nil } - return 0, errors.New("syntax error") + return 0, 0, errors.New("syntax error") } // fetchString scans a string literal token. @@ -342,43 +367,14 @@ func (r *Lexer) fetchString() { r.pos++ data := r.Data[r.pos:] - isValid, hasEscapes, length := findStringLen(data) + isValid, length := findStringLen(data) if !isValid { r.pos += length r.errParse("unterminated string literal") return } - if !hasEscapes { - r.token.byteValue = data[:length] - r.pos += length + 1 - return - } - - r.token.byteValue = make([]byte, 0, length) - p := 0 - for i := 0; i < len(data); { - switch data[i] { - case '"': - r.pos += i + 1 - r.token.byteValue = append(r.token.byteValue, data[p:i]...) - i++ - return - - case '\\': - r.token.byteValue = append(r.token.byteValue, data[p:i]...) - off, err := r.processEscape(data[i:]) - if err != nil { - r.errParse(err.Error()) - return - } - i += off - p = i - - default: - i++ - } - } - r.errParse("unterminated string literal") + r.token.byteValue = data[:length] + r.pos += length + 1 // skip closing '"' as well } // scanToken scans the next token if no token is currently available in the lexer. @@ -602,7 +598,7 @@ func (r *Lexer) Consumed() { } } -func (r *Lexer) unsafeString() (string, []byte) { +func (r *Lexer) unsafeString(skipUnescape bool) (string, []byte) { if r.token.kind == tokenUndef && r.Ok() { r.FetchToken() } @@ -610,6 +606,13 @@ func (r *Lexer) unsafeString() (string, []byte) { r.errInvalidToken("string") return "", nil } + if !skipUnescape { + if err := r.unescapeStringToken(); err != nil { + r.errInvalidToken("string") + return "", nil + } + } + bytes := r.token.byteValue ret := bytesToStr(r.token.byteValue) r.consume() @@ -621,13 +624,19 @@ func (r *Lexer) unsafeString() (string, []byte) { // Warning: returned string may point to the input buffer, so the string should not outlive // the input buffer. Intended pattern of usage is as an argument to a switch statement. func (r *Lexer) UnsafeString() string { - ret, _ := r.unsafeString() + ret, _ := r.unsafeString(false) return ret } // UnsafeBytes returns the byte slice if the token is a string literal. func (r *Lexer) UnsafeBytes() []byte { - _, ret := r.unsafeString() + _, ret := r.unsafeString(false) + return ret +} + +// UnsafeFieldName returns current member name string token +func (r *Lexer) UnsafeFieldName(skipUnescape bool) string { + ret, _ := r.unsafeString(skipUnescape) return ret } @@ -640,7 +649,16 @@ func (r *Lexer) String() string { r.errInvalidToken("string") return "" } - ret := string(r.token.byteValue) + if err := r.unescapeStringToken(); err != nil { + r.errInvalidToken("string") + return "" + } + var ret string + if r.token.byteValueCloned { + ret = bytesToStr(r.token.byteValue) + } else { + ret = string(r.token.byteValue) + } r.consume() return ret } @@ -839,7 +857,7 @@ func (r *Lexer) Int() int { } func (r *Lexer) Uint8Str() uint8 { - s, b := r.unsafeString() + s, b := r.unsafeString(false) if !r.Ok() { return 0 } @@ -856,7 +874,7 @@ func (r *Lexer) Uint8Str() uint8 { } func (r *Lexer) Uint16Str() uint16 { - s, b := r.unsafeString() + s, b := r.unsafeString(false) if !r.Ok() { return 0 } @@ -873,7 +891,7 @@ func (r *Lexer) Uint16Str() uint16 { } func (r *Lexer) Uint32Str() uint32 { - s, b := r.unsafeString() + s, b := r.unsafeString(false) if !r.Ok() { return 0 } @@ -890,7 +908,7 @@ func (r *Lexer) Uint32Str() uint32 { } func (r *Lexer) Uint64Str() uint64 { - s, b := r.unsafeString() + s, b := r.unsafeString(false) if !r.Ok() { return 0 } @@ -915,7 +933,7 @@ func (r *Lexer) UintptrStr() uintptr { } func (r *Lexer) Int8Str() int8 { - s, b := r.unsafeString() + s, b := r.unsafeString(false) if !r.Ok() { return 0 } @@ -932,7 +950,7 @@ func (r *Lexer) Int8Str() int8 { } func (r *Lexer) Int16Str() int16 { - s, b := r.unsafeString() + s, b := r.unsafeString(false) if !r.Ok() { return 0 } @@ -949,7 +967,7 @@ func (r *Lexer) Int16Str() int16 { } func (r *Lexer) Int32Str() int32 { - s, b := r.unsafeString() + s, b := r.unsafeString(false) if !r.Ok() { return 0 } @@ -966,7 +984,7 @@ func (r *Lexer) Int32Str() int32 { } func (r *Lexer) Int64Str() int64 { - s, b := r.unsafeString() + s, b := r.unsafeString(false) if !r.Ok() { return 0 } @@ -1004,7 +1022,7 @@ func (r *Lexer) Float32() float32 { } func (r *Lexer) Float32Str() float32 { - s, b := r.unsafeString() + s, b := r.unsafeString(false) if !r.Ok() { return 0 } @@ -1037,7 +1055,7 @@ func (r *Lexer) Float64() float64 { } func (r *Lexer) Float64Str() float64 { - s, b := r.unsafeString() + s, b := r.unsafeString(false) if !r.Ok() { return 0 } diff --git a/jlexer/lexer_test.go b/jlexer/lexer_test.go index fcf9780..d3ab65b 100644 --- a/jlexer/lexer_test.go +++ b/jlexer/lexer_test.go @@ -33,7 +33,7 @@ func TestString(t *testing.T) { got := l.String() if got != test.want { - t.Errorf("[%d, %q] String() = %v; want %v", i, test.toParse, got, test.want) + t.Errorf("[%d, %q] String() = '%v'; want '%v'", i, test.toParse, got, test.want) } err := l.Error() if err != nil && !test.wantError { @@ -262,12 +262,12 @@ func TestJsonNumber(t *testing.T) { {toParse: `10`, want: json.Number("10"), wantValue: int64(10)}, {toParse: `0`, want: json.Number("0"), wantValue: int64(0)}, {toParse: `0.12`, want: json.Number("0.12"), wantValue: 0.12}, - {toParse: `25E-4`, want: json.Number("25E-4"), wantValue: 25E-4}, + {toParse: `25E-4`, want: json.Number("25E-4"), wantValue: 25e-4}, {toParse: `"10"`, want: json.Number("10"), wantValue: int64(10)}, {toParse: `"0"`, want: json.Number("0"), wantValue: int64(0)}, {toParse: `"0.12"`, want: json.Number("0.12"), wantValue: 0.12}, - {toParse: `"25E-4"`, want: json.Number("25E-4"), wantValue: 25E-4}, + {toParse: `"25E-4"`, want: json.Number("25E-4"), wantValue: 25e-4}, {toParse: `"foo"`, want: json.Number("foo"), wantValueError: true}, {toParse: `null`, want: json.Number(""), wantValueError: true}, @@ -324,7 +324,7 @@ func TestFetchStringUnterminatedString(t *testing.T) { l := Lexer{Data: test.data} l.fetchString() if l.pos > len(l.Data) { - t.Errorf("fetchString(%s): pos should not be greater than length of Data", test.data) + t.Errorf("fetchString(%s): pos=%v should not be greater than length of Data = %v", test.data, l.pos, len(l.Data)) } if l.Error() == nil { t.Errorf("fetchString(%s): should add parsing error", test.data) diff --git a/tests/members_escaped.go b/tests/members_escaped.go new file mode 100644 index 0000000..0fd68e8 --- /dev/null +++ b/tests/members_escaped.go @@ -0,0 +1,6 @@ +package tests + +//easyjson:json +type MembersEscaped struct { + A string `json:"漢語"` +} diff --git a/tests/members_escaping_test.go b/tests/members_escaping_test.go new file mode 100644 index 0000000..e385648 --- /dev/null +++ b/tests/members_escaping_test.go @@ -0,0 +1,51 @@ +package tests + +import ( + "reflect" + "testing" + + "github.com/mailru/easyjson" +) + +func TestMembersEscaping(t *testing.T) { + cases := []struct { + data string + esc MembersEscaped + unesc MembersUnescaped + }{ + { + data: `{"漢語": "中国"}`, + esc: MembersEscaped{A: "中国"}, + unesc: MembersUnescaped{A: "中国"}, + }, + { + data: `{"漢語": "\u4e2D\u56fD"}`, + esc: MembersEscaped{A: "中国"}, + unesc: MembersUnescaped{A: "中国"}, + }, + { + data: `{"\u6f22\u8a9E": "中国"}`, + esc: MembersEscaped{A: "中国"}, + unesc: MembersUnescaped{A: ""}, + }, + { + data: `{"\u6f22\u8a9E": "\u4e2D\u56fD"}`, + esc: MembersEscaped{A: "中国"}, + unesc: MembersUnescaped{A: ""}, + }, + } + + for i, c := range cases { + var esc MembersEscaped + easyjson.Unmarshal([]byte(c.data), &esc) + if !reflect.DeepEqual(esc, c.esc) { + t.Errorf("[%d] TestMembersEscaping(): got=%+v, exp=%+v", i, esc, c.esc) + } + + var unesc MembersUnescaped + easyjson.Unmarshal([]byte(c.data), &unesc) + if !reflect.DeepEqual(unesc, c.unesc) { + t.Errorf("[%d] TestMembersEscaping(): no-unescaping case: got=%+v, exp=%+v", i, esc, c.esc) + } + } +} diff --git a/tests/members_unescaped.go b/tests/members_unescaped.go new file mode 100644 index 0000000..1b95721 --- /dev/null +++ b/tests/members_unescaped.go @@ -0,0 +1,6 @@ +package tests + +//easyjson:json +type MembersUnescaped struct { + A string `json:"漢語"` +}