// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package interpres import ( "context" "encoding" "fmt" "reflect" "slices" "strings" "sync" "sync/atomic" "time" ) // decoder maps a parsed TOML tree onto Go values via reflection. ctx is the // context a cancellable entry point handed in, and reaches an // UnmarshalerContext destination; entry points without one leave it nil. // nodes is the document's node index, present only when a destination can // reach an OrderedMap and the parse built the tree its key order is read // from. loc is the zone a local date-time is carried in when it decodes into // a time.Time destination; nil keeps the wrapper-only default. type decoder struct { disallowUnknown bool ctx context.Context nodes nodeIndex loc *time.Location } func newDecoder() *decoder { return &decoder{} } // ctxOrBackground returns the context the decode carries, and Background when // none was given, so a custom decoder never receives a nil context. func (d *decoder) ctxOrBackground() context.Context { if d.ctx == nil { return context.Background() } return d.ctx } var timeType = reflect.TypeFor[time.Time]() var ( unmarshalerType = reflect.TypeFor[Unmarshaler]() ctxUnmarshalerType = reflect.TypeFor[UnmarshalerContext]() textUnmarshalerType = reflect.TypeFor[encoding.TextUnmarshaler]() numberType = reflect.TypeFor[Number]() ) // The per-type flags record which interface lookups a decode into that type // can succeed at, so the hot path consults the cache instead of boxing every // value into an interface to ask. The bits name the receiver the method is // found on: the value itself, or its address. const ( flagUnmarshaler uint8 = 1 << iota flagAddrUnmarshaler flagCtxUnmarshaler flagAddrCtxUnmarshaler flagTextUnmarshaler flagAddrTextUnmarshaler ) // typeFlagCache holds one flag entry per destination type. A set is immutable // once published, the same trade-off structSchemaCache makes; the cache grows // with the number of distinct types decoded, never per document. The hint // below re-points at these published entries, so a hot lookup allocates // nothing. var typeFlagCache sync.Map // reflect.Type -> *flagHintEntry // flagHintEntry pairs a type with its cached flags for the monomorphic hint // below. Both caches share the entry shape. type flagHintEntry struct { typ reflect.Type flags uint8 } // typeFlagHint remembers the entry resolved last, because a decode walks one // type across consecutive fields and elements. A lost race loses only the // hint: every value it can hold came from the cache. var typeFlagHint atomic.Pointer[flagHintEntry] func typeFlags(t reflect.Type) uint8 { if e := typeFlagHint.Load(); e != nil && e.typ == t { return e.flags } if v, ok := typeFlagCache.Load(t); ok { entry := v.(*flagHintEntry) typeFlagHint.Store(entry) return entry.flags } var f uint8 if t.Implements(unmarshalerType) { f |= flagUnmarshaler } if t.Implements(ctxUnmarshalerType) { f |= flagCtxUnmarshaler } pt := reflect.PointerTo(t) if pt.Implements(unmarshalerType) { f |= flagAddrUnmarshaler } if pt.Implements(ctxUnmarshalerType) { f |= flagAddrCtxUnmarshaler } // The date-time types are excluded from the text path: they carry // time.Time's UnmarshalText through an embedded field while their only // accepted form is a bare timestamp. if !isDateTimeType(t) { if t.Implements(textUnmarshalerType) { f |= flagTextUnmarshaler } if pt.Implements(textUnmarshalerType) { f |= flagAddrTextUnmarshaler } } actual, _ := typeFlagCache.LoadOrStore(t, &flagHintEntry{t, f}) published := actual.(*flagHintEntry) typeFlagHint.Store(published) return published.flags } // unmarshalerOf resolves the Unmarshaler for dst through the flag cache, so // an interface value is built only where the cache says the assertion can // succeed. An interface destination is asked dynamically, because the value // it will hold may implement the interface even when the interface type // itself does not. func unmarshalerOf(dst reflect.Value) (Unmarshaler, bool) { if dst.Kind() == reflect.Interface { u, ok := dst.Interface().(Unmarshaler) return u, ok } f := typeFlags(dst.Type()) if f&flagUnmarshaler != 0 { u, ok := dst.Interface().(Unmarshaler) return u, ok } if f&flagAddrUnmarshaler != 0 && dst.CanAddr() { u, ok := dst.Addr().Interface().(Unmarshaler) return u, ok } return nil, false } // ctxUnmarshalerOf is the same resolution for UnmarshalerContext. func ctxUnmarshalerOf(dst reflect.Value) (UnmarshalerContext, bool) { if dst.Kind() == reflect.Interface { u, ok := dst.Interface().(UnmarshalerContext) return u, ok } f := typeFlags(dst.Type()) if f&flagCtxUnmarshaler != 0 { u, ok := dst.Interface().(UnmarshalerContext) return u, ok } if f&flagAddrCtxUnmarshaler != 0 && dst.CanAddr() { u, ok := dst.Addr().Interface().(UnmarshalerContext) return u, ok } return nil, false } 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 UnmarshalerContext get the context beside the parsed // data, and are responsible for setting their own state. They win over // Unmarshaler, which wins over the text path. The lookups cover both T and // *T so a pointer-receiver method is invoked on an addressable struct // field. if dst.CanInterface() { if u, ok := ctxUnmarshalerOf(dst); ok { if err := u.UnmarshalTOMLContext(d.ctxOrBackground(), data); err != nil { return fmt.Errorf("unmarshal: %w", err) } return nil } 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 } } // A TOML string fills a destination that implements // encoding.TextUnmarshaler, the rule encoding/json follows. Every other // value kind keeps its own rule, so an integer still reaches a numeric // destination. if s, isString := data.(string); isString { if tu, ok := textUnmarshalerOf(dst); ok { if err := tu.UnmarshalText([]byte(s)); err != nil { return fmt.Errorf("unmarshal text: %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: if dst.Type() == durationType { return setDuration(dst, v) } return setBasic(dst, reflect.ValueOf(v), "string") case Number: return setNumber(dst, v) case bool: return setBasic(dst, reflect.ValueOf(v), "bool") case int64: return setInt(dst, v) case float64: return setFloat(dst, v) case OffsetDateTime: return setOffsetDateTime(v, dst) case time.Time: return setDateTime(v, dst) case LocalDateTime: if dst.Type() == localDateTimeType { dst.Set(reflect.ValueOf(v)) return nil } return d.setLocalTimeValue(v.Time, dst) case LocalDate: if dst.Type() == localDateType { dst.Set(reflect.ValueOf(v)) return nil } return d.setLocalTimeValue(v.Time, dst) case LocalTime: if dst.Type() == localTimeType { dst.Set(reflect.ValueOf(v)) return nil } return d.setLocalTimeValue(v.Time, dst) 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()) } } // textUnmarshalerOf is the same resolution for encoding.TextUnmarshaler, // with the date-time types excluded for the reason typeFlags records. func textUnmarshalerOf(dst reflect.Value) (encoding.TextUnmarshaler, bool) { if !dst.CanInterface() || isDateTimeType(dst.Type()) { return nil, false } if dst.Kind() == reflect.Interface { tu, ok := dst.Interface().(encoding.TextUnmarshaler) return tu, ok } f := typeFlags(dst.Type()) if f&flagTextUnmarshaler != 0 { tu, ok := dst.Interface().(encoding.TextUnmarshaler) return tu, ok } if f&flagAddrTextUnmarshaler != 0 && dst.CanAddr() { tu, ok := dst.Addr().Interface().(encoding.TextUnmarshaler) return tu, ok } return nil, false } func (d *decoder) assignTable(tbl map[string]any, dst reflect.Value) error { if dst.Type() == orderedMapType { return d.fillOrderedMap(tbl, dst) } 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 { schema := cachedStructSchema(dst.Type()) if d.disallowUnknown { // Map iteration order is random, so pick the unknown key to report // deterministically: the smallest one. unknown := "" for key := range tbl { if _, ok := schema.byName[key]; ok { continue } if _, ok := schema.byName[strings.ToLower(key)]; ok { continue } if unknown == "" || key < unknown { unknown = key } } if unknown != "" { return fmt.Errorf("interpres: unknown field %q for %s", unknown, dst.Type()) } } // The keys that resolved to a field are remembered while the table walks, // but only a struct that demands one pays for the set. var seen map[string]bool if len(schema.required) > 0 { seen = make(map[string]bool, len(tbl)) } for key, val := range tbl { // A key that is already lowercase, which document keys usually are, // hits the map directly; only a miss pays for the case fold. resolved := key field, ok := schema.byName[key] if !ok { resolved = strings.ToLower(key) field, ok = schema.byName[resolved] } if !ok { if schema.embedMaps != nil { // Leftover keys land in an untagged embedded map, the inverse // of the encoder inlining that map's entries. mv, err := fieldByIndex(dst, schema.embedMaps[0]) if err != nil { return newDecodeError(key, err) } if err := d.assignMap(map[string]any{key: val}, mv); err != nil { return newDecodeError(key, err) } } continue } if seen != nil { seen[resolved] = true } fv, err := fieldByIndex(dst, field.index) if err != nil { return newDecodeError(key, err) } if err := d.assign(val, fv); err != nil { return newDecodeError(key, err) } } for _, key := range schema.required { if !seen[key] { return fmt.Errorf("interpres: missing required key %q", key) } } 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 newDecodeError(key, err) } dst.SetMapIndex(reflect.ValueOf(key), elem) } return nil } func (d *decoder) assignSlice(items []any, dst reflect.Value) error { switch dst.Kind() { case reflect.Slice: 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 newDecodeError(fmt.Sprintf("[%d]", i), err) } } dst.Set(out) return nil case reflect.Array: // A fixed-size array takes the elements in place; a length mismatch is // the error, because a TOML array carries no way to name a default for // the elements it is short of, and the surplus has nowhere to go. if dst.Len() != len(items) { return fmt.Errorf("interpres: cannot assign %d elements to %s", len(items), dst.Type()) } for i, item := range items { if err := d.assign(item, dst.Index(i)); err != nil { return newDecodeError(fmt.Sprintf("[%d]", i), err) } } return nil default: return fmt.Errorf("interpres: cannot assign array to %s", dst.Type()) } } func (d *decoder) assignTableSlice(items []map[string]any, dst reflect.Value) error { switch dst.Kind() { case reflect.Slice: 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 newDecodeError(fmt.Sprintf("[%d]", i), err) } } dst.Set(out) return nil case reflect.Array: if dst.Len() != len(items) { return fmt.Errorf("interpres: cannot assign %d elements to %s", len(items), dst.Type()) } for i, item := range items { if err := d.assign(item, dst.Index(i)); err != nil { return newDecodeError(fmt.Sprintf("[%d]", i), err) } } return nil default: return fmt.Errorf("interpres: cannot assign array of tables to %s", dst.Type()) } } // --- low-level setters ----------------------------------------------------- // setOffsetDateTime stores an offset date-time: in a wrapper destination as it // is, and in a plain time.Time, which takes the instant with the offset the // document wrote, so a timestamp field does not have to name the wrapper. func setOffsetDateTime(v OffsetDateTime, dst reflect.Value) error { switch dst.Type() { case offsetDateTimeType: dst.Set(reflect.ValueOf(v)) case timeType: dst.Set(reflect.ValueOf(v.Time)) default: return fmt.Errorf("interpres: cannot assign datetime to %s", dst.Type()) } return nil } // setDateTime stores a time.Time that reached the tree directly, which is the // shape a tree built by hand carries. Dates the parser produced arrive as // OffsetDateTime instead. func setDateTime(v time.Time, dst reflect.Value) error { switch dst.Type() { case timeType: dst.Set(reflect.ValueOf(v)) case offsetDateTimeType: dst.Set(reflect.ValueOf(OffsetDateTime{v})) default: return fmt.Errorf("interpres: cannot assign datetime to %s", dst.Type()) } return nil } // setLocalTimeValue stores a local date-time value into a plain time.Time // destination, which the decoder permits only when LocalTimeLocation fixed // the zone the wall-clock value is carried in; without it the wrapper types // are the only destinations a local kind fills, as they always have been. func (d *decoder) setLocalTimeValue(t time.Time, dst reflect.Value) error { if dst.Type() == timeType { if d.loc != nil { // A local value is a wall clock, so the zone choice relabels it // rather than shifting the instant: 07:32 in the document is // 07:32 in the location, not an hour later. dst.Set(reflect.ValueOf(time.Date( t.Year(), t.Month(), t.Day(), t.Hour(), t.Minute(), t.Second(), t.Nanosecond(), d.loc))) return nil } return fmt.Errorf("interpres: cannot assign local date-time to time.Time; set Decoder.LocalTimeLocation to choose the zone") } return fmt.Errorf("interpres: cannot assign local date-time to %s", dst.Type()) } 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()) } // Convert rather than assign: a value of the predeclared type is not // assignable to a defined type of the same kind, so a plain Set panics on // a destination such as `type Name string`. dst.Set(val.Convert(dst.Type())) return nil } // setDuration reads a duration literal into a time.Duration destination. TOML // has no duration type, so the encoder writes the canonical Go form and the // decoder reads that back; a bare integer stays the nanosecond count it has // always been, and reaches the destination through setInt. func setDuration(dst reflect.Value, s string) error { d, err := time.ParseDuration(s) if err != nil { return fmt.Errorf("interpres: invalid duration %q", s) } dst.SetInt(int64(d)) return nil } // setNumber stores a Number, the literal NumbersAsLiterals keeps. A Number destination // takes the literal as it is; every other destination takes the evaluated // value through the ordinary rules, so an integer field, a float field and a // duration field all read a Number the way they read the evaluated kind. func setNumber(dst reflect.Value, n Number) error { if dst.Type() == numberType { dst.SetString(string(n)) return nil } v, err := decodeNumber(string(n)) if err != nil { return fmt.Errorf("interpres: %w", err) } switch v := v.(type) { case int64: return setInt(dst, v) case float64: return setFloat(dst, v) } return fmt.Errorf("interpres: cannot assign number to %s", dst.Type()) } 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()) } // OverflowUint knows every width, uint included on platforms where it // is narrower than uint64; SetUint would silently truncate instead. if dst.OverflowUint(uint64(v)) { return fmt.Errorf("interpres: integer %d overflows %s", v, dst.Type()) } dst.SetUint(uint64(v)) case reflect.Float32, reflect.Float64: // A finite value beyond the float32 range would silently become ±Inf; // infinities and NaN themselves pass through. An int64 never // overflows either float width. f := float64(v) if dst.OverflowFloat(f) { return fmt.Errorf("interpres: integer %d overflows %s", v, dst.Type()) } dst.SetFloat(f) 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: if dst.OverflowFloat(v) { return fmt.Errorf("interpres: float %g overflows %s", v, dst.Type()) } dst.SetFloat(v) return nil default: return fmt.Errorf("interpres: cannot assign float to %s", dst.Type()) } } // structFieldLoc locates one destination field by its index path from the // struct root and by the depth the field sits at, which breaks name clashes // in favour of the shallower field. required records the tag option of the // field that won the name. type structFieldLoc struct { index []int depth int required bool } // structSchema flattens the exported fields of t for decode, mirroring the // encoder: an untagged embedded struct is inlined, so its own fields match // keys of the same table, and an untagged embedded map is recorded in // embedMaps (first declaration first) as the destination for leftover keys. // When two fields resolve to one name, the shallower wins, then the later // declaration. required holds the keys a `toml:"...,required"` tag demands. type structSchema struct { byName map[string]structFieldLoc embedMaps [][]int required []string } // structSchemaCache holds one schema per struct type. A schema is immutable // once published, so concurrent callers only race to build an identical value, // the same trade-off encoding/json's field cache makes. The cache grows with // the number of distinct types decoded or encoded, never per document. var structSchemaCache sync.Map // reflect.Type -> structSchema func cachedStructSchema(t reflect.Type) structSchema { if s, ok := structSchemaCache.Load(t); ok { return s.(structSchema) } s := newStructSchema(t) actual, _ := structSchemaCache.LoadOrStore(t, s) return actual.(structSchema) } func newStructSchema(t reflect.Type) structSchema { s := structSchema{byName: make(map[string]structFieldLoc, t.NumField())} // A struct may embed a pointer to itself, which is legal Go, so the walk // tracks the struct types on the current path and stops when one repeats; // without the guard the recursion never terminates. A self-promoted key // always loses to the shallower original, so skipping it changes nothing. visiting := map[reflect.Type]bool{} var walk func(t reflect.Type, prefix []int, depth int) walk = func(t reflect.Type, prefix []int, depth int) { visiting[t] = true defer delete(visiting, t) for i := range t.NumField() { f := t.Field(i) if f.PkgPath != "" { // unexported continue } path := append(append([]int{}, prefix...), i) name := "" required := false if tag, ok := f.Tag.Lookup("toml"); ok { var opts string name, opts, _ = strings.Cut(tag, ",") if name == "-" { continue } for opts != "" { var opt string opt, opts, _ = strings.Cut(opts, ",") if opt == "required" { required = true } } } if f.Anonymous && name == "" { ft := f.Type for ft.Kind() == reflect.Pointer { ft = ft.Elem() } switch { case ft.Kind() == reflect.Struct && !isScalarStruct(ft): if !visiting[ft] { walk(ft, path, depth+1) } continue case ft.Kind() == reflect.Map && ft.Key().Kind() == reflect.String: s.embedMaps = append(s.embedMaps, path) continue } name = f.Name } if name == "" { name = f.Name } key := strings.ToLower(name) if existing, ok := s.byName[key]; !ok || depth <= existing.depth { s.byName[key] = structFieldLoc{index: path, depth: depth, required: required} } } } walk(t, nil, 0) // The missing-key error must not depend on map order, so the demanded keys // come out sorted. for key, loc := range s.byName { if loc.required { s.required = append(s.required, key) } } slices.Sort(s.required) return s } // ownsKey reports whether the field at path is the one that resolves key. // The encoder consults it to emit exactly the field the decoder would fill, // so a struct with two fields mapping to one key does not marshal into a // duplicate TOML key. func (s structSchema) ownsKey(key string, path []int) bool { loc, ok := s.byName[key] return ok && slices.Equal(loc.index, path) } // fieldByIndex walks an index path from a struct value, allocating nil // pointers along the way so a key can reach through an embedded pointer // struct. Every field on the path is exported, so each step is settable. func fieldByIndex(v reflect.Value, path []int) (reflect.Value, error) { for i, x := range path { v = v.Field(x) if i < len(path)-1 && v.Kind() == reflect.Pointer { if v.IsNil() { if !v.CanSet() { return reflect.Value{}, fmt.Errorf("cannot allocate nil embedded pointer") } v.Set(reflect.New(v.Type().Elem())) } v = v.Elem() } } return v, nil }