diff --git a/datetime.go b/datetime.go new file mode 100644 index 0000000..6686f84 --- /dev/null +++ b/datetime.go @@ -0,0 +1,133 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package interpres + +import ( + "fmt" + "regexp" + "strings" + "time" +) + +// TOML distinguishes four date-time kinds. interpres decodes an offset +// date-time to a plain time.Time (it carries a zone), and uses the wrapper +// types below for the local variants so callers can tell them apart. + +// LocalDateTime is a TOML local date-time with no offset, e.g. +// 1979-05-27T07:32:00. The embedded time.Time is in UTC. +type LocalDateTime struct{ time.Time } + +// LocalDate is a TOML local date with no time or offset, e.g. 1979-05-27. +// The embedded time.Time is at midnight UTC. +type LocalDate struct{ time.Time } + +// LocalTime is a TOML local time with no date or offset, e.g. 07:32:00.999999. +// The embedded time.Time uses the zero date. +type LocalTime struct{ time.Time } + +// String returns the TOML-canonical rendering of the local date-time, e.g. +// "1979-05-27T07:32:00" or "...:00.000000123" when the time has a fractional +// second. The fractional component is zero-padded to nanosecond precision. +func (ldt LocalDateTime) String() string { + base := ldt.Format("2006-01-02T15:04:05") + if ns := ldt.Nanosecond(); ns > 0 { + return base + "." + fmt.Sprintf("%09d", ns) + } + return base +} + +// String returns the TOML-canonical rendering of the local date, e.g. +// "1979-05-27". +func (ld LocalDate) String() string { return ld.Format("2006-01-02") } + +// String returns the TOML-canonical rendering of the local time, e.g. +// "07:32:00" or "...:00.000000123" when the time has a fractional second. +// The fractional component is zero-padded to nanosecond precision. +func (lt LocalTime) String() string { + base := lt.Format("15:04:05") + if ns := lt.Nanosecond(); ns > 0 { + return base + "." + fmt.Sprintf("%09d", ns) + } + return base +} + +var ( + offsetDateTimeLayouts = []string{ + "2006-01-02T15:04:05.999999999Z07:00", + "2006-01-02T15:04:05Z07:00", + "2006-01-02 15:04:05.999999999Z07:00", + "2006-01-02 15:04:05Z07:00", + } + localDateTimeLayouts = []string{ + "2006-01-02T15:04:05.999999999", + "2006-01-02T15:04:05", + "2006-01-02 15:04:05.999999999", + "2006-01-02 15:04:05", + } + localTimeLayouts = []string{ + "15:04:05.999999999", + "15:04:05", + } +) + +// dateTimeShape enforces the strict TOML grammar (two-digit components) that +// time.Parse would otherwise accept loosely (e.g. a single-digit hour). +var dateTimeShape = regexp.MustCompile( + `^\d{4}-\d{2}-\d{2}([Tt ]\d{2}:\d{2}:\d{2}(\.\d+)?([Zz]|[+-]\d{2}:\d{2})?)?$` + + `|^\d{2}:\d{2}:\d{2}(\.\d+)?$`, +) + +// parseDateTime classifies and parses a bare token as a TOML date-time value. +// It returns the decoded value (time.Time, LocalDateTime, LocalDate, or +// LocalTime) and whether the token was a date-time at all. +func parseDateTime(tok string) (any, bool) { + if tok == "" || tok[0] < '0' || tok[0] > '9' { + return nil, false + } + if !strings.ContainsAny(tok, "-:") { + return nil, false + } + if !dateTimeShape.MatchString(tok) { + return nil, false + } + // The ABNF accepts lowercase "t"/"z"; time.Parse only matches uppercase. + norm := strings.ToUpper(tok) + for _, layout := range offsetDateTimeLayouts { + if t, err := time.Parse(layout, norm); err == nil { + return t, true + } + } + for _, layout := range localDateTimeLayouts { + if t, err := time.Parse(layout, norm); err == nil { + return LocalDateTime{t}, true + } + } + if t, err := time.Parse("2006-01-02", norm); err == nil { + return LocalDate{t}, true + } + for _, layout := range localTimeLayouts { + if t, err := time.Parse(layout, norm); err == nil { + return LocalTime{t}, true + } + } + return nil, false +} + +// isDateToken reports whether s is exactly a YYYY-MM-DD date, used to detect a +// space-separated date-time written as "datetime". +func isDateToken(s string) bool { + if len(s) != 10 { + return false + } + for i := range len(s) { + if i == 4 || i == 7 { + if s[i] != '-' { + return false + } + } else if !isDecDigit(s[i]) { + return false + } + } + return true +} diff --git a/decode.go b/decode.go new file mode 100644 index 0000000..c571f1f --- /dev/null +++ b/decode.go @@ -0,0 +1,244 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package interpres + +import ( + "fmt" + "math" + "reflect" + "strings" + "time" +) + +// decoder maps a parsed TOML tree onto Go values via reflection. +type decoder struct { + disallowUnknown bool +} + +func newDecoder() *decoder { return &decoder{} } + +var timeType = reflect.TypeFor[time.Time]() + +func (d *decoder) decode(tree map[string]any, v any) error { + rv := reflect.ValueOf(v) + if rv.Kind() != reflect.Pointer || rv.IsNil() { + return fmt.Errorf("interpres: decode target must be a non-nil pointer") + } + return d.assign(tree, rv.Elem()) +} + +// assign stores data into dst, converting between the TOML value kinds and the +// destination's Go type. +func (d *decoder) assign(data any, dst reflect.Value) error { + if dst.Kind() == reflect.Pointer { + if dst.IsNil() { + dst.Set(reflect.New(dst.Type().Elem())) + } + return d.assign(data, dst.Elem()) + } + + // An any destination takes the value as-is. + if dst.Kind() == reflect.Interface && dst.NumMethod() == 0 { + dst.Set(reflect.ValueOf(data)) + return nil + } + + // Types implementing Unmarshaler get the parsed data wholesale and are + // responsible for setting their own state. The decoder does not consult + // any return value; whatever the receiver stores is kept. The lookup + // covers both T and *T so a pointer-receiver UnmarshalTOML method is + // invoked on an addressable struct field. + if dst.CanInterface() { + u, ok := dst.Interface().(Unmarshaler) + if !ok && dst.CanAddr() { + u, ok = dst.Addr().Interface().(Unmarshaler) + } + if ok { + if err := u.UnmarshalTOML(data); err != nil { + return fmt.Errorf("unmarshal: %w", err) + } + return nil + } + } + + switch v := data.(type) { + case map[string]any: + return d.assignTable(v, dst) + case []map[string]any: + return d.assignTableSlice(v, dst) + case []any: + return d.assignSlice(v, dst) + case string: + return setBasic(dst, reflect.ValueOf(v), "string") + case bool: + return setBasic(dst, reflect.ValueOf(v), "bool") + case int64: + return setInt(dst, v) + case float64: + return setFloat(dst, v) + case time.Time: + if dst.Type() != timeType { + return fmt.Errorf("interpres: cannot assign datetime to %s", dst.Type()) + } + dst.Set(reflect.ValueOf(v)) + return nil + default: + rv := reflect.ValueOf(data) + if rv.IsValid() && dst.Type() == rv.Type() { + dst.Set(rv) + return nil + } + return fmt.Errorf("interpres: cannot assign %T to %s", data, dst.Type()) + } +} + +func (d *decoder) assignTable(tbl map[string]any, dst reflect.Value) error { + switch dst.Kind() { + case reflect.Struct: + return d.assignStruct(tbl, dst) + case reflect.Map: + return d.assignMap(tbl, dst) + default: + return fmt.Errorf("interpres: cannot assign table to %s", dst.Type()) + } +} + +func (d *decoder) assignStruct(tbl map[string]any, dst reflect.Value) error { + fields := structFields(dst.Type()) + for key, val := range tbl { + field, ok := fields[strings.ToLower(key)] + if !ok { + if d.disallowUnknown { + return fmt.Errorf("interpres: unknown field %q for %s", key, dst.Type()) + } + continue + } + if err := d.assign(val, dst.Field(field)); err != nil { + return fmt.Errorf("%s: %w", key, err) + } + } + return nil +} + +func (d *decoder) assignMap(tbl map[string]any, dst reflect.Value) error { + if dst.Type().Key().Kind() != reflect.String { + return fmt.Errorf("interpres: map key must be a string, got %s", dst.Type().Key()) + } + if dst.IsNil() { + dst.Set(reflect.MakeMap(dst.Type())) + } + elemType := dst.Type().Elem() + for key, val := range tbl { + elem := reflect.New(elemType).Elem() + if err := d.assign(val, elem); err != nil { + return fmt.Errorf("%s: %w", key, err) + } + dst.SetMapIndex(reflect.ValueOf(key), elem) + } + return nil +} + +func (d *decoder) assignSlice(items []any, dst reflect.Value) error { + if dst.Kind() != reflect.Slice { + return fmt.Errorf("interpres: cannot assign array to %s", dst.Type()) + } + out := reflect.MakeSlice(dst.Type(), len(items), len(items)) + for i, item := range items { + if err := d.assign(item, out.Index(i)); err != nil { + return fmt.Errorf("[%d]: %w", i, err) + } + } + dst.Set(out) + return nil +} + +func (d *decoder) assignTableSlice(items []map[string]any, dst reflect.Value) error { + if dst.Kind() != reflect.Slice { + return fmt.Errorf("interpres: cannot assign array of tables to %s", dst.Type()) + } + out := reflect.MakeSlice(dst.Type(), len(items), len(items)) + for i, item := range items { + if err := d.assign(item, out.Index(i)); err != nil { + return fmt.Errorf("[%d]: %w", i, err) + } + } + dst.Set(out) + return nil +} + +// --- low-level setters ----------------------------------------------------- + +func setBasic(dst, val reflect.Value, kind string) error { + if dst.Kind() != val.Kind() { + return fmt.Errorf("interpres: cannot assign %s to %s", kind, dst.Type()) + } + dst.Set(val) + return nil +} + +func setInt(dst reflect.Value, v int64) error { + switch dst.Kind() { + case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64: + if dst.OverflowInt(v) { + return fmt.Errorf("interpres: integer %d overflows %s", v, dst.Type()) + } + dst.SetInt(v) + case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64: + if v < 0 { + return fmt.Errorf("interpres: cannot assign negative %d to %s", v, dst.Type()) + } + var max uint64 + switch dst.Kind() { + case reflect.Uint8: + max = math.MaxUint8 + case reflect.Uint16: + max = math.MaxUint16 + case reflect.Uint32: + max = math.MaxUint32 + } + if max != 0 && uint64(v) > max { + return fmt.Errorf("interpres: integer %d overflows %s", v, dst.Type()) + } + dst.SetUint(uint64(v)) + case reflect.Float32, reflect.Float64: + dst.SetFloat(float64(v)) + default: + return fmt.Errorf("interpres: cannot assign integer to %s", dst.Type()) + } + return nil +} + +func setFloat(dst reflect.Value, v float64) error { + switch dst.Kind() { + case reflect.Float32, reflect.Float64: + dst.SetFloat(v) + return nil + default: + return fmt.Errorf("interpres: cannot assign float to %s", dst.Type()) + } +} + +// structFields builds a lower-cased lookup of field name → field index for the +// exported fields of t, honouring `toml:"name"` tags. +func structFields(t reflect.Type) map[string]int { + fields := make(map[string]int, t.NumField()) + for i := range t.NumField() { + f := t.Field(i) + if f.PkgPath != "" { // unexported + continue + } + name := f.Name + if tag, ok := f.Tag.Lookup("toml"); ok { + tag = strings.Split(tag, ",")[0] + if tag == "-" { + continue + } + if tag != "" { + name = tag + } + } + fields[strings.ToLower(name)] = i + } + return fields +} diff --git a/decode_test.go b/decode_test.go new file mode 100644 index 0000000..d5ca5b3 --- /dev/null +++ b/decode_test.go @@ -0,0 +1,496 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package interpres + +import ( + "context" + "errors" + "fmt" + "math" + "strings" + "testing" +) + +func TestSyntaxErrorMessage(t *testing.T) { + err := &SyntaxError{Line: 7, Msg: "expected '=' after key"} + want := "interpres: line 7: expected '=' after key" + if got := err.Error(); got != want { + t.Errorf("Error() = %q, want %q", got, want) + } +} + +func TestParseRejectsInvalidUTF8(t *testing.T) { + _, err := Parse([]byte("v = \"\xff\"\n")) + if err == nil { + t.Fatal("expected a UTF-8 validation error") + } + se, ok := err.(*SyntaxError) + if !ok { + t.Fatalf("err is %T, want *SyntaxError", err) + } + if !strings.Contains(se.Msg, "UTF-8") { + t.Errorf("Msg = %q, want it to mention UTF-8", se.Msg) + } + if se.Line != 1 { + t.Errorf("Line = %d, want 1", se.Line) + } +} + +func TestUnmarshalIntoMap(t *testing.T) { + var m map[string]any + if err := Unmarshal([]byte(`name = "x" +count = 3 +`), &m); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if m["name"] != "x" { + t.Errorf("name = %#v", m["name"]) + } + if m["count"] != int64(3) { + t.Errorf("count = %#v (%T)", m["count"], m["count"]) + } +} + +func TestUnmarshalIntoMapNested(t *testing.T) { + var m map[string]any + if err := Unmarshal([]byte("[a]\nb = 2\n"), &m); err != nil { + t.Fatalf("unmarshal: %v", err) + } + a, ok := m["a"].(map[string]any) + if !ok { + t.Fatalf("a = %T, want map[string]any", m["a"]) + } + if a["b"] != int64(2) { + t.Errorf("a.b = %#v", a["b"]) + } +} + +func TestUnmarshalPrefilledMap(t *testing.T) { + m := map[string]any{"keep": "yes"} + if err := Unmarshal([]byte(`name = "x"`), &m); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if m["keep"] != "yes" { + t.Errorf("keep = %#v", m["keep"]) + } + if m["name"] != "x" { + t.Errorf("name = %#v", m["name"]) + } +} + +func TestUnmarshalTableToNonStructOrMap(t *testing.T) { + var s string + if err := Unmarshal([]byte("v = 1\n"), &s); err == nil { + t.Fatal("expected an error for table-to-scalar") + } +} + +func TestUnmarshalIntoAny(t *testing.T) { + // A non-nil any destination must accept the parsed tree. + var x any + if err := Unmarshal([]byte("[a]\nb = 2\n"), &x); err != nil { + t.Fatalf("unmarshal: %v", err) + } + tree, ok := x.(map[string]any) + if !ok { + t.Fatalf("x is %T, want map[string]any", x) + } + a, ok := tree["a"].(map[string]any) + if !ok { + t.Fatalf("a is %T, want map[string]any", tree["a"]) + } + if a["b"] != int64(2) { + t.Errorf("a.b = %#v", a["b"]) + } +} + +func TestUnmarshalIntoNilAny(t *testing.T) { + // A nil any target must still receive the parsed tree without + // panicking. + var x any + if err := Unmarshal([]byte("v = 1\n"), &x); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if x == nil { + t.Fatal("x is still nil after Unmarshal") + } + tree, ok := x.(map[string]any) + if !ok { + t.Fatalf("x is %T, want map[string]any", x) + } + if tree["v"] != int64(1) { + t.Errorf("v = %#v", tree["v"]) + } +} + +func TestParseContextHonoursCancellation(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if _, err := ParseContext(ctx, []byte("a = 1\n")); !errors.Is(err, context.Canceled) { + t.Fatalf("ParseContext returned %v, want context.Canceled", err) + } +} + +func TestUnmarshalContextHonoursCancellation(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + var cfg map[string]any + if err := UnmarshalContext(ctx, []byte("a = 1\n"), &cfg); !errors.Is(err, context.Canceled) { + t.Fatalf("UnmarshalContext returned %v, want context.Canceled", err) + } +} + +func TestDecoderDecodeContextHonoursCancellation(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + var cfg map[string]any + err := NewDecoder().DecodeContext(ctx, []byte("a = 1\n"), &cfg) + if !errors.Is(err, context.Canceled) { + t.Fatalf("DecodeContext returned %v, want context.Canceled", err) + } +} + +func TestContextRoundTrip(t *testing.T) { + // The *Context variants with a Background context must produce the same + // result as the non-context variants for ordinary inputs. + in := []byte(`title = "x" +count = 3 +`) + if _, err := ParseContext(context.Background(), in); err != nil { + t.Fatalf("ParseContext: %v", err) + } + var out struct { + Title string `toml:"title"` + Count int `toml:"count"` + } + if err := UnmarshalContext(context.Background(), in, &out); err != nil { + t.Fatalf("UnmarshalContext: %v", err) + } + if out.Title != "x" || out.Count != 3 { + t.Errorf("out = %#v", out) + } + if err := NewDecoder().DecodeContext(context.Background(), in, &map[string]any{}); err != nil { + t.Fatalf("Decoder.DecodeContext: %v", err) + } +} + +func TestUnmarshalIntToUintOverflow(t *testing.T) { + type C struct { + X uint8 `toml:"x"` + } + var c C + err := Unmarshal([]byte("x = 300\n"), &c) + if err == nil { + t.Fatal("expected overflow error") + } + if !strings.Contains(err.Error(), "overflow") { + t.Errorf("err = %v, want substring 'overflow'", err) + } +} + +func TestUnmarshalIntToUint16Boundary(t *testing.T) { + type C struct { + X uint16 `toml:"x"` + } + // Exactly 65535 fits, 65536 does not. + var ok C + if err := Unmarshal([]byte("x = 65535\n"), &ok); err != nil { + t.Fatalf("65535 should fit uint16, got %v", err) + } + if ok.X != 65535 { + t.Errorf("X = %d, want 65535", ok.X) + } + + var bad C + if err := Unmarshal([]byte("x = 65536\n"), &bad); err == nil { + t.Fatal("65536 must not fit uint16") + } +} + +func TestUnmarshalIntToUint64FitsMaxInt64(t *testing.T) { + type C struct { + X uint64 `toml:"x"` + } + var c C + tok := "x = 9223372036854775807\n" // math.MaxInt64 + if err := Unmarshal([]byte(tok), &c); err != nil { + t.Fatalf("MaxInt64 should fit uint64, got %v", err) + } + if c.X != math.MaxInt64 { + t.Errorf("X = %d, want %d", uint64(c.X), uint64(math.MaxInt64)) + } +} + +func TestUnmarshalNegativeIntToUint(t *testing.T) { + type C struct { + X uint8 `toml:"x"` + } + var c C + err := Unmarshal([]byte("x = -1\n"), &c) + if err == nil { + t.Fatal("expected negative-to-uint error") + } + if !strings.Contains(err.Error(), "negative") { + t.Errorf("err = %v, want substring 'negative'", err) + } +} + +func TestUnmarshalIntToFloat(t *testing.T) { + type C struct { + X float64 `toml:"x"` + } + var c C + if err := Unmarshal([]byte("x = 5\n"), &c); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if c.X != 5.0 { + t.Errorf("X = %v, want 5.0", c.X) + } +} + +func TestUnmarshalIntToStringFails(t *testing.T) { + type C struct { + X string `toml:"x"` + } + var c C + if err := Unmarshal([]byte(`x = 5`), &c); err == nil { + t.Fatal("expected type-mismatch error") + } +} + +func TestUnmarshalStringToIntFails(t *testing.T) { + type C struct { + X int `toml:"x"` + } + var c C + if err := Unmarshal([]byte(`x = "hello"`), &c); err == nil { + t.Fatal("expected type-mismatch error") + } +} + +func TestUnmarshalArrayToScalarFails(t *testing.T) { + type C struct { + X int `toml:"x"` + } + var c C + if err := Unmarshal([]byte("x = [1, 2]\n"), &c); err == nil { + t.Fatal("expected type-mismatch error") + } +} + +func TestUnmarshalArrayOfTablesToScalarFails(t *testing.T) { + type Item struct { + Name string `toml:"name"` + } + type C struct { + X Item `toml:"x"` + } + var c C + if err := Unmarshal([]byte("[[x]]\nname = \"a\"\n"), &c); err == nil { + t.Fatal("expected type-mismatch error") + } +} + +func TestUnmarshalFloatToIntFails(t *testing.T) { + type C struct { + X int `toml:"x"` + } + var c C + if err := Unmarshal([]byte("x = 1.5\n"), &c); err == nil { + t.Fatal("expected type-mismatch error") + } +} + +func TestUnmarshalBoolToIntFails(t *testing.T) { + type C struct { + X int `toml:"x"` + } + var c C + if err := Unmarshal([]byte("x = true\n"), &c); err == nil { + t.Fatal("expected type-mismatch error") + } +} + +func TestUnmarshalNonPointerRejected(t *testing.T) { + var v int + if err := Unmarshal([]byte("x = 1\n"), v); err == nil { + t.Fatal("expected error for non-pointer target") + } +} + +func TestUnmarshalAssignErrorWrapped(t *testing.T) { + // assignStruct wraps inner assignment errors with the key name. + type C struct { + Inner struct { + X int `toml:"x"` + } `toml:"inner"` + } + var c C + err := Unmarshal([]byte("[inner]\nx = \"oops\"\n"), &c) + if err == nil { + t.Fatal("expected an assignment error") + } + if !strings.Contains(err.Error(), "x") { + t.Errorf("err = %v, want it to mention key x", err) + } +} + +func TestUnmarshalAssignMapErrorWrapped(t *testing.T) { + // assignMap wraps inner errors with the key of the bad element. + m := map[string]int{} + err := Unmarshal([]byte("[s]\nx = 1\n"), &m) + if err == nil { + t.Fatal("expected an assignment error") + } + if !strings.Contains(err.Error(), "s") { + t.Errorf("err = %v, want it to mention key 's'", err) + } +} + +func TestUnmarshalAssignSliceErrorWrapped(t *testing.T) { + // assignSlice wraps inner errors with the bad element's index. + type C struct { + Items []int `toml:"items"` + } + var c C + err := Unmarshal([]byte(`items = [1, "oops"]`), &c) + if err == nil { + t.Fatal("expected an assignment error") + } + if !strings.Contains(err.Error(), "[1]") { + t.Errorf("err = %v, want it to mention index [1]", err) + } +} + +func TestUnmarshalAssignTimeToWrongTypeFails(t *testing.T) { + // assign maps time.Time to time.Time only. + type C struct { + T string `toml:"t"` + } + var c C + err := Unmarshal([]byte("t = 2026-01-01T00:00:00Z\n"), &c) + if err == nil { + t.Fatal("expected time-to-string assignment to fail") + } +} + +func TestMarshalerUsesCustomEncoding(t *testing.T) { + // A type that implements Marshaler can be encoded through a wrapping struct. + val := inlineMarshaler(func() (any, error) { + return map[string]any{"k": "v"}, nil + }) + type wrap struct { + Inner inlineMarshaler `toml:"inner"` + } + out, err := Marshal(wrap{Inner: val}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + got := string(out) + if !strings.Contains(got, "[inner]") || !strings.Contains(got, "k = \"v\"") { + t.Errorf("out = %q, want a [inner] table with k = \"v\"", got) + } +} + +type inlineMarshaler func() (any, error) + +func (i inlineMarshaler) MarshalTOML() (any, error) { return i() } + +// --- Unmarshaler ----------------------------------------------------------- + +// receiverSetter implements *Unmarshaler by reshaping a parsed table. +type receiverSetter struct { + Field string + Received any +} + +func (r *receiverSetter) UnmarshalTOML(data any) error { + m, ok := data.(map[string]any) + if !ok { + return fmt.Errorf("interpres: receiverSetter expects table, got %T", data) + } + if v, ok := m["field"].(string); ok { + r.Field = v + } + r.Received = data + return nil +} + +func TestUnmarshalerByPointer(t *testing.T) { + type Cfg struct { + R receiverSetter `toml:"r"` + } + var cfg Cfg + if err := Unmarshal([]byte(`[r] +field = "x" +`), &cfg); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if cfg.R.Field != "x" { + t.Errorf("Field = %q, want \"x\"", cfg.R.Field) + } + if cfg.R.Received == nil { + t.Error("Receiver did not see parsed data") + } +} + +type scalarUnmarshaler struct{ val string } + +func (s *scalarUnmarshaler) UnmarshalTOML(data any) error { + str, ok := data.(string) + if !ok { + return fmt.Errorf("interpres: scalarUnmarshaler expects string, got %T", data) + } + s.val = str + return nil +} + +func TestUnmarshalerReceivesRawScalar(t *testing.T) { + // Place the Unmarshaler-implementing field inside a wrapper struct so + // the decoder dispatches the scalar value to its UnmarshalTOML. + type Cfg struct { + S scalarUnmarshaler `toml:"s"` + } + var cfg Cfg + if err := Unmarshal([]byte(`s = "hello"`), &cfg); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if cfg.S.val != "hello" { + t.Errorf("cfg.S.val = %q, want \"hello\"", cfg.S.val) + } +} + +type failingUnmarshaler struct{} + +func (f *failingUnmarshaler) UnmarshalTOML(_ any) error { return errors.New("boom") } + +func TestUnmarshalerErrorPropagates(t *testing.T) { + type Cfg struct { + F failingUnmarshaler `toml:"f"` + } + var cfg Cfg + if err := Unmarshal([]byte("f = 1"), &cfg); err == nil { + t.Fatal("expected error from UnmarshalTOML") + } else if !strings.Contains(err.Error(), "boom") { + t.Errorf("err = %v, want substring \"boom\"", err) + } +} + +func TestUnmarshalerTakesPrecedenceOverDefault(t *testing.T) { + // Even when the field is a scalar type and the value is a table, + // UnmarshalTOML wins, because the receiver decides. + type Cfg struct { + R receiverSetter `toml:"r"` + } + var cfg Cfg + in := []byte(`[r] +field = "y" +`) + if err := Unmarshal(in, &cfg); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if cfg.R.Field != "y" { + t.Errorf("Field = %q, want \"y\"", cfg.R.Field) + } +} diff --git a/encode.go b/encode.go new file mode 100644 index 0000000..2b2bd33 --- /dev/null +++ b/encode.go @@ -0,0 +1,752 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package interpres + +import ( + "bytes" + "context" + "fmt" + "math" + "reflect" + "slices" + "strconv" + "strings" + "time" + "unicode/utf8" +) + +var ( + localDateTimeType = reflect.TypeFor[LocalDateTime]() + localDateType = reflect.TypeFor[LocalDate]() + localTimeType = reflect.TypeFor[LocalTime]() + timeGoType = reflect.TypeFor[time.Time]() +) + +// encoder produces a TOML document from a Go value via a small intermediate +// representation that preserves the order in which fields were declared. +type encoder struct { + buf bytes.Buffer + ctx context.Context + opts Encoder +} + +func newEncoder() *encoder { return &encoder{} } + +func (e *encoder) bytes() []byte { return e.buf.Bytes() } + +func (e *encoder) checkCtx() error { + if e.ctx == nil { + return nil + } + return e.ctx.Err() +} + +// encode converts v into a TOML document. v must be a struct or a +// map[string]V (or a non-nil pointer to one). +func (e *encoder) encode(v any) error { + if err := e.checkCtx(); err != nil { + return err + } + rv := reflect.ValueOf(v) + if !rv.IsValid() { + return fmt.Errorf("interpres: cannot marshal nil value") + } + if rv.Kind() == reflect.Pointer { + if rv.IsNil() { + return fmt.Errorf("interpres: cannot marshal nil pointer") + } + rv = rv.Elem() + } + doc := &tomlDoc{ctx: e.ctx, opts: e.opts} + switch rv.Kind() { + case reflect.Struct: + if err := buildStructDoc(rv, doc, ""); err != nil { + return err + } + case reflect.Map: + if err := buildMapDoc(rv, doc, ""); err != nil { + return err + } + default: + return fmt.Errorf("interpres: top-level value must be a struct or map[string]V, got %s", rv.Type()) + } + return e.emitDoc(doc, nil) +} + +// --- intermediate representation ----------------------------------------- + +// entryKind discriminates the three forms an entry in a tomlDoc may take. +type entryKind int + +const ( + entryScalar entryKind = iota + entryTable + entryArray +) + +// entry is one binding in a tomlDoc. entries live in a single slice in the +// order they were added; emission either walks that order directly +// (Encoder with GroupByKind(false)) or partitions by kind first +// (Encoder with GroupByKind(true), the default). +type entry struct { + kind entryKind + key string + val any // entryScalar + doc *tomlDoc // entryTable + docs []*tomlDoc +} + +// tomlDoc holds the entries of one TOML table in declaration order. +type tomlDoc struct { + entries []entry + ctx context.Context // inherited from encoder; nil-safe + opts Encoder // inherited from encoder; options drive emit-time behaviour +} + +func (d *tomlDoc) checkCtx() error { + if d.ctx == nil { + return nil + } + return d.ctx.Err() +} + +func (d *tomlDoc) addScalar(key string, val any) { + d.entries = append(d.entries, entry{kind: entryScalar, key: key, val: val}) +} + +func (d *tomlDoc) addTable(key string, sub *tomlDoc) { + d.entries = append(d.entries, entry{kind: entryTable, key: key, doc: sub}) +} + +func (d *tomlDoc) addArray(key string, subs []*tomlDoc) { + d.entries = append(d.entries, entry{kind: entryArray, key: key, docs: subs}) +} + +// partitionedEntries returns the entries grouped by kind, preserving each +// group's relative order. The only allocation is the three slice headers. +func (d *tomlDoc) partitionedEntries() (scalars []entry, tables []entry, arrays []entry) { + for _, e := range d.entries { + switch e.kind { + case entryScalar: + scalars = append(scalars, e) + case entryTable: + tables = append(tables, e) + case entryArray: + arrays = append(arrays, e) + } + } + return +} + +// --- reflection walk: struct --------------------------------------------- + +func buildStructDoc(v reflect.Value, doc *tomlDoc, ctx string) error { + t := v.Type() + for i := range t.NumField() { + if i%ctxCheckInterval == 0 { + if err := doc.checkCtx(); err != nil { + return err + } + } + f := t.Field(i) + if f.PkgPath != "" { + continue + } + if f.Anonymous { + tag, _ := f.Tag.Lookup("toml") + if tag == "-" { + continue + } + if tag == "" { + fv := followPtr(v.Field(i)) + if !fv.IsValid() { + continue + } + switch fv.Kind() { + case reflect.Struct: + if isScalarStruct(fv.Type()) { + name := strings.ToLower(f.Name) + if err := doc.appendScalar(name, fv.Interface(), ctx); err != nil { + return err + } + continue + } + if err := buildStructDoc(fv, doc, ctx); err != nil { + return err + } + continue + case reflect.Map: + if err := buildMapDoc(fv, doc, ctx); err != nil { + return err + } + continue + } + } + } + name := fieldName(f) + if name == "-" { + continue + } + if err := addField(doc, name, v.Field(i), ctx); err != nil { + return err + } + } + return nil +} + +// fieldName returns the TOML key for a struct field, honouring the `toml` +// tag (name or `-`) and falling back to a lower-cased field name. +func fieldName(f reflect.StructField) string { + if tag, ok := f.Tag.Lookup("toml"); ok { + name, _, _ := strings.Cut(tag, ",") + if name == "-" { + return "-" + } + if name != "" { + return name + } + } + return strings.ToLower(f.Name) +} + +// --- reflection walk: map ------------------------------------------------ + +func buildMapDoc(v reflect.Value, doc *tomlDoc, ctx string) error { + if v.Type().Key().Kind() != reflect.String { + return fmt.Errorf("interpres: map key must be string, got %s", v.Type().Key()) + } + keys := v.MapKeys() + slices.SortFunc(keys, func(a, b reflect.Value) int { + return strings.Compare(a.String(), b.String()) + }) + for i, k := range keys { + if i%ctxCheckInterval == 0 { + if err := doc.checkCtx(); err != nil { + return err + } + } + if err := addField(doc, k.String(), v.MapIndex(k), ctx); err != nil { + return err + } + } + return nil +} + +// --- reflection walk: field dispatch ------------------------------------- + +func addField(doc *tomlDoc, name string, v reflect.Value, ctx string) error { + if v.CanInterface() { + if m, ok := v.Interface().(Marshaler); ok { + mv, err := m.MarshalTOML() + if err != nil { + return fmt.Errorf("interpres: %s.%s: %w", ctx, name, err) + } + v = reflect.ValueOf(mv) + } + } + v = followPtr(v) + if !v.IsValid() { + return nil + } + if v.Kind() == reflect.Interface { + if v.IsNil() { + return nil + } + v = v.Elem() + } + switch v.Kind() { + case reflect.Struct: + if isScalarStruct(v.Type()) { + return doc.appendScalar(name, v.Interface(), ctx) + } + return addSubTable(doc, name, v, ctx) + case reflect.Map: + return addSubTable(doc, name, v, ctx) + case reflect.Slice, reflect.Array: + return addArrayValue(doc, name, v, ctx) + default: + val, err := normaliseValue(v) + if err != nil { + return fmt.Errorf("interpres: %s.%s: %w", ctx, name, err) + } + return doc.appendScalar(name, val, ctx) + } +} + +// appendScalar wraps addScalar with a uniform error path. +func (d *tomlDoc) appendScalar(name string, val any, ctx string) error { + d.addScalar(name, val) + return nil +} + +func addSubTable(doc *tomlDoc, name string, v reflect.Value, ctx string) error { + sub := &tomlDoc{ctx: doc.ctx, opts: doc.opts} + switch v.Kind() { + case reflect.Struct: + if err := buildStructDoc(v, sub, joinKey(ctx, name)); err != nil { + return err + } + case reflect.Map: + if err := buildMapDoc(v, sub, joinKey(ctx, name)); err != nil { + return err + } + } + doc.addTable(name, sub) + return nil +} + +func addArrayValue(doc *tomlDoc, name string, v reflect.Value, ctx string) error { + if v.Kind() == reflect.Slice && v.IsNil() { + // A nil slice has no explicit representation in TOML, so it is skipped. + return nil + } + n := v.Len() + if n == 0 { + if isTableElementType(v.Type().Elem()) { + // Empty array of tables has no valid TOML form, so it is skipped. + return nil + } + if doc.opts.omitEmptyArrays { + return nil + } + return doc.appendScalar(name, []any{}, ctx) + } + + if isTableElementValue(v.Index(0)) { + subs := make([]*tomlDoc, n) + for i := range n { + if i%ctxCheckInterval == 0 { + if err := doc.checkCtx(); err != nil { + return err + } + } + ev := followPtr(v.Index(i)) + if !ev.IsValid() { + return fmt.Errorf("interpres: %s.%s[%d]: nil element", ctx, name, i) + } + sub := &tomlDoc{ctx: doc.ctx, opts: doc.opts} + switch ev.Kind() { + case reflect.Struct: + if isScalarStruct(ev.Type()) { + return fmt.Errorf("interpres: %s.%s[%d]: heterogeneous array contains scalar", ctx, name, i) + } + if err := buildStructDoc(ev, sub, joinKey(ctx, fmt.Sprintf("%s[%d]", name, i))); err != nil { + return err + } + case reflect.Map: + if err := buildMapDoc(ev, sub, joinKey(ctx, fmt.Sprintf("%s[%d]", name, i))); err != nil { + return err + } + default: + return fmt.Errorf("interpres: %s.%s: heterogeneous array, expected table", ctx, name) + } + subs[i] = sub + } + doc.addArray(name, subs) + return nil + } + + // Regular array of scalars. + items := make([]any, n) + for i := range n { + if i%ctxCheckInterval == 0 { + if err := doc.checkCtx(); err != nil { + return err + } + } + ev := followPtr(v.Index(i)) + if !ev.IsValid() { + return fmt.Errorf("interpres: %s.%s[%d]: nil element", ctx, name, i) + } + if ev.CanInterface() { + if m, ok := ev.Interface().(Marshaler); ok { + mv, err := m.MarshalTOML() + if err != nil { + return fmt.Errorf("interpres: %s.%s[%d]: %w", ctx, name, i, err) + } + ev = reflect.ValueOf(mv) + ev = followPtr(ev) + } + } + val, err := normaliseValue(ev) + if err != nil { + return fmt.Errorf("interpres: %s.%s[%d]: %w", ctx, name, i, err) + } + items[i] = val + } + return doc.appendScalar(name, items, ctx) +} + +// normaliseValue converts a reflect.Value into one of the canonical scalar or +// nested-array representations the emitter understands. Slices and arrays are +// recursively normalised so that nested arrays (e.g. [][]int) work. +func normaliseValue(v reflect.Value) (any, error) { + if v.CanInterface() { + if m, ok := v.Interface().(Marshaler); ok { + return m.MarshalTOML() + } + } + switch v.Kind() { + case reflect.String: + return v.String(), nil + case reflect.Bool: + return v.Bool(), nil + case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64: + return v.Int(), nil + case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64: + u := v.Uint() + if u > math.MaxInt64 { + return nil, fmt.Errorf("unsigned value %d overflows int64", u) + } + return int64(u), nil + case reflect.Float32, reflect.Float64: + return v.Float(), nil + case reflect.Slice, reflect.Array: + items := make([]any, v.Len()) + for i := range v.Len() { + val, err := normaliseValue(v.Index(i)) + if err != nil { + return nil, fmt.Errorf("[%d]: %w", i, err) + } + items[i] = val + } + return items, nil + } + if !v.IsValid() { + return nil, fmt.Errorf("invalid value") + } + return nil, fmt.Errorf("cannot encode %s", v.Type()) +} + +// followPtr unwraps pointer and interface layers. Returns a zero Value if a +// nil pointer or nil interface is encountered. +func followPtr(v reflect.Value) reflect.Value { + for { + switch v.Kind() { + case reflect.Pointer, reflect.Interface: + if v.IsNil() { + return reflect.Value{} + } + v = v.Elem() + continue + } + return v + } +} + +// isScalarStruct reports whether t is a struct type that the encoder treats +// as a TOML scalar (time.Time, LocalDateTime, LocalDate, LocalTime). +func isScalarStruct(t reflect.Type) bool { + return t == timeGoType || isLocalDateType(t) +} + +func isLocalDateType(t reflect.Type) bool { + return t == localDateTimeType || t == localDateType || t == localTimeType +} + +func isTableElementType(t reflect.Type) bool { + switch t.Kind() { + case reflect.Struct: + return !isScalarStruct(t) + case reflect.Map: + return t.Key().Kind() == reflect.String + } + return false +} + +func isTableElementValue(v reflect.Value) bool { + v = followPtr(v) + if !v.IsValid() { + return false + } + return isTableElementType(v.Type()) +} + +func joinKey(ctx, name string) string { + if ctx == "" { + return name + } + return ctx + "." + name +} + +// --- emission ------------------------------------------------------------ + +// writeBlankLine writes a single newline before a table or array-of-tables +// header so the output has a blank line between sections, unless the buffer +// is empty (i.e. this is the very first header). +func (e *encoder) writeBlankLine() { + if e.buf.Len() == 0 { + return + } + e.buf.WriteByte('\n') +} + +func (e *encoder) emitDoc(doc *tomlDoc, prefix []string) error { + if e.opts.groupByKind { + scalars, tables, arrays := doc.partitionedEntries() + for _, kv := range scalars { + if err := e.writeKV(kv.key, kv.val); err != nil { + return err + } + } + for _, t := range tables { + path := append(append([]string{}, prefix...), t.key) + e.writeBlankLine() + e.buf.WriteByte('[') + writeKeyPath(&e.buf, path) + e.buf.WriteString("]\n") + if err := e.emitDoc(t.doc, path); err != nil { + return err + } + } + for _, a := range arrays { + path := append(append([]string{}, prefix...), a.key) + for _, sub := range a.docs { + e.writeBlankLine() + e.buf.WriteString("[[") + writeKeyPath(&e.buf, path) + e.buf.WriteString("]]\n") + if err := e.emitDoc(sub, path); err != nil { + return err + } + } + } + return nil + } + + // Preserve declaration order. Scalars and table/array headers may now + // interleave, which means each table/array header must include only its + // own section content; the emitter still writes sub-documents as separate + // nested blocks, so a "" sub-keyed scalar following a header for the same + // section is impossible in practice (struct fields are visited in order). + for _, ent := range doc.entries { + switch ent.kind { + case entryScalar: + if err := e.writeKV(ent.key, ent.val); err != nil { + return err + } + case entryTable: + path := append(append([]string{}, prefix...), ent.key) + e.writeBlankLine() + e.buf.WriteByte('[') + writeKeyPath(&e.buf, path) + e.buf.WriteString("]\n") + if err := e.emitDoc(ent.doc, path); err != nil { + return err + } + case entryArray: + path := append(append([]string{}, prefix...), ent.key) + for _, sub := range ent.docs { + e.writeBlankLine() + e.buf.WriteString("[[") + writeKeyPath(&e.buf, path) + e.buf.WriteString("]]\n") + if err := e.emitDoc(sub, path); err != nil { + return err + } + } + } + } + return nil +} + +func (e *encoder) writeKV(key string, val any) error { + if !utf8.ValidString(key) { + return fmt.Errorf("interpres: key %q is not valid UTF-8", key) + } + e.writeKey(key) + e.buf.WriteString(" = ") + if err := e.writeValue(val); err != nil { + return err + } + e.buf.WriteByte('\n') + return nil +} + +func writeKeyPath(buf *bytes.Buffer, path []string) { + for i, p := range path { + if i > 0 { + buf.WriteByte('.') + } + if isBareKey(p) { + buf.WriteString(p) + continue + } + writeQuotedString(buf, p) + } +} + +func (e *encoder) writeKey(key string) { + if isBareKey(key) { + e.buf.WriteString(key) + return + } + writeQuotedString(&e.buf, key) +} + +// writeQuotedString writes s as a TOML basic string (double-quoted) to buf. +// Returns an error only if s is not valid UTF-8; invalid byte sequences +// within a valid UTF-8 string are encoded as \ufffd replacement characters. +func writeQuotedString(buf *bytes.Buffer, s string) error { + if !utf8.ValidString(s) { + return fmt.Errorf("interpres: string is not valid UTF-8") + } + buf.WriteByte('"') + for i := 0; i < len(s); { + r, size := utf8.DecodeRuneInString(s[i:]) + if r == utf8.RuneError && size == 1 { + buf.WriteString(`\ufffd`) + i++ + continue + } + i += size + writeEscapedRune(buf, r) + } + buf.WriteByte('"') + return nil +} + +// writeEscapedRune writes a single rune to buf, escaping it as required by +// TOML basic-string rules. +func writeEscapedRune(buf *bytes.Buffer, r rune) { + switch r { + case '\\': + buf.WriteString(`\\`) + case '"': + buf.WriteString(`\"`) + case '\b': + buf.WriteString(`\b`) + case '\t': + buf.WriteString(`\t`) + case '\n': + buf.WriteString(`\n`) + case '\f': + buf.WriteString(`\f`) + case '\r': + buf.WriteString(`\r`) + default: + if r < 0x20 || r == 0x7f { + fmt.Fprintf(buf, `\u%04X`, r) + } else { + buf.WriteRune(r) + } + } +} + +func isBareKey(s string) bool { + if s == "" { + return false + } + for i := range len(s) { + c := s[i] + if !((c >= 'A' && c <= 'Z') || (c >= 'a' && c <= 'z') || (c >= '0' && c <= '9') || c == '_' || c == '-') { + return false + } + } + return true +} + +func (e *encoder) writeValue(val any) error { + switch v := val.(type) { + case string: + return e.writeStringVal(v) + case bool: + e.buf.WriteString(strconv.FormatBool(v)) + return nil + case int64: + e.buf.WriteString(strconv.FormatInt(v, 10)) + return nil + case float64: + return e.writeFloat(v) + case time.Time: + e.buf.WriteString(v.Format(time.RFC3339Nano)) + return nil + case LocalDateTime: + e.buf.WriteString(v.String()) + return nil + case LocalDate: + e.buf.WriteString(v.String()) + return nil + case LocalTime: + e.buf.WriteString(v.String()) + return nil + case []any: + e.buf.WriteByte('[') + for i, item := range v { + if i > 0 { + e.buf.WriteString(", ") + } + if err := e.writeValue(item); err != nil { + return err + } + } + e.buf.WriteByte(']') + return nil + case nil: + return fmt.Errorf("interpres: cannot encode nil value") + default: + return fmt.Errorf("interpres: cannot encode %T", val) + } +} + +func (e *encoder) writeStringVal(s string) error { + if e.opts.literalMultilineAt > 0 && strings.ContainsRune(s, '\n') && len(s) >= e.opts.literalMultilineAt { + return writeLiteralMultilineString(&e.buf, s) + } + return writeQuotedString(&e.buf, s) +} + +// writeLiteralMultilineString writes s as a TOML literal multi-line string, +// surrounded by triple single quotes. The opening delimiter is followed by a +// newline that the reader trims, so we always include one. The closing +// delimiter sits on its own line; if the value does not end in a newline, one +// is inserted before the closing delimiter. +func writeLiteralMultilineString(buf *bytes.Buffer, s string) error { + if !utf8.ValidString(s) { + return fmt.Errorf("interpres: string is not valid UTF-8") + } + buf.WriteString("'''\n") + buf.WriteString(s) + if !strings.HasSuffix(s, "\n") { + buf.WriteByte('\n') + } + buf.WriteString("'''") + return nil +} + +func (e *encoder) writeFloat(v float64) error { + switch { + case math.IsNaN(v): + e.buf.WriteString("nan") + case math.IsInf(v, 1): + e.buf.WriteString("inf") + case math.IsInf(v, -1): + e.buf.WriteString("-inf") + case v == 0: + // Normalise negative zero to positive zero (TOML has no -0). + e.buf.WriteString("0.0") + default: + s := strconv.FormatFloat(v, 'g', -1, 64) + // TOML forbids leading zeros in the exponent digits. + if idx := strings.LastIndexAny(s, "eE"); idx >= 0 { + mant := s[:idx] + exp := s[idx+1:] // e.g. "+06", "-05" + sign := "" + if len(exp) > 0 && (exp[0] == '+' || exp[0] == '-') { + sign = string(exp[0]) + exp = exp[1:] + } + exp = strings.TrimLeft(exp, "0") + if exp == "" { + exp = "0" + } + s = mant + "e" + sign + exp + } + if !strings.ContainsAny(s, ".eE") { + s += ".0" + } + e.buf.WriteString(s) + } + return nil +} diff --git a/encode_test.go b/encode_test.go new file mode 100644 index 0000000..9c87143 --- /dev/null +++ b/encode_test.go @@ -0,0 +1,1010 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package interpres + +import ( + "bytes" + "context" + "errors" + "math" + "reflect" + "strings" + "testing" + "time" +) + +func TestMarshalScalars(t *testing.T) { + type Cfg struct { + Title string `toml:"title"` + Count int `toml:"count"` + Unsigned uint64 `toml:"unsigned"` + Ratio float64 `toml:"ratio"` + Enabled bool `toml:"enabled"` + Disabled bool `toml:"disabled"` + } + out, err := Marshal(Cfg{ + Title: "demo", + Count: 42, + Unsigned: 99, + Ratio: 3.14, + Enabled: true, + Disabled: false, + }) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "title = \"demo\"\ncount = 42\nunsigned = 99\nratio = 3.14\nenabled = true\ndisabled = false\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalFloatSpecials(t *testing.T) { + type Cfg struct { + PosInf float64 `toml:"pos_inf"` + NegInf float64 `toml:"neg_inf"` + NaN float64 `toml:"nan"` + Zero float64 `toml:"zero"` + IntVal float64 `toml:"int_val"` + NegZ float64 `toml:"neg_zero"` + } + out, err := Marshal(Cfg{ + PosInf: math.Inf(1), + NegInf: math.Inf(-1), + NaN: math.NaN(), + Zero: 0, + IntVal: 7, + NegZ: math.Copysign(0, -1), + }) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "pos_inf = inf\nneg_inf = -inf\nnan = nan\nzero = 0.0\nint_val = 7.0\nneg_zero = 0.0\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalFloatNormalizesNegativeZero(t *testing.T) { + // TOML has no -0; the emitter must normalise negative zero to "0.0". + type Cfg struct { + Z float64 `toml:"z"` + } + out, err := Marshal(Cfg{Z: math.Copysign(0, -1)}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "z = 0.0\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalContextHonoursCancellation(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + type C struct { + A int `toml:"a"` + } + if _, err := MarshalContext(ctx, C{A: 1}); !errors.Is(err, context.Canceled) { + t.Fatalf("MarshalContext returned %v, want context.Canceled", err) + } + if _, err := NewEncoder().MarshalContext(ctx, C{A: 1}); !errors.Is(err, context.Canceled) { + t.Fatalf("Encoder.MarshalContext returned %v, want context.Canceled", err) + } +} + +func TestEncoderGroupByKindDefault(t *testing.T) { + // NewEncoder must default to GroupByKind=true so legacy callers keep the + // scalars-first ordering. + type Cfg struct { + Name string `toml:"name"` + S struct { + Host string `toml:"host"` + } `toml:"s"` + } + out, err := NewEncoder().Marshal(Cfg{Name: "x", S: struct { + Host string `toml:"host"` + }{Host: "h"}}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "name = \"x\"\n\n[s]\nhost = \"h\"\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestEncoderGroupByKindFalsePreservesOrder(t *testing.T) { + type Inner struct { + Host string `toml:"host"` + } + type Cfg struct { + Name string `toml:"name"` + Server Inner `toml:"server"` + Debug bool `toml:"debug"` + } + in := Cfg{ + Name: "x", + Server: Inner{Host: "h"}, + Debug: true, + } + out, err := NewEncoder().GroupByKind(false).Marshal(in) + if err != nil { + t.Fatalf("marshal: %v", err) + } + // With GroupByKind(false) the encoder walks entries in declaration order. + // The output is still parseable, but a scalar that follows a header is + // parsed as a sub-table key. That is the user's trade-off; see + // docs/API.md. + want := "name = \"x\"\n\n[server]\nhost = \"h\"\ndebug = true\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestEncoderGroupByKindTrueDefaultOrder(t *testing.T) { + // The default (GroupByKind=true) must lift the trailing scalar ahead of + // the [server] block so the document round-trips losslessly. + type Inner struct { + Host string `toml:"host"` + } + type Cfg struct { + Name string `toml:"name"` + Server Inner `toml:"server"` + Debug bool `toml:"debug"` + } + in := Cfg{ + Name: "x", + Server: Inner{Host: "h"}, + Debug: true, + } + out, err := NewEncoder().Marshal(in) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "name = \"x\"\ndebug = true\n\n[server]\nhost = \"h\"\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestEncoderOmitEmptyArrays(t *testing.T) { + type Cfg struct { + Tags []string `toml:"tags"` + Secrets []string `toml:"secrets"` + } + out, err := NewEncoder().OmitEmptyArrays().Marshal(Cfg{ + Tags: []string{"a", "b"}, + Secrets: []string{}, + }) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "tags = [\"a\", \"b\"]\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestEncoderDefaultEmitsEmptyArray(t *testing.T) { + type Cfg struct { + Tags []string `toml:"tags"` + } + out, err := NewEncoder().Marshal(Cfg{Tags: []string{}}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "tags = []\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestEncoderOmitEmptyArrayOfTablesStillSkipped(t *testing.T) { + type Item struct { + Name string `toml:"name"` + } + type Cfg struct { + Title string `toml:"title"` + Items []Item `toml:"items"` + } + out, err := NewEncoder().OmitEmptyArrays().Marshal(Cfg{ + Title: "demo", + Items: nil, + }) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "title = \"demo\"\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestEncoderUseLiteralMultiline(t *testing.T) { + type Cfg struct { + Long string `toml:"long"` + } + long := strings.Repeat("a", 50) + "\nline two\nline three" + out, err := NewEncoder().UseLiteralMultiline(20).Marshal(Cfg{Long: long}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "long = '''\n" + long + "\n'''\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestEncoderUseLiteralMultilineBelowThreshold(t *testing.T) { + // A multi-line value shorter than the threshold must remain escaped. + type Cfg struct { + Short string `toml:"short"` + } + out, err := NewEncoder().UseLiteralMultiline(1000).Marshal(Cfg{Short: "one\ntwo"}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "short = \"one\\ntwo\"\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestEncoderUseLiteralMultilineThresholdZero(t *testing.T) { + // UseLiteralMultiline(0) disables the literal form entirely. + type Cfg struct { + S string `toml:"s"` + } + out, err := NewEncoder().UseLiteralMultiline(0).Marshal(Cfg{S: "a\nb\nc\nd"}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + if !bytes.HasPrefix(out, []byte("s = \"")) { + t.Errorf("output mismatch, expected basic quoted form:\ngot: %q", out) + } +} + +// marshalerFunc adapts a plain function value to the Marshaler interface. +// Tests use it to express "this field produces this TOML value" without a +// dedicated struct definition. +type marshalerFunc func() (any, error) + +func (f marshalerFunc) MarshalTOML() (any, error) { return f() } + +// failingMarshalerFunc invokes MarshalTOML to a fixed error; it lets us check +// that a MarshalTOML failure propagates back to Marshal. +type failingMarshalerFunc struct{} + +func (failingMarshalerFunc) MarshalTOML() (any, error) { + return nil, errors.New("oops") +} + +func TestMarshalerReturningTime(t *testing.T) { + // A Marshaler may return a date-time scalar; the encoder must emit it + // using its canonical form. + when := time.Date(2026, 6, 26, 10, 0, 0, 0, time.UTC) + type Cfg struct { + M marshalerFunc `toml:"m"` + } + out, err := Marshal(Cfg{M: marshalerFunc(func() (any, error) { return when, nil })}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "m = 2026-06-26T10:00:00Z\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalerReturningDifferentStruct(t *testing.T) { + // A Marshaler returning a struct (not a scalar) is treated as a sub-table + // by the encoder. + type Inner struct { + V string `toml:"v"` + } + type Cfg struct { + P marshalerFunc `toml:"p"` + } + out, err := Marshal(Cfg{P: marshalerFunc(func() (any, error) { + return Inner{V: "x"}, nil + })}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "[p]\nv = \"x\"\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalerReturningSliceOfMaps(t *testing.T) { + // A Marshaler returning []map[string]any becomes an array of tables. + type Cfg struct { + Items marshalerFunc `toml:"items"` + } + out, err := Marshal(Cfg{Items: marshalerFunc(func() (any, error) { + return []map[string]any{ + {"k": "a"}, + {"k": "b"}, + }, nil + })}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "[[items]]\nk = \"a\"\n\n[[items]]\nk = \"b\"\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalerErrorPropagates(t *testing.T) { + type Cfg struct { + F failingMarshalerFunc `toml:"f"` + } + if _, err := Marshal(Cfg{F: failingMarshalerFunc{}}); err == nil { + t.Fatal("expected an error from MarshalTOML") + } else if !strings.Contains(err.Error(), "oops") { + t.Errorf("err = %v, want substring \"oops\"", err) + } +} + +func TestMarshalEmbeddedScalarStruct(t *testing.T) { + // A field declared directly as a scalar-struct type (here LocalDateTime) + // must be encoded as a TOML scalar at the parent level, not rendered as + // a sub-table. + ldt := LocalDateTime{Time: time.Date(2026, 6, 26, 0, 0, 0, 0, time.UTC)} + type Cfg struct { + Name string `toml:"name"` + S LocalDateTime `toml:"s"` + } + out, err := Marshal(Cfg{Name: "x", S: ldt}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "name = \"x\"\ns = 2026-06-26T00:00:00\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestEncoderChainedOptions(t *testing.T) { + // All chainable options combined; verify they compose without errors. + type Inner struct { + V string `toml:"v"` + } + type Cfg struct { + S string `toml:"s"` + I Inner `toml:"i"` + } + long := strings.Repeat("x", 200) + out, err := NewEncoder(). + GroupByKind(false). + OmitEmptyArrays(). + UseLiteralMultiline(50). + Marshal(Cfg{S: "short", I: Inner{V: long}}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + _ = out // success path is enough; per-option correctness is exercised above. +} + +func TestMarshalStringEscapes(t *testing.T) { + cases := []struct { + name string + in string + want string // the TOML scalar value (without "s = " prefix) + }{ + {"plain", "hello", `"hello"`}, + {"quote", `say "hi"`, `"say \"hi\""`}, + {"backslash", `a\b`, `"a\\b"`}, + {"newline", "line1\nline2", `"line1\nline2"`}, + {"tab", "col1\tcol2", `"col1\tcol2"`}, + {"cr", "line\rmore", `"line\rmore"`}, + {"control", "a\x01b", `"a\u0001b"`}, + {"unicode", "\u201csmart\u201d", `"“smart”"`}, // printable unicode; not escaped + {"empty", "", `""`}, + {"slash_only", "a/b", `"a/b"`}, + } + for _, c := range cases { + out, err := Marshal(struct { + S string `toml:"s"` + }{S: c.in}) + if err != nil { + t.Fatalf("%s: marshal: %v", c.name, err) + } + got := strings.TrimSuffix(string(out), "\n") + want := "s = " + c.want + if got != want { + t.Errorf("%s:\ngot: %s\nwant: %s", c.name, got, want) + } + } +} + +func TestMarshalDateTime(t *testing.T) { + type Cfg struct { + Offset time.Time `toml:"offset"` + Local LocalDateTime `toml:"local"` + Day LocalDate `toml:"day"` + Clock LocalTime `toml:"clock"` + } + out, err := Marshal(Cfg{ + Offset: time.Date(2026, 6, 26, 10, 0, 0, 0, time.UTC), + Local: LocalDateTime{Time: time.Date(2026, 6, 26, 7, 32, 0, 0, time.UTC)}, + Day: LocalDate{Time: time.Date(2026, 6, 26, 0, 0, 0, 0, time.UTC)}, + Clock: LocalTime{Time: time.Date(0, 1, 1, 7, 32, 0, 0, time.UTC)}, + }) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "offset = 2026-06-26T10:00:00Z\nlocal = 2026-06-26T07:32:00\nday = 2026-06-26\nclock = 07:32:00\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalDateTimeFractional(t *testing.T) { + out, err := Marshal(struct { + LDT LocalDateTime `toml:"ldt"` + LT LocalTime `toml:"lt"` + }{ + LDT: LocalDateTime{Time: time.Date(2026, 6, 26, 7, 32, 0, 123456789, time.UTC)}, + LT: LocalTime{Time: time.Date(0, 1, 1, 7, 32, 0, 123, time.UTC)}, + }) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "ldt = 2026-06-26T07:32:00.123456789\nlt = 07:32:00.000000123\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalArraysOfScalars(t *testing.T) { + type Cfg struct { + Tags []string `toml:"tags"` + Ports []int `toml:"ports"` + Mixed []any `toml:"mixed"` + Empty []int `toml:"empty"` + EmptyS []string `toml:"empty_s"` + } + out, err := Marshal(Cfg{ + Tags: []string{"a", "b"}, + Ports: []int{80, 443}, + Mixed: []any{int64(1), "x", true}, + Empty: nil, + EmptyS: []string{}, + }) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "tags = [\"a\", \"b\"]\nports = [80, 443]\nmixed = [1, \"x\", true]\nempty_s = []\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalNestedArrays(t *testing.T) { + type Cfg struct { + Matrix [][]int `toml:"matrix"` + Words [][]string `toml:"words"` + } + out, err := Marshal(Cfg{ + Matrix: [][]int{{1, 2}, {3, 4}}, + Words: [][]string{{"a", "b"}, {"c"}}, + }) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "matrix = [[1, 2], [3, 4]]\nwords = [[\"a\", \"b\"], [\"c\"]]\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalFloatExponentNoLeadingZero(t *testing.T) { + // strconv.FormatFloat with 'g' would produce "1e+06" (leading zero in + // exponent). The encoder must strip it so the output is "1e+6". + type Cfg struct { + Large float64 `toml:"large"` + Small float64 `toml:"small"` + } + out, err := Marshal(Cfg{Large: 1e6, Small: 1e-5}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + // Parse to check the output is valid TOML (for a strict parser that + // rejects leading zeros in exponents). + if _, err := Parse(out); err != nil { + t.Fatalf("marshalled output is not valid TOML:\n%s\nerror: %v", out, err) + } + if string(out) != "large = 1e+6\nsmall = 1e-5\n" { + t.Errorf("output mismatch:\ngot: %q", out) + } +} + +func TestMarshalStructAsTable(t *testing.T) { + type Server struct { + Host string `toml:"host"` + Port int `toml:"port"` + } + type Cfg struct { + Title string `toml:"title"` + Server Server `toml:"server"` + } + out, err := Marshal(Cfg{ + Title: "demo", + Server: Server{Host: "127.0.0.1", Port: 9090}, + }) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "title = \"demo\"\n\n[server]\nhost = \"127.0.0.1\"\nport = 9090\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalArrayOfTables(t *testing.T) { + type Item struct { + Name string `toml:"name"` + Qty int `toml:"qty"` + } + type Cfg struct { + Items []Item `toml:"items"` + } + out, err := Marshal(Cfg{ + Items: []Item{ + {Name: "a", Qty: 1}, + {Name: "b", Qty: 2}, + }, + }) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "[[items]]\nname = \"a\"\nqty = 1\n\n[[items]]\nname = \"b\"\nqty = 2\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalEmptyArrayOfTablesIsSkipped(t *testing.T) { + type Item struct { + Name string `toml:"name"` + } + type Cfg struct { + Title string `toml:"title"` + Items []Item `toml:"items"` + } + out, err := Marshal(Cfg{ + Title: "demo", + Items: nil, + }) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "title = \"demo\"\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalNestedTablesAndArrays(t *testing.T) { + type SMTP struct { + Host string `toml:"host"` + Port int `toml:"port"` + } + type Form struct { + Name string `toml:"name"` + SMTP SMTP `toml:"smtp"` + } + type Cfg struct { + Port int `toml:"port"` + Forms []Form `toml:"forms"` + } + out, err := Marshal(Cfg{ + Port: 8080, + Forms: []Form{ + {Name: "contact", SMTP: SMTP{Host: "h1", Port: 587}}, + {Name: "feedback", SMTP: SMTP{Host: "h2", Port: 25}}, + }, + }) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "port = 8080\n\n[[forms]]\nname = \"contact\"\n\n[forms.smtp]\nhost = \"h1\"\nport = 587\n\n[[forms]]\nname = \"feedback\"\n\n[forms.smtp]\nhost = \"h2\"\nport = 25\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalStructTags(t *testing.T) { + type Cfg struct { + Keep string `toml:"keep"` + Rename string `toml:"renamed"` + Skip string `toml:"-"` + Untagged string + } + out, err := Marshal(Cfg{ + Keep: "k", Rename: "r", Skip: "s", Untagged: "u", + }) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "keep = \"k\"\nrenamed = \"r\"\nuntagged = \"u\"\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalEmbeddedStructPromoted(t *testing.T) { + type Base struct { + ID int `toml:"id"` + } + type Derived struct { + Base + Name string `toml:"name"` + } + out, err := Marshal(Derived{ID: 1, Name: "x"}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "id = 1\nname = \"x\"\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalEmbeddedStructAsTable(t *testing.T) { + type Inner struct { + Host string `toml:"host"` + } + type Cfg struct { + Inner Inner `toml:"inner"` + Name string `toml:"name"` + } + out, err := Marshal(Cfg{Inner: Inner{Host: "h"}, Name: "n"}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "name = \"n\"\n\n[inner]\nhost = \"h\"\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalMapKeysSorted(t *testing.T) { + m := map[string]any{ + "zeta": 1, + "alpha": 2, + "mu": 3, + } + out, err := Marshal(m) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "alpha = 2\nmu = 3\nzeta = 1\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalMapWithSubMap(t *testing.T) { + m := map[string]any{ + "meta": map[string]any{"x": 1, "y": 2}, + "a": "z", + } + out, err := Marshal(m) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "a = \"z\"\n\n[meta]\nx = 1\ny = 2\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalMarshaler(t *testing.T) { + type Port int + type Cfg struct { + P Port `toml:"p"` + } + out, err := Marshal(Cfg{P: 8080}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "p = 8080\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalMarshalerReturningScalar(t *testing.T) { + type Wrapped struct { + Value string `toml:"value"` + } + type Alias struct{} + out, err := Marshal(struct { + W Wrapped `toml:"w"` + }{W: Wrapped{Value: "hello"}}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "[w]\nvalue = \"hello\"\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } + _ = Alias{} +} + +func TestMarshalMarshalerReturningDifferentShape(t *testing.T) { + out, err := Marshal(struct { + C Custom `toml:"c"` + }{C: Custom{tag: "x"}}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + // Custom returns a string from MarshalTOML. + want := "c = \"x\"\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalNilPointerFieldSkipped(t *testing.T) { + type Cfg struct { + Name string `toml:"name"` + Hidden *string `toml:"hidden"` + } + out, err := Marshal(Cfg{Name: "x"}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "name = \"x\"\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalNonNilPointerFollowed(t *testing.T) { + v := "v" + type Cfg struct { + Name string `toml:"name"` + Hidden *string `toml:"hidden"` + } + out, err := Marshal(Cfg{Name: "n", Hidden: &v}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "name = \"n\"\nhidden = \"v\"\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalTopLevelMustBeStructOrMap(t *testing.T) { + if _, err := Marshal(42); err == nil { + t.Errorf("expected error marshalling int at top level") + } + if _, err := Marshal("hello"); err == nil { + t.Errorf("expected error marshalling string at top level") + } + if _, err := Marshal(nil); err == nil { + t.Errorf("expected error marshalling nil") + } +} + +func TestMarshalMapKeyMustBeString(t *testing.T) { + m := map[int]any{1: "x"} + if _, err := Marshal(m); err == nil { + t.Errorf("expected error for non-string map key") + } +} + +func TestMarshalUnexportedFieldSkipped(t *testing.T) { + type Cfg struct { + Pub string `toml:"pub"` + priv string + } + out, err := Marshal(Cfg{Pub: "p", priv: "s"}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "pub = \"p\"\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalBareAndQuotedKeys(t *testing.T) { + type Cfg struct { + Bare string `toml:"bare_key"` + Dash string `toml:"with-dash"` + Num string `toml:"num123"` + Q string `toml:"needs space"` + Dot string `toml:"needs.dot"` + } + out, err := Marshal(Cfg{ + Bare: "a", Dash: "b", Num: "c", Q: "d", Dot: "e", + }) + if err != nil { + t.Fatalf("marshal: %v", err) + } + want := "bare_key = \"a\"\nwith-dash = \"b\"\nnum123 = \"c\"\n\"needs space\" = \"d\"\n\"needs.dot\" = \"e\"\n" + if string(out) != want { + t.Errorf("output mismatch:\ngot: %q\nwant: %q", out, want) + } +} + +func TestMarshalThenParseRoundTrip(t *testing.T) { + type Server struct { + Host string `toml:"host"` + Port int `toml:"port"` + Enabled bool `toml:"enabled"` + Tags []string `toml:"tags"` + } + type Form struct { + Name string `toml:"name"` + Allowed []string `toml:"allowed"` + } + type Cfg struct { + Title string `toml:"title"` + Count int `toml:"count"` + Ratio float64 `toml:"ratio"` + Server Server `toml:"server"` + Forms []Form `toml:"forms"` + Due time.Time `toml:"due"` + Day LocalDate `toml:"day"` + } + in := Cfg{ + Title: "demo", + Count: 42, + Ratio: 3.14, + Server: Server{ + Host: "127.0.0.1", Port: 9090, Enabled: true, + Tags: []string{"a", "b"}, + }, + Forms: []Form{ + {Name: "contact", Allowed: []string{"x"}}, + {Name: "feedback", Allowed: nil}, + }, + Due: time.Date(2026, 6, 26, 10, 0, 0, 0, time.UTC), + Day: LocalDate{Time: time.Date(2026, 6, 26, 0, 0, 0, 0, time.UTC)}, + } + out, err := Marshal(in) + if err != nil { + t.Fatalf("marshal: %v", err) + } + tree1, err := Parse(out) + if err != nil { + t.Fatalf("parse of marshalled: %v\noutput:\n%s", err, out) + } + // Decode back into the struct. + var out2 Cfg + if err := Unmarshal(out, &out2); err != nil { + t.Fatalf("unmarshal of marshalled: %v", err) + } + if !reflect.DeepEqual(in, out2) { + t.Errorf("round-trip mismatch:\nin: %#v\nout: %#v", in, out2) + } + _ = tree1 +} + +func TestMarshalRoundTripFromUntypedTree(t *testing.T) { + src := []byte(`title = "demo" +count = 42 +ratio = 3.14 +enabled = true + +[server] +host = "127.0.0.1" +port = 9090 + +[[items]] +name = "a" +qty = 1 + +[[items]] +name = "b" +qty = 2 + +[meta] +created = 2026-06-26T10:00:00Z +`) + tree1, err := Parse(src) + if err != nil { + t.Fatalf("parse src: %v", err) + } + out, err := Marshal(tree1) + if err != nil { + t.Fatalf("marshal: %v", err) + } + tree2, err := Parse(out) + if err != nil { + t.Fatalf("re-parse marshalled: %v\noutput:\n%s", err, out) + } + if !reflect.DeepEqual(tree1, tree2) { + t.Errorf("round-trip mismatch:\nbefore: %#v\nafter: %#v", tree1, tree2) + } +} + +func TestMarshalUintOverflow(t *testing.T) { + type Cfg struct { + Big uint64 `toml:"big"` + } + if _, err := Marshal(Cfg{Big: 1<<63 + 1}); err == nil { + t.Errorf("expected overflow error") + } +} + +func TestMarshalKeyRequiresUTF8(t *testing.T) { + m := map[string]any{"\xff": "x"} + if _, err := Marshal(m); err == nil { + t.Errorf("expected error for invalid UTF-8 key") + } +} + +func TestMarshalStringRequiresUTF8(t *testing.T) { + type Cfg struct { + S string `toml:"s"` + } + if _, err := Marshal(Cfg{S: "abc\xff"}); err == nil { + t.Errorf("expected error for invalid UTF-8 string") + } +} + +func TestLocalDateString(t *testing.T) { + ld := LocalDate{Time: time.Date(1979, 5, 27, 0, 0, 0, 0, time.UTC)} + if got := ld.String(); got != "1979-05-27" { + t.Errorf("LocalDate.String() = %q, want 1979-05-27", got) + } +} + +func TestLocalDateTimeString(t *testing.T) { + ldt := LocalDateTime{Time: time.Date(1979, 5, 27, 7, 32, 0, 0, time.UTC)} + if got := ldt.String(); got != "1979-05-27T07:32:00" { + t.Errorf("LocalDateTime.String() = %q, want 1979-05-27T07:32:00", got) + } + ldt2 := LocalDateTime{Time: time.Date(1979, 5, 27, 7, 32, 0, 5, time.UTC)} + if got := ldt2.String(); got != "1979-05-27T07:32:00.000000005" { + t.Errorf("LocalDateTime.String() = %q, want 1979-05-27T07:32:00.000000005", got) + } + ldt3 := LocalDateTime{Time: time.Date(1979, 5, 27, 7, 32, 0, 500, time.UTC)} + if got := ldt3.String(); got != "1979-05-27T07:32:00.000000500" { + t.Errorf("LocalDateTime.String() = %q, want 1979-05-27T07:32:00.000000500", got) + } +} + +func TestLocalTimeString(t *testing.T) { + lt := LocalTime{Time: time.Date(0, 1, 1, 7, 32, 0, 0, time.UTC)} + if got := lt.String(); got != "07:32:00" { + t.Errorf("LocalTime.String() = %q, want 07:32:00", got) + } +} + +func TestEncoderEquivalenceToMarshal(t *testing.T) { + type Cfg struct { + Title string `toml:"title"` + Count int `toml:"count"` + } + in := Cfg{Title: "x", Count: 7} + a, err := Marshal(in) + if err != nil { + t.Fatalf("marshal: %v", err) + } + b, err := NewEncoder().Marshal(in) + if err != nil { + t.Fatalf("encoder marshal: %v", err) + } + if !reflect.DeepEqual(a, b) { + t.Errorf("Marshal and Encoder disagree:\n%s\n%s", a, b) + } +} + +type Custom struct { + tag string +} + +func (c Custom) MarshalTOML() (any, error) { return c.tag, nil } diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..dadfc31 --- /dev/null +++ b/go.mod @@ -0,0 +1,3 @@ +module sourcedock.dev/petrbalvin/interpres + +go 1.27.0 diff --git a/interpres.go b/interpres.go new file mode 100644 index 0000000..3b44376 --- /dev/null +++ b/interpres.go @@ -0,0 +1,252 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +// Package interpres is a dependency-free TOML parser for Go. +// +// interpres reads and writes TOML documents using only the standard library. +// It exposes a small, encoding/json-style API: +// +// var cfg Config +// err := interpres.Unmarshal(data, &cfg) +// +// out, err := interpres.Marshal(cfg) +// +// or, for an untyped tree: +// +// tree, err := interpres.Parse(data) +// +// A Decoder allows strict decoding that rejects keys without a matching +// struct field, mirroring (*json.Decoder).DisallowUnknownFields. +package interpres + +import ( + "context" + "fmt" + "unicode/utf8" +) + +// A SyntaxError describes a malformed TOML document, including the 1-based +// line on which the problem was detected. +type SyntaxError struct { + Line int + Msg string +} + +func (e *SyntaxError) Error() string { + return fmt.Sprintf("interpres: line %d: %s", e.Line, e.Msg) +} + +// Parse decodes a TOML document into a nested map[string]any. +// +// Values are mapped to Go types as follows: strings to string, integers to +// int64, floats to float64, booleans to bool, date-times to time.Time, arrays +// to []any, and tables (including inline tables) to map[string]any. +// +// Parse is equivalent to ParseContext with context.Background. +func Parse(data []byte) (map[string]any, error) { + return ParseContext(context.Background(), data) +} + +// ParseContext decodes a TOML document into a nested map[string]any, obeying +// ctx. The context is checked between top-level statements so cancellation is +// honoured before the parser has done substantial work. +func ParseContext(ctx context.Context, data []byte) (map[string]any, error) { + if err := ctx.Err(); err != nil { + return nil, err + } + if !utf8.Valid(data) { + return nil, &SyntaxError{Line: 1, Msg: "input is not valid UTF-8"} + } + p := &parser{src: []rune(string(data)), line: 1, ctx: ctx} + return p.parse() +} + +// Unmarshal parses a TOML document and stores the result in the value pointed +// to by v. v is typically a pointer to a struct or to a map[string]any. +// +// Struct fields are matched to TOML keys by the `toml:"name"` tag, or by a +// case-insensitive match on the field name when no tag is present. A tag of +// "-" skips the field. +// +// Unmarshal is equivalent to UnmarshalContext with context.Background. +func Unmarshal(data []byte, v any) error { + return UnmarshalContext(context.Background(), data, v) +} + +// UnmarshalContext is the cancellable variant of Unmarshal. +func UnmarshalContext(ctx context.Context, data []byte, v any) error { + tree, err := ParseContext(ctx, data) + if err != nil { + return err + } + return newDecoder().decode(tree, v) +} + +// A Decoder decodes a TOML document into a Go value with configurable +// strictness. +type Decoder struct { + disallowUnknown bool +} + +// NewDecoder returns a Decoder. +func NewDecoder() *Decoder { return &Decoder{} } + +// DisallowUnknownFields causes Decode to return an error when the document +// contains a key with no matching destination struct field. +func (d *Decoder) DisallowUnknownFields() *Decoder { + d.disallowUnknown = true + return d +} + +// Decode parses data and stores the result in the value pointed to by v, +// honouring the decoder's strictness settings. +// +// Decode is equivalent to DecodeContext with context.Background. +func (d *Decoder) Decode(data []byte, v any) error { + return d.DecodeContext(context.Background(), data, v) +} + +// DecodeContext is the cancellable variant of Decode. +func (d *Decoder) DecodeContext(ctx context.Context, data []byte, v any) error { + tree, err := ParseContext(ctx, data) + if err != nil { + return err + } + dec := newDecoder() + dec.disallowUnknown = d.disallowUnknown + return dec.decode(tree, v) +} + +// Marshaler is the interface implemented by types that can produce a custom +// TOML representation of themselves. MarshalTOML returns a value that Marshal +// then encodes as if the returned value had been passed in its place, which +// is useful for emitting a Go type as a different TOML shape (for example, a +// struct as an inline table or a primitive alias as a richer value). +type Marshaler interface { + MarshalTOML() (any, error) +} + +// Unmarshaler is the inverse of Marshaler: a type that wants control over +// how it is decoded from a TOML value may implement UnmarshalTOML. The data +// argument is whatever the parser produced for that key: one of string, +// bool, int64, float64, time.Time, LocalDateTime, LocalDate, LocalTime, +// []any, or map[string]any. UnmarshalTOML may parse, inspect, or transform +// the value however it likes, then store the result by mutating its +// receiver through the standard pointer-indirection rules of the reflect +// package (i.e. via reflect.Value.Set or by reassigning fields through a +// pointer the receiver holds). +// +// UnmarshalTOML is invoked from (*Decoder).Decode / Unmarshal when the +// destination type implements the interface. The decoder does not need to +// consult the concrete return value; whatever the receiver stores is kept. +type Unmarshaler interface { + UnmarshalTOML(data any) error +} + +// Marshal returns the TOML 1.0 encoding of v. +// +// Marshal traverses v using reflection and applies the following rules: +// +// - The top-level value must be a struct or a map[string]V. Pointers are +// followed; a nil top-level pointer is an error. +// - Struct fields are matched by `toml:"name"` tag (case-insensitive +// fallback to field name; `-` skips). Anonymous (embedded) fields without +// a tag are inlined. +// - Maps use sorted keys for deterministic output. +// - Slices and arrays of structs or maps become TOML arrays of tables; a +// nil or empty array of tables is omitted (TOML forbids an empty `[[a]]`), +// while other empty arrays emit as `key = []`. +// - Other slices and arrays become TOML arrays. +// - Scalars encode as TOML scalars: bool, int64, float64, string, time.Time +// (offset date-time), and LocalDateTime/LocalDate/LocalTime (local +// variants). +// - Values implementing Marshaler are encoded by calling MarshalTOML and +// using its result. +// - nil pointer fields are omitted. +// +// Marshal cannot encode cyclic data structures; passing one will loop until +// the stack overflows. The output is not guaranteed to be byte-identical to +// the input that produced v: comments, whitespace, key order (for maps), +// string quoting style, and the choice between `[table]` headers and inline +// tables are not preserved. +// +// Marshal is equivalent to MarshalContext with context.Background. +func Marshal(v any) ([]byte, error) { + return MarshalContext(context.Background(), v) +} + +// MarshalContext is the cancellable variant of Marshal. +func MarshalContext(ctx context.Context, v any) ([]byte, error) { + if err := ctx.Err(); err != nil { + return nil, err + } + return NewEncoder().MarshalContext(ctx, v) +} + +// An Encoder encodes Go values into TOML. +// +// All options default to behaviour that preserves byte-for-byte compatibility +// with previous releases and passes the toml-test compliance suite: +// +// GroupByKind: true (scalars first, then tables, then arrays of tables) +// OmitEmptyArrays: false (a nil/empty []string slice emits [] as a value; +// a nil/empty []Item struct slice is still skipped) +// LiteralMultilineAt: 0 (always emit basic multi-line strings with +// escape sequences, never literal ones) +// +// Use the chainable option methods to opt out. The option state is private; +// callers that need the underlying knobs reach for the methods rather than +// reading or mutating fields. +type Encoder struct { + groupByKind bool // default true; set via (*Encoder).GroupByKind + omitEmptyArrays bool // default false; set via (*Encoder).OmitEmptyArrays + literalMultilineAt int // default 0; set via (*Encoder).UseLiteralMultiline +} + +// NewEncoder returns an Encoder with default options. +func NewEncoder() *Encoder { return &Encoder{groupByKind: true} } + +// GroupByKind toggles whether fields at the same TOML level are reordered +// into the group-by-kind layout (scalars first, then tables, then arrays of +// tables). When set to false, the emitter preserves the source declaration +// order (struct field order, or sorted key order for maps). +func (e *Encoder) GroupByKind(v bool) *Encoder { + e.groupByKind = v + return e +} + +// OmitEmptyArrays opts in to skipping empty (non-nil, length 0) TOML arrays +// of scalars. The default emits them as "key = []". Nil slices and empty +// arrays of tables are already always omitted. +func (e *Encoder) OmitEmptyArrays() *Encoder { + e.omitEmptyArrays = true + return e +} + +// UseLiteralMultiline sets the length threshold at which a multi-line string +// is emitted as a literal triple-quoted string instead of the escaped form. +// Use 0 or any negative value to disable (always escaped). The literal form +// is selected only when the value contains an internal newline; otherwise the +// single-line basic form is used regardless of this setting. +func (e *Encoder) UseLiteralMultiline(threshold int) *Encoder { + e.literalMultilineAt = threshold + return e +} + +// Marshal encodes v to TOML bytes. It is equivalent to calling Marshal with v. +// +// Marshal is equivalent to MarshalContext with context.Background. +func (e *Encoder) Marshal(v any) ([]byte, error) { + return e.MarshalContext(context.Background(), v) +} + +// MarshalContext is the cancellable variant of Marshal. +func (e *Encoder) MarshalContext(ctx context.Context, v any) ([]byte, error) { + enc := newEncoder() + enc.ctx = ctx + enc.opts = *e + if err := enc.encode(v); err != nil { + return nil, err + } + return enc.bytes(), nil +} diff --git a/interpres_test.go b/interpres_test.go new file mode 100644 index 0000000..c2c5c63 --- /dev/null +++ b/interpres_test.go @@ -0,0 +1,513 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package interpres + +import ( + "math" + "testing" + "time" +) + +func TestParseScalars(t *testing.T) { + tree, err := Parse([]byte(` +title = "interpres" +count = 42 +ratio = 3.14 +enabled = true +disabled = false +hexv = 0xFF +octv = 0o755 +binv = 0b1010 +grouped = 1_000_000 +neg = -17 +expv = 1e3 +`)) + if err != nil { + t.Fatalf("parse: %v", err) + } + + cases := map[string]any{ + "title": "interpres", + "count": int64(42), + "ratio": 3.14, + "enabled": true, + "disabled": false, + "hexv": int64(255), + "octv": int64(493), + "binv": int64(10), + "grouped": int64(1000000), + "neg": int64(-17), + "expv": 1000.0, + } + for k, want := range cases { + if got := tree[k]; got != want { + t.Errorf("%s = %#v (%T), want %#v (%T)", k, got, got, want, want) + } + } +} + +func TestParseInfNan(t *testing.T) { + tree, err := Parse([]byte("pos = inf\nneg = -inf\nbad = nan\n")) + if err != nil { + t.Fatalf("parse: %v", err) + } + if v := tree["pos"].(float64); !math.IsInf(v, 1) { + t.Errorf("pos = %v, want +Inf", v) + } + if v := tree["neg"].(float64); !math.IsInf(v, -1) { + t.Errorf("neg = %v, want -Inf", v) + } + if v := tree["bad"].(float64); !math.IsNaN(v) { + t.Errorf("bad = %v, want NaN", v) + } +} + +func TestParseStrings(t *testing.T) { + tree, err := Parse([]byte(` +basic = "a\tb\nc" +literal = 'C:\path\no\escape' +quote = "say \"hi\"" +unicode = "\u00e9" +`)) + if err != nil { + t.Fatalf("parse: %v", err) + } + if tree["basic"] != "a\tb\nc" { + t.Errorf("basic = %q", tree["basic"]) + } + if tree["literal"] != `C:\path\no\escape` { + t.Errorf("literal = %q", tree["literal"]) + } + if tree["quote"] != `say "hi"` { + t.Errorf("quote = %q", tree["quote"]) + } + if tree["unicode"] != "é" { + t.Errorf("unicode = %q", tree["unicode"]) + } +} + +func TestParseMultilineString(t *testing.T) { + tree, err := Parse([]byte("text = \"\"\"\nfirst\nsecond\"\"\"\n")) + if err != nil { + t.Fatalf("parse: %v", err) + } + if tree["text"] != "first\nsecond" { + t.Errorf("text = %q, want %q", tree["text"], "first\nsecond") + } +} + +func TestParseMultilineLineEndingBackslash(t *testing.T) { + tree, err := Parse([]byte("text = \"\"\"\\\n one \\\n two\"\"\"\n")) + if err != nil { + t.Fatalf("parse: %v", err) + } + if tree["text"] != "one two" { + t.Errorf("text = %q, want %q", tree["text"], "one two") + } +} + +func TestParseTablesAndDottedKeys(t *testing.T) { + tree, err := Parse([]byte(` +owner.name = "Petr" + +[server] +host = "localhost" +port = 9090 + +[server.tls] +enabled = true +`)) + if err != nil { + t.Fatalf("parse: %v", err) + } + server := tree["server"].(map[string]any) + if server["host"] != "localhost" || server["port"] != int64(9090) { + t.Errorf("server = %#v", server) + } + tls := server["tls"].(map[string]any) + if tls["enabled"] != true { + t.Errorf("tls = %#v", tls) + } + owner := tree["owner"].(map[string]any) + if owner["name"] != "Petr" { + t.Errorf("owner = %#v", owner) + } +} + +func TestParseArrayOfTables(t *testing.T) { + tree, err := Parse([]byte(` +[[forms]] +name = "contact" + +[[forms]] +name = "feedback" +`)) + if err != nil { + t.Fatalf("parse: %v", err) + } + forms := tree["forms"].([]map[string]any) + if len(forms) != 2 { + t.Fatalf("len(forms) = %d, want 2", len(forms)) + } + if forms[0]["name"] != "contact" || forms[1]["name"] != "feedback" { + t.Errorf("forms = %#v", forms) + } +} + +func TestParseArraysAndInlineTables(t *testing.T) { + tree, err := Parse([]byte(` +ports = [80, 443] +mixed = [ + "a", + "b", +] +point = { x = 1, y = 2 } +`)) + if err != nil { + t.Fatalf("parse: %v", err) + } + ports := tree["ports"].([]any) + if len(ports) != 2 || ports[0] != int64(80) || ports[1] != int64(443) { + t.Errorf("ports = %#v", ports) + } + mixed := tree["mixed"].([]any) + if len(mixed) != 2 || mixed[0] != "a" || mixed[1] != "b" { + t.Errorf("mixed = %#v", mixed) + } + point := tree["point"].(map[string]any) + if point["x"] != int64(1) || point["y"] != int64(2) { + t.Errorf("point = %#v", point) + } +} + +func TestParseDateTime(t *testing.T) { + tree, err := Parse([]byte(` +offset = 1979-05-27T07:32:00Z +local = 1979-05-27T07:32:00 +day = 1979-05-27 +clock = 07:32:00 +`)) + if err != nil { + t.Fatalf("parse: %v", err) + } + if off, ok := tree["offset"].(time.Time); !ok || off.Year() != 1979 || off.Hour() != 7 { + t.Errorf("offset = %#v (%T)", tree["offset"], tree["offset"]) + } + if ldt, ok := tree["local"].(LocalDateTime); !ok || ldt.Year() != 1979 || ldt.Hour() != 7 { + t.Errorf("local = %#v (%T)", tree["local"], tree["local"]) + } + if d, ok := tree["day"].(LocalDate); !ok || d.Month() != time.May || d.Day() != 27 { + t.Errorf("day = %#v (%T)", tree["day"], tree["day"]) + } + if clk, ok := tree["clock"].(LocalTime); !ok || clk.Hour() != 7 || clk.Minute() != 32 { + t.Errorf("clock = %#v (%T)", tree["clock"], tree["clock"]) + } +} + +func TestDateTimeFormats(t *testing.T) { + tree, err := Parse([]byte("a = 1987-07-05 17:45:00Z\nb = 1987-07-05t17:45:00z\nc = 1977-12-21T10:32:00.555\n")) + if err != nil { + t.Fatalf("parse: %v", err) + } + if _, ok := tree["a"].(time.Time); !ok { + t.Errorf("a is %T, want time.Time", tree["a"]) + } + if _, ok := tree["b"].(time.Time); !ok { + t.Errorf("b is %T, want time.Time", tree["b"]) + } + if _, ok := tree["c"].(LocalDateTime); !ok { + t.Errorf("c is %T, want LocalDateTime", tree["c"]) + } +} + +func TestUnmarshalDateTime(t *testing.T) { + type Doc struct { + Created time.Time `toml:"created"` + Day LocalDate `toml:"day"` + } + var d Doc + if err := Unmarshal([]byte("created = 2026-06-20T10:00:00Z\nday = 2026-06-20\n"), &d); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if d.Created.Year() != 2026 || d.Created.Hour() != 10 { + t.Errorf("Created = %v", d.Created) + } + if d.Day.Year() != 2026 || d.Day.Month() != time.June || d.Day.Day() != 20 { + t.Errorf("Day = %v", d.Day) + } +} + +func TestUnmarshalStruct(t *testing.T) { + type SMTP struct { + Host string `toml:"host"` + Port int `toml:"port"` + } + type Form struct { + Name string `toml:"name"` + SMTP SMTP `toml:"smtp"` + Origins []string `toml:"allowed_origins"` + } + type Config struct { + Port int `toml:"port"` + Forms []Form `toml:"forms"` + } + + data := []byte(` +port = 8080 + +[[forms]] +name = "contact" +allowed_origins = ["https://example.com"] + +[forms.smtp] +host = "smtp.example.com" +port = 587 +`) + + var cfg Config + if err := Unmarshal(data, &cfg); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if cfg.Port != 8080 { + t.Errorf("Port = %d", cfg.Port) + } + if len(cfg.Forms) != 1 { + t.Fatalf("len(Forms) = %d", len(cfg.Forms)) + } + f := cfg.Forms[0] + if f.Name != "contact" || f.SMTP.Host != "smtp.example.com" || f.SMTP.Port != 587 { + t.Errorf("form = %#v", f) + } + if len(f.Origins) != 1 || f.Origins[0] != "https://example.com" { + t.Errorf("origins = %#v", f.Origins) + } +} + +func TestUnmarshalCaseInsensitiveAndUntagged(t *testing.T) { + type Config struct { + Title string + Count int + } + var cfg Config + if err := Unmarshal([]byte("title = \"x\"\ncount = 3\n"), &cfg); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if cfg.Title != "x" || cfg.Count != 3 { + t.Errorf("cfg = %#v", cfg) + } +} + +func TestDisallowUnknownFields(t *testing.T) { + type C struct { + Known string `toml:"known"` + } + data := []byte("known = \"x\"\nbogus = 1\n") + + var lenient C + if err := Unmarshal(data, &lenient); err != nil { + t.Fatalf("lenient unmarshal: %v", err) + } + + var strict C + err := NewDecoder().DisallowUnknownFields().Decode(data, &strict) + if err == nil { + t.Fatal("expected error for unknown field, got nil") + } +} + +func TestSkippedFieldTag(t *testing.T) { + type C struct { + Keep string `toml:"keep"` + Skip string `toml:"-"` + } + var c C + if err := Unmarshal([]byte("keep = \"y\"\n"), &c); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if c.Keep != "y" || c.Skip != "" { + t.Errorf("c = %#v", c) + } +} + +func TestSyntaxErrorReportsLine(t *testing.T) { + _, err := Parse([]byte("a = 1\nb = \nc = 3\n")) + if err == nil { + t.Fatal("expected a syntax error") + } + se, ok := err.(*SyntaxError) + if !ok { + t.Fatalf("error is %T, want *SyntaxError", err) + } + if se.Line != 2 { + t.Errorf("Line = %d, want 2", se.Line) + } +} + +func TestComments(t *testing.T) { + tree, err := Parse([]byte(` +# a leading comment +key = "value" # trailing comment +# another +`)) + if err != nil { + t.Fatalf("parse: %v", err) + } + if tree["key"] != "value" { + t.Errorf("key = %q", tree["key"]) + } +} + +func TestDuplicateKeyRejected(t *testing.T) { + _, err := Parse([]byte("a = 1\na = 2\n")) + if err == nil { + t.Fatal("expected duplicate key error") + } +} + +func TestRejectsInvalidNumbers(t *testing.T) { + for _, tok := range []string{ + "01", "-01", "00", + "1__0", "_1", "1_", "0x_1", "1_.0", + "1.", ".5", "1.2.3", "1.e2", + "0x", "0o", "0b", "0b2", "0o8", "0xG", + "+0x1", + } { + if _, err := Parse([]byte("v = " + tok + "\n")); err == nil { + t.Errorf("%q: expected an error, got none", tok) + } + } +} + +func TestAcceptsNumberEdgeCases(t *testing.T) { + cases := map[string]any{ + "0": int64(0), + "-0": int64(0), + "+99": int64(99), + "1_000": int64(1000), + "0xDEAD_BEEF": int64(0xDEADBEEF), + "0o755": int64(493), + "0b1010": int64(10), + "0.0": 0.0, + "3.14": 3.14, + "6.022e23": 6.022e23, + "1e10": 1e10, + "-2.5E-3": -2.5e-3, + } + for tok, want := range cases { + tree, err := Parse([]byte("v = " + tok + "\n")) + if err != nil { + t.Errorf("%q: %v", tok, err) + continue + } + if got := tree["v"]; got != want { + t.Errorf("%q = %#v (%T), want %#v", tok, got, got, want) + } + } +} + +func TestRejectsTableRedefinition(t *testing.T) { + _, err := Parse([]byte("[a]\nx = 1\n\n[a]\ny = 2\n")) + if err == nil { + t.Fatal("expected a table-redefinition error") + } +} + +func TestAllowsImplicitThenExplicitTable(t *testing.T) { + tree, err := Parse([]byte("[a.b]\nx = 1\n\n[a]\ny = 2\n")) + if err != nil { + t.Fatalf("parse: %v", err) + } + a := tree["a"].(map[string]any) + if a["y"] != int64(2) { + t.Errorf("a.y = %#v", a["y"]) + } + if b := a["b"].(map[string]any); b["x"] != int64(1) { + t.Errorf("a.b.x = %#v", b["x"]) + } +} + +func TestRejectsControlCharInString(t *testing.T) { + if _, err := Parse([]byte("v = \"a\x01b\"\n")); err == nil { + t.Fatal("expected a control-character error") + } +} + +func TestAllowsEscapedControlChar(t *testing.T) { + tree, err := Parse([]byte(`v = "\u0000"`)) + if err != nil { + t.Fatalf("parse: %v", err) + } + if tree["v"] != "\x00" { + t.Errorf("v = %q", tree["v"]) + } +} + +func TestMultilineQuotesAtDelimiter(t *testing.T) { + tree, err := Parse([]byte("a = '''''two quotes'''''\n")) + if err != nil { + t.Fatalf("parse: %v", err) + } + if tree["a"] != "''two quotes''" { + t.Errorf("a = %q", tree["a"]) + } +} + +func TestRejectsInlineTableExtension(t *testing.T) { + cases := map[string]string{ + "by header": "a = { b = 1 }\n[a.c]\nx = 2\n", + "by dotted key": "a = { b = 1 }\na.c = 2\n", + "header over it": "a = { b = 1 }\n[a]\nx = 2\n", + } + for name, doc := range cases { + if _, err := Parse([]byte(doc)); err == nil { + t.Errorf("%s: expected an inline-table extension error", name) + } + } +} + +func TestRejectsSpecInvalid(t *testing.T) { + cases := map[string]string{ + "single-digit hour": "a = 2023-10-01T1:32:00Z\n", + "inline duplicate key": "a = { b = 1, b = 2 }\n", + "inline dotted overwrite": "a = { b = 1, b.c = 2 }\n", + "dotted over header": "[a.b]\nx = 1\n[a]\nb.y = 2\n", + "table over array": "[[t]]\n[t]\n", + "truncated datetime": "a = 2026-01-02T\n", + "datetime no seconds": "a = 2026-01-02T07:32\n", + } + for name, doc := range cases { + if _, err := Parse([]byte(doc)); err == nil { + t.Errorf("%s: expected an error", name) + } + } +} + +func TestArrayOfTablesPerElementSubtable(t *testing.T) { + tree, err := Parse([]byte(` +[[forms]] +name = "a" + +[forms.smtp] +host = "h1" + +[[forms]] +name = "b" + +[forms.smtp] +host = "h2" +`)) + if err != nil { + t.Fatalf("parse: %v", err) + } + forms := tree["forms"].([]map[string]any) + if len(forms) != 2 { + t.Fatalf("len(forms) = %d", len(forms)) + } + if h := forms[0]["smtp"].(map[string]any)["host"]; h != "h1" { + t.Errorf("forms[0].smtp.host = %v", h) + } + if h := forms[1]["smtp"].(map[string]any)["host"]; h != "h2" { + t.Errorf("forms[1].smtp.host = %v", h) + } +} diff --git a/number.go b/number.go new file mode 100644 index 0000000..f3d350b --- /dev/null +++ b/number.go @@ -0,0 +1,167 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package interpres + +import ( + "fmt" + "math" + "strconv" + "strings" +) + +// decodeNumber parses a bare numeric token under strict TOML rules: no leading +// zeros, underscores only between digits, prefixed radixes without a sign, and +// floats with explicit fraction/exponent digits. +func decodeNumber(tok string) (any, error) { + switch tok { + case "inf", "+inf": + return math.Inf(1), nil + case "-inf": + return math.Inf(-1), nil + case "nan", "+nan", "-nan": + return math.NaN(), nil + } + + if len(tok) >= 2 && tok[0] == '0' && (tok[1] == 'x' || tok[1] == 'o' || tok[1] == 'b') { + return decodeRadix(tok) + } + if strings.ContainsAny(tok, ".eE") { + return decodeFloat(tok) + } + return decodeDecimalInt(tok) +} + +func decodeDecimalInt(tok string) (any, error) { + sign, body := splitSign(tok) + digits, err := joinDigits(body, isDecDigit) + if err != nil { + return nil, err + } + if err := checkNoLeadingZero(digits); err != nil { + return nil, err + } + i, err := strconv.ParseInt(sign+digits, 10, 64) + if err != nil { + return nil, fmt.Errorf("integer %q out of range", tok) + } + return i, nil +} + +func decodeRadix(tok string) (any, error) { + var base int + var isDigit func(byte) bool + switch tok[1] { + case 'x': + base, isDigit = 16, isHexDigit + case 'o': + base, isDigit = 8, isOctDigit + case 'b': + base, isDigit = 2, isBinDigit + } + digits, err := joinDigits(tok[2:], isDigit) + if err != nil { + return nil, err + } + i, err := strconv.ParseInt(digits, base, 64) + if err != nil { + return nil, fmt.Errorf("integer %q out of range", tok) + } + return i, nil +} + +func decodeFloat(tok string) (any, error) { + sign, s := splitSign(tok) + + mantissa, exp := s, "" + if i := strings.IndexAny(s, "eE"); i >= 0 { + mantissa, exp = s[:i], s[i+1:] + } + + intPart, frac, hasDot := mantissa, "", false + if i := strings.IndexByte(mantissa, '.'); i >= 0 { + intPart, frac, hasDot = mantissa[:i], mantissa[i+1:], true + } + if !hasDot && exp == "" { + return nil, fmt.Errorf("invalid float %q", tok) + } + + ip, err := joinDigits(intPart, isDecDigit) + if err != nil { + return nil, err + } + if err := checkNoLeadingZero(ip); err != nil { + return nil, err + } + build := sign + ip + + if hasDot { + fp, err := joinDigits(frac, isDecDigit) + if err != nil { + return nil, err + } + build += "." + fp + } + if exp != "" { + esign, edigits := splitSign(exp) + ed, err := joinDigits(edigits, isDecDigit) + if err != nil { + return nil, err + } + build += "e" + esign + ed + } + + f, err := strconv.ParseFloat(build, 64) + if err != nil { + return nil, fmt.Errorf("invalid float %q", tok) + } + return f, nil +} + +// joinDigits validates that every rune is a digit (per isDigit) and that each +// underscore sits between two digits, returning the digits with underscores +// removed. +func joinDigits(s string, isDigit func(byte) bool) (string, error) { + if s == "" { + return "", fmt.Errorf("number is missing digits") + } + var b strings.Builder + for i := range len(s) { + c := s[i] + if c == '_' { + if i == 0 || i == len(s)-1 || !isDigit(s[i-1]) || !isDigit(s[i+1]) { + return "", fmt.Errorf("misplaced underscore in number %q", s) + } + continue + } + if !isDigit(c) { + return "", fmt.Errorf("invalid character %q in number", string(c)) + } + b.WriteByte(c) + } + return b.String(), nil +} + +func checkNoLeadingZero(digits string) error { + if len(digits) > 1 && digits[0] == '0' { + return fmt.Errorf("leading zeros are not allowed in numbers") + } + return nil +} + +func splitSign(tok string) (sign, rest string) { + if tok != "" && (tok[0] == '+' || tok[0] == '-') { + if tok[0] == '-' { + return "-", tok[1:] + } + return "", tok[1:] + } + return "", tok +} + +func isDecDigit(c byte) bool { return c >= '0' && c <= '9' } +func isOctDigit(c byte) bool { return c >= '0' && c <= '7' } +func isBinDigit(c byte) bool { return c == '0' || c == '1' } +func isHexDigit(c byte) bool { + return isDecDigit(c) || (c >= 'a' && c <= 'f') || (c >= 'A' && c <= 'F') +} diff --git a/parser.go b/parser.go new file mode 100644 index 0000000..cc40128 --- /dev/null +++ b/parser.go @@ -0,0 +1,929 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package interpres + +import ( + "context" + "fmt" + "strconv" + "strings" +) + +// ctxCheckInterval is the number of top-level parser iterations between +// context-cancellation checks. A small interval keeps the response snappy on +// cancellation; a too-small one wastes cycles on a non-cancelled run. +const ctxCheckInterval = 64 + +// parser is a recursive-descent TOML parser producing a map[string]any tree. +type parser struct { + src []rune + pos int + line int + ctx context.Context + + root map[string]any + current map[string]any + headers map[string]bool + frozen map[string]bool + dotted map[string]bool + arrays map[string]bool + + currentPath []string +} + +func (p *parser) parse() (map[string]any, error) { + p.root = map[string]any{} + p.current = p.root + p.headers = map[string]bool{} + p.frozen = map[string]bool{} + p.dotted = map[string]bool{} + p.arrays = map[string]bool{} + p.currentPath = nil + + for i := 0; ; i++ { + if i%ctxCheckInterval == 0 { + if err := p.checkCtx(); err != nil { + return nil, err + } + } + if err := p.skipBlank(); err != nil { + return nil, err + } + if p.eof() { + break + } + c := p.peek() + switch { + case c == '[': + if err := p.parseTableHeader(); err != nil { + return nil, err + } + default: + if err := p.parseKeyValue(); err != nil { + return nil, err + } + } + if err := p.expectLineEnd(); err != nil { + return nil, err + } + } + return p.root, nil +} + +// checkCtx returns ctx.Err() when the context has been cancelled, nil +// otherwise. The call is a no-op when ctx is nil or the zero Background +// context, both of which never cancel. +func (p *parser) checkCtx() error { + if p.ctx == nil { + return nil + } + return p.ctx.Err() +} + +// --- table headers --------------------------------------------------------- + +func (p *parser) parseTableHeader() error { + array := false + p.next() // consume '[' + if !p.eof() && p.peek() == '[' { + array = true + p.next() + } + + key, err := p.parseKeyPath() + if err != nil { + return err + } + + p.skipInline() + if p.eof() || p.peek() != ']' { + return p.errf("expected ']' to close table header") + } + p.next() + if array { + if p.eof() || p.peek() != ']' { + return p.errf("expected ']]' to close array-of-tables header") + } + p.next() + } + + if array { + tbl, err := p.appendArrayTable(key) + if err != nil { + return err + } + // A new array-of-tables element starts a fresh scope: sub-table headers + // and inline-table freezes from the previous element no longer apply. + p.resetScopeUnder(key) + p.arrays[pathKey(key)] = true + p.current = tbl + p.currentPath = key + return nil + } + + pk := pathKey(key) + if p.headers[pk] || p.dotted[pk] || p.arrays[pk] { + return p.errf("table %q is defined more than once", strings.Join(key, ".")) + } + p.headers[pk] = true + + tbl, err := p.tableAt(key) + if err != nil { + return err + } + p.current = tbl + p.currentPath = key + return nil +} + +// tableAt walks (creating intermediate tables) to the table named by key, +// relative to the document root, rejecting any step into a frozen inline table. +func (p *parser) tableAt(key []string) (map[string]any, error) { + cur := p.root + path := make([]string, 0, len(key)) + for _, k := range key { + path = append(path, k) + if p.frozen[pathKey(path)] { + return nil, p.errf("cannot extend inline table %q", strings.Join(path, ".")) + } + existing, ok := cur[k] + if !ok { + next := map[string]any{} + cur[k] = next + cur = next + continue + } + switch v := existing.(type) { + case map[string]any: + cur = v + case []map[string]any: + if len(v) == 0 { + return nil, p.errf("key %q is an empty array of tables", k) + } + cur = v[len(v)-1] + default: + return nil, p.errf("key %q is not a table", k) + } + } + return cur, nil +} + +func (p *parser) appendArrayTable(key []string) (map[string]any, error) { + parent := p.root + for _, k := range key[:len(key)-1] { + existing, ok := parent[k] + if !ok { + next := map[string]any{} + parent[k] = next + parent = next + continue + } + switch v := existing.(type) { + case map[string]any: + parent = v + case []map[string]any: + parent = v[len(v)-1] + default: + return nil, p.errf("key %q is not a table", k) + } + } + + leaf := key[len(key)-1] + tbl := map[string]any{} + switch existing := parent[leaf].(type) { + case nil: + parent[leaf] = []map[string]any{tbl} + case []map[string]any: + parent[leaf] = append(existing, tbl) + default: + return nil, p.errf("key %q is not an array of tables", leaf) + } + return tbl, nil +} + +// --- key/value ------------------------------------------------------------- + +func (p *parser) parseKeyValue() error { + key, err := p.parseKeyPath() + if err != nil { + return err + } + p.skipInline() + if p.eof() || p.peek() != '=' { + return p.errf("expected '=' after key") + } + p.next() + p.skipInline() + + val, err := p.parseValue() + if err != nil { + return err + } + + dest := p.current + abs := append([]string{}, p.currentPath...) + for _, k := range key[:len(key)-1] { + abs = append(abs, k) + if p.frozen[pathKey(abs)] { + return p.errf("cannot extend inline table %q", strings.Join(abs, ".")) + } + if p.headers[pathKey(abs)] { + return p.errf("cannot extend table %q with a dotted key", strings.Join(abs, ".")) + } + p.dotted[pathKey(abs)] = true + existing, ok := dest[k] + if !ok { + next := map[string]any{} + dest[k] = next + dest = next + continue + } + m, ok := existing.(map[string]any) + if !ok { + return p.errf("key %q is not a table", k) + } + dest = m + } + leaf := key[len(key)-1] + abs = append(abs, leaf) + if _, exists := dest[leaf]; exists { + return p.errf("duplicate key %q", leaf) + } + dest[leaf] = val + p.freezeInline(abs, val) + return nil +} + +// freezeInline marks the path of an inline table (and any nested inline tables) +// as immutable, so a later header or dotted key cannot extend it. +func (p *parser) freezeInline(path []string, val any) { + m, ok := val.(map[string]any) + if !ok { + return + } + p.frozen[pathKey(path)] = true + for k, v := range m { + child := append(append([]string{}, path...), k) + p.freezeInline(child, v) + } +} + +// resetScopeUnder forgets the header and freeze records nested under key, which +// belong to the previous element of an array of tables. +func (p *parser) resetScopeUnder(key []string) { + prefix := pathKey(key) + "\x00" + for k := range p.headers { + if strings.HasPrefix(k, prefix) { + delete(p.headers, k) + } + } + for k := range p.frozen { + if strings.HasPrefix(k, prefix) { + delete(p.frozen, k) + } + } +} + +// parseKeyPath parses a dotted key into its components. +func (p *parser) parseKeyPath() ([]string, error) { + var parts []string + for { + p.skipInline() + part, err := p.parseKeyComponent() + if err != nil { + return nil, err + } + parts = append(parts, part) + p.skipInline() + if !p.eof() && p.peek() == '.' { + p.next() + continue + } + break + } + return parts, nil +} + +func (p *parser) parseKeyComponent() (string, error) { + if p.eof() { + return "", p.errf("expected a key") + } + switch c := p.peek(); c { + case '"': + if p.lookahead(`"""`) { + return "", p.errf("multiline strings are not allowed in keys") + } + return p.parseBasicString() + case '\'': + if p.lookahead(`'''`) { + return "", p.errf("multiline strings are not allowed in keys") + } + return p.parseLiteralString() + default: + start := p.pos + for !p.eof() { + c := p.peek() + if (c >= 'A' && c <= 'Z') || (c >= 'a' && c <= 'z') || + (c >= '0' && c <= '9') || c == '_' || c == '-' { + p.next() + continue + } + break + } + if p.pos == start { + return "", p.errf("invalid key character %q", string(p.peek())) + } + return string(p.src[start:p.pos]), nil + } +} + +// --- values ---------------------------------------------------------------- + +func (p *parser) parseValue() (any, error) { + if p.eof() { + return nil, p.errf("expected a value") + } + switch c := p.peek(); { + case c == '"': + return p.parseBasicString() + case c == '\'': + return p.parseLiteralString() + case c == '[': + return p.parseArray() + case c == '{': + return p.parseInlineTable() + case c == 't' || c == 'f': + return p.parseBool() + default: + return p.parseAtom() + } +} + +func (p *parser) parseBool() (any, error) { + if p.match("true") { + return true, nil + } + if p.match("false") { + return false, nil + } + return nil, p.errf("invalid value") +} + +// parseAtom handles numbers, inf/nan, and date-times. +func (p *parser) parseAtom() (any, error) { + start := p.pos + p.scanBareToken() + tok := string(p.src[start:p.pos]) + if tok == "" { + return nil, p.errf("expected a value") + } + // A date may be followed by a space and a time, forming one date-time. + if isDateToken(tok) && !p.eof() && p.peek() == ' ' { + if next, ok := p.peekAt(1); ok && next >= '0' && next <= '9' { + p.next() // consume the separating space + timeStart := p.pos + p.scanBareToken() + tok = tok + " " + string(p.src[timeStart:p.pos]) + } + } + if v, ok := parseDateTime(tok); ok { + return v, nil + } + v, err := decodeNumber(tok) + if err != nil { + return nil, p.errf("%s", err) + } + return v, nil +} + +// scanBareToken advances past a bare value token (number, bool, or date-time), +// stopping at whitespace, a separator, or a comment. +func (p *parser) scanBareToken() { + for !p.eof() { + c := p.peek() + if c == ' ' || c == '\t' || c == '\n' || c == '\r' || + c == ',' || c == ']' || c == '}' || c == '#' { + return + } + p.next() + } +} + +// --- strings --------------------------------------------------------------- + +func (p *parser) parseBasicString() (string, error) { + if p.lookahead(`"""`) { + return p.parseMultilineString('"', true) + } + p.next() // opening quote + var b strings.Builder + for { + if p.eof() { + return "", p.errf("unterminated string") + } + c := p.next() + switch c { + case '"': + return b.String(), nil + case '\n': + return "", p.errf("unterminated string") + case '\r': + return "", p.errf("bare carriage return is not allowed in a string") + case '\\': + r, err := p.readEscape() + if err != nil { + return "", err + } + b.WriteRune(r) + default: + if isControlRune(c) { + return "", p.errf("control character U+%04X is not allowed in a string", c) + } + b.WriteRune(c) + } + } +} + +func (p *parser) parseLiteralString() (string, error) { + if p.lookahead(`'''`) { + return p.parseMultilineString('\'', false) + } + p.next() // opening quote + var b strings.Builder + for { + if p.eof() { + return "", p.errf("unterminated literal string") + } + c := p.next() + if c == '\'' { + return b.String(), nil + } + if c == '\n' { + return "", p.errf("unterminated literal string") + } + if c == '\r' { + return "", p.errf("bare carriage return is not allowed in a string") + } + if isControlRune(c) { + return "", p.errf("control character U+%04X is not allowed in a string", c) + } + b.WriteRune(c) + } +} + +func (p *parser) parseMultilineString(quote rune, escapes bool) (string, error) { + p.skipN(3) // opening delimiter + // A newline immediately after the opening delimiter is trimmed. + if !p.eof() && p.peek() == '\r' { + p.next() + } + if !p.eof() && p.peek() == '\n' { + p.line++ + p.next() + } + + var b strings.Builder + for { + if p.eof() { + return "", p.errf("unterminated multiline string") + } + if p.peek() == quote { + // Count the run of delimiter characters. The last three close the + // string; up to two extra ones belong to the content. + n := 0 + for p.pos+n < len(p.src) && p.src[p.pos+n] == quote { + n++ + } + if n >= 3 { + if n > 5 { + return "", p.errf("too many '%c' before the closing delimiter", quote) + } + for range n - 3 { + b.WriteRune(quote) + } + p.skipN(n) + return b.String(), nil + } + for range n { + b.WriteRune(quote) + p.next() + } + continue + } + c := p.next() + if c == '\n' { + p.line++ + b.WriteRune(c) + continue + } + if c == '\r' { + if !p.eof() && p.peek() == '\n' { + b.WriteRune(c) + continue + } + return "", p.errf("bare carriage return is not allowed in a string") + } + if escapes && c == '\\' { + // Line-ending backslash trims the following whitespace/newlines. + if p.trimLineEndingBackslash() { + continue + } + r, err := p.readEscape() + if err != nil { + return "", err + } + b.WriteRune(r) + continue + } + if isControlRune(c) { + return "", p.errf("control character U+%04X is not allowed in a string", c) + } + b.WriteRune(c) + } +} + +// trimLineEndingBackslash consumes whitespace through the next newline (and the +// blank lines that follow) when a backslash is the last token on a line. +// It reports whether it did so. +func (p *parser) trimLineEndingBackslash() bool { + save, saveLine := p.pos, p.line + for !p.eof() { + c := p.peek() + if c == ' ' || c == '\t' || c == '\r' { + p.next() + continue + } + if c == '\n' { + break + } + // Not a line-ending backslash; restore. + p.pos, p.line = save, saveLine + return false + } + if p.eof() { + p.pos, p.line = save, saveLine + return false + } + // Consume the newline and all following whitespace. + for !p.eof() { + c := p.peek() + if c == '\n' { + p.line++ + p.next() + continue + } + if c == ' ' || c == '\t' || c == '\r' { + p.next() + continue + } + break + } + return true +} + +func (p *parser) readEscape() (rune, error) { + if p.eof() { + return 0, p.errf("unterminated escape sequence") + } + c := p.next() + switch c { + case 'b': + return '\b', nil + case 't': + return '\t', nil + case 'n': + return '\n', nil + case 'f': + return '\f', nil + case 'r': + return '\r', nil + case '"': + return '"', nil + case '\\': + return '\\', nil + case 'u': + return p.readUnicode(4) + case 'U': + return p.readUnicode(8) + default: + return 0, p.errf("invalid escape sequence \\%c", c) + } +} + +func (p *parser) readUnicode(n int) (rune, error) { + if p.pos+n > len(p.src) { + return 0, p.errf("invalid unicode escape") + } + hex := string(p.src[p.pos : p.pos+n]) + p.pos += n + v, err := strconv.ParseInt(hex, 16, 64) + if err != nil { + return 0, p.errf("invalid unicode escape \\%s", hex) + } + if v > 0x10FFFF || (v >= 0xD800 && v <= 0xDFFF) { + return 0, p.errf("escape \\%s is not a valid Unicode scalar value", hex) + } + return rune(v), nil +} + +// --- arrays and inline tables --------------------------------------------- + +func (p *parser) parseArray() (any, error) { + p.next() // '[' + arr := []any{} + for { + if err := p.skipArraySpace(); err != nil { + return nil, err + } + if p.eof() { + return nil, p.errf("unterminated array") + } + if p.peek() == ']' { + p.next() + return arr, nil + } + v, err := p.parseValue() + if err != nil { + return nil, err + } + arr = append(arr, v) + if err := p.skipArraySpace(); err != nil { + return nil, err + } + if p.eof() { + return nil, p.errf("unterminated array") + } + switch p.peek() { + case ',': + p.next() + case ']': + p.next() + return arr, nil + default: + return nil, p.errf("expected ',' or ']' in array") + } + } +} + +func (p *parser) parseInlineTable() (any, error) { + p.next() // '{' + tbl := map[string]any{} + assigned := map[string]bool{} + p.skipInline() + if !p.eof() && p.peek() == '}' { + p.next() + return tbl, nil + } + for { + p.skipInline() + key, err := p.parseKeyPath() + if err != nil { + return nil, err + } + p.skipInline() + if p.eof() || p.peek() != '=' { + return nil, p.errf("expected '=' in inline table") + } + p.next() + p.skipInline() + val, err := p.parseValue() + if err != nil { + return nil, err + } + + dest := tbl + path := make([]string, 0, len(key)) + for _, k := range key[:len(key)-1] { + path = append(path, k) + if assigned[pathKey(path)] { + return nil, p.errf("key %q is already defined", strings.Join(path, ".")) + } + existing, ok := dest[k] + if !ok { + m := map[string]any{} + dest[k] = m + dest = m + continue + } + m, isMap := existing.(map[string]any) + if !isMap { + return nil, p.errf("key %q is already defined", k) + } + dest = m + } + leaf := key[len(key)-1] + path = append(path, leaf) + if _, exists := dest[leaf]; exists { + return nil, p.errf("duplicate key %q in inline table", leaf) + } + dest[leaf] = val + assigned[pathKey(path)] = true + + p.skipInline() + if p.eof() { + return nil, p.errf("unterminated inline table") + } + switch p.peek() { + case ',': + p.next() + case '}': + p.next() + return tbl, nil + default: + return nil, p.errf("expected ',' or '}' in inline table") + } + } +} + +// --- scanning helpers ------------------------------------------------------ + +func (p *parser) eof() bool { return p.pos >= len(p.src) } +func (p *parser) peek() rune { return p.src[p.pos] } + +// peekAt returns the rune at offset n from the current position and whether the +// offset is within the source. Use it instead of indexing p.src directly when +// the offset may sit past the end. +func (p *parser) peekAt(n int) (rune, bool) { + i := p.pos + n + if i < 0 || i >= len(p.src) { + return 0, false + } + return p.src[i], true +} + +func (p *parser) next() rune { + c := p.src[p.pos] + p.pos++ + return c +} + +func (p *parser) skipN(n int) { + for i := 0; i < n && !p.eof(); i++ { + p.next() + } +} + +func (p *parser) match(word string) bool { + if p.lookahead(word) { + p.skipN(len([]rune(word))) + return true + } + return false +} + +func (p *parser) lookahead(s string) bool { + r := []rune(s) + if p.pos+len(r) > len(p.src) { + return false + } + for i, c := range r { + if p.src[p.pos+i] != c { + return false + } + } + return true +} + +// skipInline consumes spaces and tabs only. +func (p *parser) skipInline() { + for !p.eof() { + if c := p.peek(); c == ' ' || c == '\t' { + p.next() + continue + } + break + } +} + +// skipArraySpace consumes whitespace, newlines, and comments inside arrays. +func (p *parser) skipArraySpace() error { + for !p.eof() { + switch p.peek() { + case ' ', '\t': + p.next() + case '\r': + if err := p.expectCRLF(); err != nil { + return err + } + case '\n': + p.line++ + p.next() + case '#': + if err := p.skipComment(); err != nil { + return err + } + default: + return nil + } + } + return nil +} + +// skipBlank consumes whitespace, blank lines, and comments between statements. +func (p *parser) skipBlank() error { + for !p.eof() { + switch p.peek() { + case ' ', '\t': + p.next() + case '\r': + if err := p.expectCRLF(); err != nil { + return err + } + case '\n': + p.line++ + p.next() + case '#': + if err := p.skipComment(); err != nil { + return err + } + default: + return nil + } + } + return nil +} + +func (p *parser) skipComment() error { + p.next() // consume '#' + for !p.eof() { + c := p.peek() + switch { + case c == '\n': + return nil + case c == '\r': + if p.pos+1 < len(p.src) && p.src[p.pos+1] == '\n' { + return nil + } + return p.errf("bare carriage return is not allowed") + case c == '\t': + p.next() + case c < 0x20 || c == 0x7f: + return p.errf("control character U+%04X is not allowed in a comment", c) + default: + p.next() + } + } + return nil +} + +// expectCRLF consumes a carriage return that must be immediately followed by a +// line feed; a bare CR is invalid. +func (p *parser) expectCRLF() error { + if p.pos+1 < len(p.src) && p.src[p.pos+1] == '\n' { + p.next() // consume CR; the LF is handled by the caller + return nil + } + return p.errf("bare carriage return is not allowed") +} + +// expectLineEnd consumes trailing inline whitespace and an optional comment, +// then requires a newline or end of input. +func (p *parser) expectLineEnd() error { + p.skipInline() + if p.eof() { + return nil + } + if p.peek() == '#' { + if err := p.skipComment(); err != nil { + return err + } + } + if p.eof() { + return nil + } + if p.peek() == '\r' { + if err := p.expectCRLF(); err != nil { + return err + } + } + if p.eof() { + return nil + } + if p.peek() == '\n' { + p.line++ + p.next() + return nil + } + return p.errf("unexpected %q after value", string(p.peek())) +} + +func (p *parser) errf(format string, args ...any) error { + return &SyntaxError{Line: p.line, Msg: fmt.Sprintf(format, args...)} +} + +// pathKey joins key components with a NUL separator so a dotted path can be +// used as a map key for tracking defined tables. +func pathKey(parts []string) string { + return strings.Join(parts, "\x00") +} + +// isControlRune reports whether r is a control character disallowed in a string +// literal. Tab, line feed, and carriage return are permitted (handled +// elsewhere); everything else below U+0020, plus U+007F, is rejected. +func isControlRune(r rune) bool { + if r == '\t' || r == '\n' || r == '\r' { + return false + } + return r < 0x20 || r == 0x7f +}