Files
interpres/decode.go
T

328 lines
9.2 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 (
"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 {
schema := newStructSchema(dst.Type())
2026-08-19 09:47:00 +02:00
for key, val := range tbl {
field, ok := schema.byName[strings.ToLower(key)]
2026-08-19 09:47:00 +02:00
if !ok {
if d.disallowUnknown {
return fmt.Errorf("interpres: unknown field %q for %s", key, dst.Type())
}
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)
}
}
2026-08-19 09:47:00 +02:00
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)
2026-08-19 09:47:00 +02:00
}
}
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)
2026-08-19 09:47:00 +02:00
}
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)
2026-08-19 09:47:00 +02:00
}
}
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)
2026-08-19 09:47:00 +02:00
}
}
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())
}
}
// 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
}
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
2026-08-19 09:47:00 +02:00
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}
2026-08-19 09:47:00 +02:00
}
}
}
walk(t, nil, 0)
return s
}
// 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
2026-08-19 09:47:00 +02:00
}