552 lines
17 KiB
Go
552 lines
17 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|
// SPDX-License-Identifier: MIT
|
|
|
|
package interpres
|
|
|
|
import (
|
|
"encoding"
|
|
"fmt"
|
|
"reflect"
|
|
"slices"
|
|
"strings"
|
|
"sync"
|
|
"sync/atomic"
|
|
"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]()
|
|
|
|
var (
|
|
unmarshalerType = reflect.TypeFor[Unmarshaler]()
|
|
textUnmarshalerType = reflect.TypeFor[encoding.TextUnmarshaler]()
|
|
)
|
|
|
|
// 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
|
|
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
|
|
}
|
|
pt := reflect.PointerTo(t)
|
|
if pt.Implements(unmarshalerType) {
|
|
f |= flagAddrUnmarshaler
|
|
}
|
|
// 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
|
|
}
|
|
|
|
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
|
|
}
|
|
}
|
|
|
|
// 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 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)
|
|
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 {
|
|
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())
|
|
}
|
|
}
|
|
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.
|
|
field, ok := schema.byName[key]
|
|
if !ok {
|
|
field, ok = schema.byName[strings.ToLower(key)]
|
|
}
|
|
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
|
|
}
|
|
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)
|
|
}
|
|
}
|
|
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 {
|
|
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 newDecodeError(fmt.Sprintf("[%d]", 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 newDecodeError(fmt.Sprintf("[%d]", i), err)
|
|
}
|
|
}
|
|
dst.Set(out)
|
|
return nil
|
|
}
|
|
|
|
// --- 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
|
|
}
|
|
|
|
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
|
|
}
|
|
|
|
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.
|
|
type structFieldLoc struct {
|
|
index []int
|
|
depth int
|
|
}
|
|
|
|
// 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.
|
|
type structSchema struct {
|
|
byName map[string]structFieldLoc
|
|
embedMaps [][]int
|
|
}
|
|
|
|
// 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 := ""
|
|
if tag, ok := f.Tag.Lookup("toml"); ok {
|
|
name, _, _ = strings.Cut(tag, ",")
|
|
if name == "-" {
|
|
continue
|
|
}
|
|
}
|
|
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}
|
|
}
|
|
}
|
|
}
|
|
walk(t, nil, 0)
|
|
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
|
|
}
|