Files
interpres/encode.go
T

861 lines
22 KiB
Go
Raw Normal View History

2026-08-19 09:47:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package interpres
import (
"bytes"
"context"
"errors"
2026-08-19 09:47:00 +02:00
"fmt"
"maps"
2026-08-19 09:47:00 +02:00
"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 fieldOmitted(f, v.Field(i)) {
continue
}
2026-08-19 09:47:00 +02:00
if err := addField(doc, name, v.Field(i), ctx); err != nil {
return err
}
}
return nil
}
// isZeroer mirrors encoding/json's omitzero: a type that knows its own zero
// state decides through that method before reflection is consulted.
type isZeroer interface{ IsZero() bool }
// fieldOmitted reports whether the field's tag options drop it from the
// output: omitzero skips the zero value of the field's type, omitempty skips
// an empty collection (slice, array, or map). The decoder ignores both
// options; they shape emission only.
func fieldOmitted(f reflect.StructField, v reflect.Value) bool {
tag, ok := f.Tag.Lookup("toml")
if !ok {
return false
}
_, opts, _ := strings.Cut(tag, ",")
for opts != "" {
var opt string
opt, opts, _ = strings.Cut(opts, ",")
switch opt {
case "omitzero":
if isZeroValue(v) {
return true
}
case "omitempty":
switch v.Kind() {
case reflect.Slice, reflect.Array, reflect.Map:
if v.Len() == 0 {
return true
}
}
}
}
return false
}
func isZeroValue(v reflect.Value) bool {
if v.CanInterface() {
if z, ok := v.Interface().(isZeroer); ok {
return z.IsZero()
}
}
return v.IsZero()
}
2026-08-19 09:47:00 +02:00
// 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 &EncodeError{Path: joinKey(ctx, name), Err: err}
2026-08-19 09:47:00 +02:00
}
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)
}
// An array keeps the [[header]] form only when every element is a table.
// TOML lets one array mix tables with scalars, and that mix renders as a
// value array with the table elements written inline.
allTables := true
for i := range n {
if !isTableElementValue(v.Index(i)) {
allTables = false
break
}
}
if allTables {
2026-08-19 09:47:00 +02:00
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 &EncodeError{Path: fmt.Sprintf("%s[%d]", joinKey(ctx, name), i), Err: errors.New("nil element")}
2026-08-19 09:47:00 +02:00
}
sub := &tomlDoc{ctx: doc.ctx, opts: doc.opts}
switch ev.Kind() {
case reflect.Struct:
if isScalarStruct(ev.Type()) {
return &EncodeError{Path: fmt.Sprintf("%s[%d]", joinKey(ctx, name), i), Err: errors.New("heterogeneous array contains scalar")}
2026-08-19 09:47:00 +02:00
}
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 &EncodeError{Path: joinKey(ctx, name), Err: errors.New("heterogeneous array, expected table")}
2026-08-19 09:47:00 +02:00
}
subs[i] = sub
}
doc.addArray(name, subs)
return nil
}
// Value array. Table elements normalise to map[string]any and the emitter
// writes them as inline tables.
2026-08-19 09:47:00 +02:00
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 &EncodeError{Path: fmt.Sprintf("%s[%d]", joinKey(ctx, name), i), Err: errors.New("nil element")}
2026-08-19 09:47:00 +02:00
}
if ev.CanInterface() {
if m, ok := ev.Interface().(Marshaler); ok {
mv, err := m.MarshalTOML()
if err != nil {
return &EncodeError{Path: fmt.Sprintf("%s[%d]", joinKey(ctx, name), i), Err: err}
2026-08-19 09:47:00 +02:00
}
ev = reflect.ValueOf(mv)
ev = followPtr(ev)
}
}
val, err := normaliseValue(ev)
if err != nil {
return &EncodeError{Path: fmt.Sprintf("%s[%d]", joinKey(ctx, name), i), Err: err}
2026-08-19 09:47:00 +02:00
}
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) {
// Map and slice elements arrive wrapped in interface{}; look through them.
for v.Kind() == reflect.Interface && !v.IsNil() {
v = v.Elem()
}
if v.Kind() == reflect.Interface {
return nil, fmt.Errorf("cannot encode nil value")
}
2026-08-19 09:47:00 +02:00
if v.CanInterface() {
if m, ok := v.Interface().(Marshaler); ok {
return m.MarshalTOML()
}
}
// The datetime structs are TOML scalars; the emitter renders each of them.
if t := v.Type(); t == timeGoType || isLocalDateType(t) {
return v.Interface(), nil
}
2026-08-19 09:47:00 +02:00
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.Map:
// A table nested in a value array has no header form, so it renders
// inline; the keys normalise to strings for the emitter.
if v.Type().Key().Kind() != reflect.String {
return nil, fmt.Errorf("map key must be string, got %s", v.Type().Key())
}
out := make(map[string]any, v.Len())
for _, k := range v.MapKeys() {
val, err := normaliseValue(v.MapIndex(k))
if err != nil {
return nil, fmt.Errorf("[%s]: %w", k.String(), err)
}
out[k.String()] = val
}
return out, nil
2026-08-19 09:47:00 +02:00
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 map[string]any:
return e.writeInlineTable(v)
2026-08-19 09:47:00 +02:00
case nil:
return fmt.Errorf("interpres: cannot encode nil value")
default:
return fmt.Errorf("interpres: cannot encode %T", val)
}
}
// writeInlineTable renders m as a TOML inline table with sorted keys, the
// order buildMapDoc uses for header tables. It backs the table elements of a
// value array, where the [[header]] form is not available.
func (e *encoder) writeInlineTable(m map[string]any) error {
keys := slices.Sorted(maps.Keys(m))
e.buf.WriteByte('{')
for i, k := range keys {
if i > 0 {
e.buf.WriteString(", ")
}
e.writeKey(k)
e.buf.WriteString(" = ")
if err := e.writeValue(m[k]); err != nil {
return err
}
}
e.buf.WriteByte('}')
return nil
}
2026-08-19 09:47:00 +02:00
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, the contract the output
// rules in the documentation state.
2026-08-19 09:47:00 +02:00
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
}