diff options
| -rw-r--r-- | asspb.go | 214 | ||||
| -rw-r--r-- | asspb_test.go | 356 |
2 files changed, 356 insertions, 214 deletions
@@ -157,10 +157,10 @@ package asspb import ( "bytes" - "encoding/json" "errors" "fmt" "iter" + "reflect" "regexp" "strconv" "strings" @@ -189,21 +189,32 @@ func (e *syntaxError) Error() string { return fmt.Sprintf("%d:%d syntax error: %s", e.line, e.col, e.reason) } -func appendAnyRepeated(prev any, new ...any) []any { - if prev == nil { - return new +func (p *parser) fieldMap(s reflect.Value) (map[string]reflect.Value, error) { + if !(s.Kind() == reflect.Struct || s.Kind() == reflect.Pointer && s.Type().Elem().Kind() == reflect.Struct) { + return nil, p.error("field must be a struct or pointer to struct (got %T)", s.Interface()) } - if prevList, ok := prev.([]any); ok { - return append(prevList, new...) + if s.Kind() == reflect.Pointer { + if s.IsNil() { + s.Set(reflect.New(s.Type().Elem())) + } + s = s.Elem() } - return append([]any{prev}, new...) -} - -func appendAny(prev any, new any) any { - if prev == nil { - return new + m := make(map[string]reflect.Value) + for i := range s.NumField() { + field := s.Type().Field(i) + if !field.IsExported() { + continue + } + fieldName := field.Name + if tag, ok := field.Tag.Lookup("ccl"); ok { + fieldName, _, _ = strings.Cut(tag, ",") + if fieldName == "-" { + continue + } + } + m[fieldName] = s.Field(i) } - return appendAnyRepeated(prev, new) + return m, nil } type parser struct { @@ -249,29 +260,41 @@ func (p *parser) next() ([]byte, error) { var numRE = regexp.MustCompile(`^[-+]?(0[xX][0-9a-fA-F]+|((0|[1-9][0-9]*)(\.[0-9]*)?|\.[0-9]+)([eE][-+]?[0-9]+)?)$`) -func parseNum(numBytes []byte) (any, bool) { +func (p *parser) parseNum(out reflect.Value, numBytes []byte) error { if !numRE.Match(numBytes) { - return nil, false + return p.error("invalid number") } if bytes.HasPrefix(numBytes, []byte("0x")) || bytes.HasPrefix(numBytes, []byte("0X")) { + if out.Kind() != reflect.Int64 { + return p.error("field must be int64") + } n, err := strconv.ParseInt(string(numBytes[2:]), 16, 64) if err != nil { - return nil, false + return p.error("invalid number") } - return n, true + out.SetInt(n) + return nil } if bytes.ContainsAny(numBytes, ".eE") { + if out.Kind() != reflect.Float64 { + return p.error("field must be float64") + } n, err := strconv.ParseFloat(string(numBytes), 64) if err != nil { - return nil, false + return p.error("invalid number") } - return n, true + out.SetFloat(n) + return nil + } + if out.Kind() != reflect.Int64 { + return p.error("field must be int64") } n, err := strconv.ParseInt(string(numBytes), 10, 64) if err != nil { - return nil, false + return p.error("invalid number") } - return n, true + out.SetInt(n) + return nil } var escapesRE = regexp.MustCompile(`(?s)\\(.|\r\n|[0-7]{3}|x[0-9a-fA-F]{2}|u[0-9a-fA-F]{4}|U[0-9a-fA-F]{8})`) @@ -350,146 +373,171 @@ func (p *parser) unescape(rawStr []byte) ([]byte, error) { return escaped, nil } -func (p *parser) parseString(tok []byte) (string, error) { +func (p *parser) parseString(out reflect.Value, tok []byte) error { + if out.Kind() != reflect.String { + return p.error("field must be string") + } s := new(strings.Builder) for { ss, err := p.unescape(tok[1 : len(tok)-1]) if err != nil { - return "", err + return err } s.Write(ss) nextTok, err := p.peek() if err != nil || nextTok[0] != '\'' && nextTok[0] != '"' { - return s.String(), nil + out.SetString(s.String()) + return nil } p.next() tok = nextTok } } -func (p *parser) parseMessage() (map[string]any, error) { - m := make(map[string]any) +func (p *parser) parseMessage(out reflect.Value) error { + fieldMap, err := p.fieldMap(out) + if err != nil { + return err + } for { tok, err := p.next() if err != nil || tok[0] == '}' { - return m, err + return err } - if err := p.parseFieldVal(m, tok); err != nil { - return nil, err + if err := p.parseFieldVal(fieldMap, tok); err != nil { + return err } } } -func (p *parser) parseVal(tok []byte) (any, error) { - var ( - val any - err error - ) +func (p *parser) parseVal(out reflect.Value, tok []byte) error { + var err error switch tok[0] { case '{': - val, err = p.parseMessage() + err = p.parseMessage(out) case '[': - val, err = p.parseList() + err = p.parseList(out) case '\'', '"': - val, err = p.parseString(tok) + err = p.parseString(out, tok) default: switch string(tok) { case "true", "yes", "on": - val = true + if out.Kind() != reflect.Bool { + return p.error("field must be bool") + } + out.SetBool(true) case "false", "no", "off": - val = false - default: - var ok bool - val, ok = parseNum(tok) - if !ok { - return nil, p.error("expecting field value") + if out.Kind() != reflect.Bool { + return p.error("field must be bool") } + out.SetBool(false) + default: + err = p.parseNum(out, tok) } } - return val, err + return err } -func (p *parser) parseList() ([]any, error) { - l := []any{} +func (p *parser) parseList(out reflect.Value) error { + if out.Kind() != reflect.Slice { + return p.error("field must be slice") + } + if out.IsNil() { + // Not technically necessary since nil slices are usually treated the + // same as any empty slice, but usually nil would mean that the user + // didn't set the value, at least for types that have a nil. + out.Set(reflect.MakeSlice(out.Type(), 0, 0)) + } for i := 0; ; i++ { tok, err := p.next() if err != nil || tok[0] == ']' { - return l, err + return err } if i > 0 { if tok[0] != ',' { - return nil, p.error("expecting comma") + return p.error("expecting comma") } tok, err = p.next() if err != nil || tok[0] == ']' { // allow trailing comma - return l, err + return err } } - val, err := p.parseVal(tok) - if err != nil { - return nil, err + out.Set(reflect.Append(out, reflect.Zero(out.Type().Elem()))) + if err := p.parseVal(out.Index(out.Len()-1), tok); err != nil { + return err + } + } +} + +func (p *parser) appendVal(out reflect.Value, tok []byte) error { + if tok[0] == '[' || out.Kind() != reflect.Slice { + if err := p.parseVal(out, tok); err != nil { + return err + } + } else { + out.Set(reflect.Append(out, reflect.Zero(out.Type().Elem()))) + if err := p.parseVal(out.Index(out.Len()-1), tok); err != nil { + return err } - l = append(l, val) } + return nil } -func (p *parser) parseFieldVal(out map[string]any, field []byte) error { +func (p *parser) parseFieldVal(fieldMap map[string]reflect.Value, field []byte) error { if b := field[0]; !(b == '_' || 'a' <= b && b <= 'z' || 'A' <= b && b <= 'Z') { return p.error("expecting field") } + structField, ok := fieldMap[string(field)] + if !ok { + return p.error("no field named %q", field) + } tok, err := p.next() if err != nil { return err } - var val any switch tok[0] { - case '{': - if val, err = p.parseMessage(); err != nil { + case '{', '[': + if err := p.appendVal(structField, tok); err != nil { return err } - case '[': - val, err := p.parseList() - if err != nil { - return err - } - out[string(field)] = appendAnyRepeated(out[string(field)], val...) - return nil case ':': tok, err := p.next() if err != nil { return err } - if val, err = p.parseVal(tok); err != nil { + if err := p.appendVal(structField, tok); err != nil { return err } - if l, ok := val.([]any); ok { - out[string(field)] = appendAnyRepeated(out[string(field)], l...) - return nil - } default: return p.error("expecting colon") } - out[string(field)] = appendAny(out[string(field)], val) return nil } -func (p *parser) parse() (map[string]any, error) { - m := make(map[string]any) +func (p *parser) parse(v any) error { + sp := reflect.ValueOf(v) + if sp.Kind() != reflect.Pointer || sp.IsNil() { + return p.error("value must be a non-nil pointer") + } + fieldMap, err := p.fieldMap(sp.Elem()) + if err != nil { + return err + } for { tok, err := p.next() if err != nil { if err == errEOF { - return m, nil + return nil } - return nil, err + return err } - if err := p.parseFieldVal(m, tok); err != nil { - return nil, err + if err := p.parseFieldVal(fieldMap, tok); err != nil { + return err } } } -// Unmarshal parses an asspb message and writes the result into v. Unmarshal +// Unmarshal parses a ccl message and writes the result into v. Unmarshal // internally calls json.Unmarshal for the reflection-based struct unpacking, so // feel free to use json struct tags on v, or implement UnmarshalJSON to control // the unmarshalling behavior. @@ -502,13 +550,5 @@ func (p *parser) parse() (map[string]any, error) { func Unmarshal(data []byte, v any) error { nextToken, stop := iter.Pull2(tokens(data)) defer stop() - m, err := (&parser{nextTok: nextToken, data: data}).parse() - if err != nil { - return err - } - jsonBytes, err := json.Marshal(m) - if err != nil { - return err - } - return json.Unmarshal(jsonBytes, v) + return (&parser{nextTok: nextToken, data: data}).parse(v) } diff --git a/asspb_test.go b/asspb_test.go index 63361ef..7deb8f3 100644 --- a/asspb_test.go +++ b/asspb_test.go @@ -9,221 +9,240 @@ import ( func TestUnmarshal(t *testing.T) { t.Parallel() + type nestedMessage struct { + Field int64 `ccl:"field"` + } + type message struct { + String string `ccl:"string"` + String2 string `ccl:"string2"` + Int int64 `ccl:"int"` + Float float64 `ccl:"float"` + Bool bool `ccl:"bool"` + Bool2 bool `ccl:"bool2"` + Message *nestedMessage `ccl:"message"` + Repeated []int64 `ccl:"repeated"` + RepeatedMessage []nestedMessage `ccl:"repeated_message"` + + Ignore map[int]int `ccl:"-,"` // unlike JSON this also means ignore + unexported int64 + } + for _, tc := range []struct { desc string msg string - want map[string]any + want message }{{ desc: "Complete", msg: `# This is a comment -field_string: 'asdf\n' # comment end of line -field_doublestring: "asdf\n" -field_int: 10 -field_float: 10.5e13 -field_true: true -field_false: false -field_nested { asdf: 10 } -field_repeated [1, 2, 3] -field_repeated: 4 -field_repeated [5, 6] +string: 'asdf\n' # comment end of line +string2: "asdf\n" +int: 10 +float: 10.5e13 +bool: true +bool2: false +message { field: 10 } +repeated [1, 2, 3] +repeated: 4 +repeated [5, 6] `, - want: map[string]any{ - "field_string": "asdf\n", - "field_doublestring": "asdf\n", - "field_int": 10., - "field_float": 10.5e13, - "field_true": true, - "field_false": false, - "field_nested": map[string]any{"asdf": 10.}, - "field_repeated": []any{1., 2., 3., 4., 5., 6.}, + want: message{ + String: "asdf\n", + String2: "asdf\n", + Int: 10, + Float: 10.5e13, + Bool: true, + Bool2: false, + Message: &nestedMessage{Field: 10}, + Repeated: []int64{1, 2, 3, 4, 5, 6}, }, }, { desc: "MultilineString", - msg: `field: "strings + msg: `string: "strings can just span multiple lines"`, - want: map[string]any{"field": "strings\ncan just span multiple lines"}, + want: message{String: "strings\ncan just span multiple lines"}, }, { desc: "Zero", - msg: `field: 0`, - want: map[string]any{"field": 0.}, + msg: `int: 0`, + want: message{Int: 0}, }, { desc: "Hex", - msg: `field: 0xff`, - want: map[string]any{"field": 255.}, + msg: `int: 0xff`, + want: message{Int: 255}, }, { desc: "CapitalHex", - msg: `field: 0XfF`, - want: map[string]any{"field": 255.}, + msg: `int: 0XfF`, + want: message{Int: 255}, }, { desc: "HexLeadingZero", - msg: `field: 0x0f`, - want: map[string]any{"field": 15.}, + msg: `int: 0x0f`, + want: message{Int: 15}, }, { desc: "Float", - msg: `field: 1.5e10`, - want: map[string]any{"field": 1.5e10}, + msg: `float: 1.5e10`, + want: message{Float: 1.5e10}, }, { desc: "FloatCapitalE", - msg: `field: 1.5E10`, - want: map[string]any{"field": 1.5e10}, + msg: `float: 1.5E10`, + want: message{Float: 1.5e10}, }, { desc: "NegativeFloat", - msg: `field: -1.5e-10`, - want: map[string]any{"field": -1.5e-10}, + msg: `float: -1.5e-10`, + want: message{Float: -1.5e-10}, }, { desc: "PositiveFloat", - msg: `field: +1.5e+10`, - want: map[string]any{"field": 1.5e10}, + msg: `float: +1.5e+10`, + want: message{Float: 1.5e10}, }, { desc: "Int", - msg: `field: 10`, - want: map[string]any{"field": 10.}, + msg: `int: 10`, + want: message{Int: 10}, }, { desc: "NegativeInt", - msg: `field: -10`, - want: map[string]any{"field": -10.}, + msg: `int: -10`, + want: message{Int: -10}, }, { desc: "PositiveInt", - msg: `field: +10`, - want: map[string]any{"field": 10.}, + msg: `int: +10`, + want: message{Int: 10}, }, { desc: "String", - msg: `field: 'asdf'`, - want: map[string]any{"field": "asdf"}, + msg: `string: 'asdf'`, + want: message{String: "asdf"}, }, { desc: "DoubleString", - msg: `field: "asdf"`, - want: map[string]any{"field": "asdf"}, + msg: `string: "asdf"`, + want: message{String: "asdf"}, }, { desc: "StringEscapeSingle", - msg: `you: 'ain\'t'`, - want: map[string]any{"you": "ain't"}, + msg: `string: 'ain\'t'`, + want: message{String: "ain't"}, }, { desc: "DoubleStringEscapeSingle", - msg: `I: "won\'t"`, - want: map[string]any{"I": "won't"}, + msg: `string: "won\'t"`, + want: message{String: "won't"}, }, { desc: "StringEscapeDouble", - msg: `field: '\"'`, - want: map[string]any{"field": `"`}, + msg: `string: '\"'`, + want: message{String: `"`}, }, { desc: "DoubleStringEscapeDouble", - msg: `field: "\""`, - want: map[string]any{"field": `"`}, + msg: `string: "\""`, + want: message{String: `"`}, }, { desc: "StringEscapeQuestionMark", - msg: `field: "\?"`, - want: map[string]any{"field": "?"}, + msg: `string: "\?"`, + want: message{String: "?"}, }, { desc: "StringEscapeBackslash", - msg: `field: '\\'`, - want: map[string]any{"field": `\`}, + msg: `string: '\\'`, + want: message{String: `\`}, }, { desc: "StringEscapeA", - msg: `field: '\a'`, - want: map[string]any{"field": "\a"}, + msg: `string: '\a'`, + want: message{String: "\a"}, }, { desc: "StringEscapeB", - msg: `field: '\b'`, - want: map[string]any{"field": "\b"}, + msg: `string: '\b'`, + want: message{String: "\b"}, }, { desc: "StringEscapeF", - msg: `field: '\f'`, - want: map[string]any{"field": "\f"}, + msg: `string: '\f'`, + want: message{String: "\f"}, }, { desc: "StringEscapeN", - msg: `field: '\n'`, - want: map[string]any{"field": "\n"}, + msg: `string: '\n'`, + want: message{String: "\n"}, }, { desc: "StringEscapeR", - msg: `field: '\r'`, - want: map[string]any{"field": "\r"}, + msg: `string: '\r'`, + want: message{String: "\r"}, }, { desc: "StringEscapeT", - msg: `field: '\t'`, - want: map[string]any{"field": "\t"}, + msg: `string: '\t'`, + want: message{String: "\t"}, }, { desc: "StringEscapeV", - msg: `field: '\v'`, - want: map[string]any{"field": "\v"}, + msg: `string: '\v'`, + want: message{String: "\v"}, }, { desc: "StringHex", - msg: `field: '\x0a'`, - want: map[string]any{"field": "\n"}, + msg: `string: '\x0a'`, + want: message{String: "\n"}, }, { desc: "StringHexHighByte", - msg: `field: "\xe4\xb8\x96"`, - want: map[string]any{"field": "世"}, + msg: `string: "\xe4\xb8\x96"`, + want: message{String: "世"}, }, { desc: "StringUnicode", - msg: `field: '\u2014'`, - want: map[string]any{"field": "—"}, + msg: `string: '\u2014'`, + want: message{String: "—"}, }, { desc: "StringOctal", - msg: `field: '\033'`, - want: map[string]any{"field": "\033"}, + msg: `string: '\033'`, + want: message{String: "\033"}, }, { desc: "Message", - msg: `field { nested_field: 10 }`, - want: map[string]any{"field": map[string]any{"nested_field": 10.}}, + msg: `message { field: 10 }`, + want: message{Message: &nestedMessage{Field: 10}}, }, { desc: "EmptyMessage", - msg: `field {}`, - want: map[string]any{"field": map[string]any{}}, + msg: `message {}`, + want: message{Message: &nestedMessage{}}, }, { desc: "Repeated", - msg: `field: 1 -field: 2`, - want: map[string]any{"field": []any{1., 2.}}, + msg: ` + repeated: 1 + repeated: 2`, + want: message{Repeated: []int64{1, 2}}, }, { desc: "RepeatedList", - msg: `field: [1, 2]`, - want: map[string]any{"field": []any{1., 2.}}, + msg: `repeated: [1, 2]`, + want: message{Repeated: []int64{1, 2}}, }, { desc: "EmptyList", - msg: `field: []`, - want: map[string]any{"field": []any{}}, + msg: `repeated: []`, + want: message{Repeated: []int64{}}, }, { desc: "RepeatedListTrailingComma", - msg: `field: [ - 1, - 2, - ]`, - want: map[string]any{"field": []any{1., 2.}}, + msg: `repeated: [ + 1, + 2, + ]`, + want: message{Repeated: []int64{1, 2}}, }, { desc: "ListOfMessage", - msg: `field: [{}]`, - want: map[string]any{"field": []any{map[string]any{}}}, + msg: `repeated_message: [{}]`, + want: message{RepeatedMessage: []nestedMessage{{}}}, }, { desc: "CStyleComment", - msg: `field: /** inline comment **/ {}`, - want: map[string]any{"field": map[string]any{}}, + msg: `message: /** inline comment **/ {}`, + want: message{Message: &nestedMessage{}}, }, { desc: "CStyleLineComment", - msg: `field: {} // line comment`, - want: map[string]any{"field": map[string]any{}}, + msg: `message: {} // line comment`, + want: message{Message: &nestedMessage{}}, }, { desc: "ConcatStrings", - msg: `field: 'that'"'"'s cool'`, - want: map[string]any{"field": "that's cool"}, + msg: `string: 'that'"'"'s cool'`, + want: message{String: "that's cool"}, }, { desc: "RemoveNewline", - msg: `field: 'remove newline \ + msg: `string: 'remove newline \ from string'`, - want: map[string]any{"field": "remove newline from string"}, + want: message{String: "remove newline from string"}, }, { desc: "RemoveNewlineWindows", - msg: "field: 'remove newline \\\r\nfrom string'", - want: map[string]any{"field": "remove newline from string"}, + msg: "string: 'remove newline \\\r\nfrom string'", + want: message{String: "remove newline from string"}, }} { t.Run(tc.desc, func(t *testing.T) { t.Parallel() - got := make(map[string]any) + var got message if err := Unmarshal([]byte(tc.msg), &got); err != nil { t.Fatalf("Unmarshal(%q) failed: %s\n", tc.msg, err) } - if diff := cmp.Diff(tc.want, got); diff != "" { + if diff := cmp.Diff(tc.want, got, cmp.AllowUnexported(message{})); diff != "" { t.Errorf("Unmarshal(%q) returned unexpected diff (-want +got):\n%s", tc.msg, diff) } }) @@ -233,56 +252,70 @@ from string'`, func TestUnmarshal_Invalid(t *testing.T) { t.Parallel() + type nestedMessage struct { + Int int `ccl:"int"` + } + type message struct { + Int int64 `ccl:"int"` + String string `ccl:"string"` + Msg nestedMessage `ccl:"msg"` + Repeated []int64 `ccl:"repeated"` + RepeatedMsg []nestedMessage `ccl:"repeated_msg"` + } + for _, tc := range []struct { desc string msg string }{{ desc: "BadNum", - msg: `field: .`, + msg: `int: .`, }, { desc: "BadStringEscape", - msg: `field: '\g'`, + msg: `string: '\g'`, }, { desc: "BadDoubleStringEscape", - msg: `field: "\g"`, + msg: `string: "\g"`, }, { desc: "UnterminatedString", - msg: `field: '`, + msg: `string: '`, }, { desc: "UnterminatedDoubleString", - msg: `field: "`, + msg: `string: "`, }, { desc: "NoFieldName", msg: `10`, }, { desc: "MsgNoFieldName", - msg: `field {10}`, + msg: `msg {10}`, }, { desc: "ListMissingComma", - msg: `field [1 2]`, + msg: `repeated [1 2]`, }, { desc: "ListBadVal", - msg: `field [asdf]`, + msg: `repeated [asdf]`, }, { desc: "ListBadMsgVal", - msg: `field [{asdf}]`, + msg: `repeated_msg [{asdf}]`, }, { desc: "IntLeadingZero", - msg: `field: 0644`, + msg: `int: 0644`, }, { desc: "InvalidOctal", - msg: `field: "\777"`, + msg: `string: "\777"`, }, { desc: "InvalidUTF8", - msg: `field: "\x80"`, + msg: `string: "\x80"`, }, { desc: "FieldMissingVal", - msg: `field`, + msg: `string`, + }, { + desc: "FieldMissingColon", + msg: `string "abc"`, }} { t.Run(tc.desc, func(t *testing.T) { t.Parallel() - got := make(map[string]any) + var got message err := Unmarshal([]byte(tc.msg), &got) if err == nil { t.Errorf("Unmarshal(%q) returned success, want error", tc.msg) @@ -291,14 +324,83 @@ func TestUnmarshal_Invalid(t *testing.T) { } } +func TestUnmarshal_InvalidType(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + desc string + msg string + out any + }{{ + desc: "Nil", + msg: `field {}`, + out: (*struct{})(nil), + }, { + desc: "Struct", + msg: `field {}`, + out: new(int), + }, { + desc: "Int", + msg: `Field: 123`, + out: new(struct{ Field string }), + }, { + desc: "IntHex", + msg: `Field: 0x123`, + out: new(struct{ Field string }), + }, { + desc: "Float", + msg: `Field: 123.`, + out: new(struct{ Field int64 }), + }, { + desc: "NestedMessage", + msg: `Field {Field {}}`, + out: new(struct{ Field int64 }), + }, { + desc: "True", + msg: `F:on`, + out: new(struct{ F int64 }), + }, { + desc: "False", + msg: `F:no`, + out: new(struct{ F int64 }), + }, { + desc: "List", + msg: `F:[]`, + out: new(struct{ F int64 }), + }, { + desc: "RepeatedBool", + msg: `F:on F:no`, + out: new(struct{ F []int64 }), + }, { + desc: "String", + msg: `F:"abc"`, + out: new(struct{ F int64 }), + }} { + t.Run(tc.desc, func(t *testing.T) { + t.Parallel() + + err := Unmarshal([]byte(tc.msg), tc.out) + if err == nil { + t.Errorf("Unmarshal(%+v) returned success, want error", tc.out) + } + }) + } +} + func TestUnmarshal_ErrorLineCol(t *testing.T) { + t.Parallel() + + type message struct { + Secret int64 `ccl:"secret"` + } + msg := ` ###### This is a very important file please do not modify ######################################################### ################ The more ## I put the more secure it is###### secret:12345; # oops typo ` - err := Unmarshal([]byte(msg), new(map[string]any)) + err := Unmarshal([]byte(msg), new(message)) syntaxErr, ok := err.(*syntaxError) if !ok { t.Errorf("Unmarshal(%q): expected *syntaxError, got error %T %[2]v", msg, err) |
