369 lines
11 KiB
Go
369 lines
11 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|
// SPDX-License-Identifier: MIT
|
|
|
|
package interpres
|
|
|
|
import (
|
|
"fmt"
|
|
"reflect"
|
|
"slices"
|
|
"strings"
|
|
"sync"
|
|
"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 {
|
|
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[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 {
|
|
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 -----------------------------------------------------
|
|
|
|
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())
|
|
}
|
|
// 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
|
|
}
|