feat: add the required tag option and UnmarshalerContext
Test / test (push) Canceled after 39s

Assisted-by: GLM 5.3 Flash
This commit is contained in:
2026-09-22 00:04:16 +02:00
parent 2e5dfc54c9
commit 0ba145ba0c
5 changed files with 250 additions and 16 deletions
+95 -13
View File
@@ -4,6 +4,7 @@
package interpres
import (
"context"
"encoding"
"fmt"
"reflect"
@@ -14,17 +15,30 @@ import (
"time"
)
// decoder maps a parsed TOML tree onto Go values via reflection.
// 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.
type decoder struct {
disallowUnknown bool
ctx context.Context
}
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]()
)
@@ -36,6 +50,8 @@ var (
const (
flagUnmarshaler uint8 = 1 << iota
flagAddrUnmarshaler
flagCtxUnmarshaler
flagAddrCtxUnmarshaler
flagTextUnmarshaler
flagAddrTextUnmarshaler
)
@@ -72,10 +88,16 @@ func typeFlags(t reflect.Type) 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.
@@ -115,6 +137,24 @@ func unmarshalerOf(dst reflect.Value) (Unmarshaler, bool) {
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() {
@@ -139,12 +179,18 @@ func (d *decoder) assign(data any, dst reflect.Value) error {
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.
// 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)
@@ -258,12 +304,20 @@ func (d *decoder) assignStruct(tbl map[string]any, dst reflect.Value) error {
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 {
field, ok = schema.byName[strings.ToLower(key)]
resolved = strings.ToLower(key)
field, ok = schema.byName[resolved]
}
if !ok {
if schema.embedMaps != nil {
@@ -279,6 +333,9 @@ func (d *decoder) assignStruct(tbl map[string]any, dst reflect.Value) error {
}
continue
}
if seen != nil {
seen[resolved] = true
}
fv, err := fieldByIndex(dst, field.index)
if err != nil {
return newDecodeError(key, err)
@@ -287,6 +344,11 @@ func (d *decoder) assignStruct(tbl map[string]any, dst reflect.Value) error {
return newDecodeError(key, err)
}
}
for _, key := range schema.required {
if !seen[key] {
return fmt.Errorf("interpres: missing required key %q", key)
}
}
return nil
}
@@ -488,10 +550,12 @@ func setFloat(dst reflect.Value, v float64) error {
// 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.
// 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
index []int
depth int
required bool
}
// structSchema flattens the exported fields of t for decode, mirroring the
@@ -499,10 +563,11 @@ type structFieldLoc struct {
// 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.
// 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
@@ -538,11 +603,20 @@ func newStructSchema(t reflect.Type) structSchema {
}
path := append(append([]int{}, prefix...), i)
name := ""
required := false
if tag, ok := f.Tag.Lookup("toml"); ok {
name, _, _ = strings.Cut(tag, ",")
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
@@ -566,11 +640,19 @@ func newStructSchema(t reflect.Type) structSchema {
}
key := strings.ToLower(name)
if existing, ok := s.byName[key]; !ok || depth <= existing.depth {
s.byName[key] = structFieldLoc{index: path, depth: 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
}