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
+10
View File
@@ -46,6 +46,16 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
(10000 levels, which no hand-written document approaches): a document that
nests arrays or inline tables deeper used to run the stack out and is now
rejected with a `SyntaxError` naming the limit.
- The `toml` tag gained the `required` option: a field tagged
`toml:"host,required"` makes the decode fail with
`missing required key "host"` when the document carries no key that
resolves to it. The option shapes decoding only, and the encoder ignores
it.
- `UnmarshalerContext`, the custom-decode interface that hands the decode's
context to the method, `UnmarshalTOMLContext(ctx, data)`. It wins over
`UnmarshalTOML` when a type implements both, so a long custom decode can
abort on cancellation; a non-cancellable entry point hands in
`context.Background`, never nil.
- A TOML array decodes into a Go fixed-size array, `[N]T`, where only a slice
was accepted before; the encoder could already encode one. A length mismatch
is an error wrapped with the key path.
+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
}
+107
View File
@@ -1316,3 +1316,110 @@ func TestDecodeFixedArray(t *testing.T) {
}
})
}
func TestRequiredTag(t *testing.T) {
type Config struct {
Host string `toml:"host,required"`
Radius int `toml:"radius"`
}
t.Run("a present key satisfies the tag", func(t *testing.T) {
var cfg Config
if err := Unmarshal([]byte("radius = 2\nhost = \"example.org\"\n"), &cfg); err != nil {
t.Fatal(err)
}
if cfg.Host != "example.org" || cfg.Radius != 2 {
t.Errorf("decoded %+v", cfg)
}
})
t.Run("a missing key is an error", func(t *testing.T) {
var cfg Config
err := Unmarshal([]byte("radius = 2\n"), &cfg)
want := `interpres: missing required key "host"`
if err == nil || err.Error() != want {
t.Errorf("err = %v, want %q", err, want)
}
})
t.Run("the error carries the key path", func(t *testing.T) {
var outer struct {
Server Config `toml:"server"`
}
err := Unmarshal([]byte("[server]\nradius = 1\n"), &outer)
want := `server: interpres: missing required key "host"`
if err == nil || err.Error() != want {
t.Errorf("err = %v, want %q", err, want)
}
})
t.Run("case-insensitive match satisfies the tag", func(t *testing.T) {
var cfg Config
if err := Unmarshal([]byte("HOST = \"x\"\n"), &cfg); err != nil {
t.Errorf("err = %v, want nil", err)
}
})
}
type ctxRecorder struct {
got context.Context
value any
}
func (r *ctxRecorder) UnmarshalTOMLContext(ctx context.Context, data any) error {
r.got = ctx
r.value = data
return nil
}
func TestUnmarshalerContext(t *testing.T) {
t.Run("the context reaches the method", func(t *testing.T) {
type keyT struct{}
ctx := context.WithValue(context.Background(), keyT{}, "sentinel")
var r ctxRecorder
if err := UnmarshalContext(ctx, []byte("a = 1\n"), &r); err != nil {
t.Fatal(err)
}
if v, _ := r.got.Value(keyT{}).(string); v != "sentinel" {
t.Errorf("ctx = %v, want the caller's context", r.got)
}
tree, isMap := r.value.(map[string]any)
if !isMap || tree["a"] != int64(1) {
t.Errorf("value = %#v, want the tree with a = 1", r.value)
}
})
t.Run("the context wins over Unmarshaler", func(t *testing.T) {
var v struct {
R ctxBoth `toml:"r"`
}
if err := Unmarshal([]byte("r = 1\n"), &v); err != nil {
t.Fatal(err)
}
if !v.R.ctxCalled {
t.Error("UnmarshalTOMLContext was not called")
}
if v.R.plainCalled {
t.Error("UnmarshalTOML was called although the context method exists")
}
})
t.Run("a non-cancellable entry point hands in Background", func(t *testing.T) {
var r ctxRecorder
if err := Unmarshal([]byte("a = 1\n"), &r); err != nil {
t.Fatal(err)
}
if r.got != context.Background() {
t.Errorf("ctx = %v, want context.Background", r.got)
}
})
}
type ctxBoth struct {
ctxCalled bool
plainCalled bool
}
func (b *ctxBoth) UnmarshalTOMLContext(ctx context.Context, data any) error {
b.ctxCalled = true
return nil
}
func (b *ctxBoth) UnmarshalTOML(data any) error {
b.plainCalled = true
return nil
}
+23 -1
View File
@@ -225,6 +225,11 @@ one declared later wins.
Unknown keys are ignored by default, landing in an untagged embedded map when
the struct has one; [Strict decoding](#strict-decoding) rejects them instead.
The tag may carry the `required` option, `toml:"host,required"`: the decode
fails with `missing required key "host"` when no key of the document resolved
to the field. The check runs after the table is read, so the other fields
carry their values whether the required one is present or not.
### Numeric conversion
The parser produces `int64` for every integer and `float64` for every float.
@@ -314,6 +319,21 @@ automatically, and a nil pointer destination is allocated first. An error
returned from `UnmarshalTOML` halts the decode and propagates wrapped with the
key path, for example `addr: unmarshal: not a string`.
### Custom decoding: `UnmarshalerContext`
`UnmarshalerContext` is `Unmarshaler` with the decode's context handed in:
```go
type UnmarshalerContext interface {
UnmarshalTOMLContext(ctx context.Context, data any) error
}
```
A type that implements both gets `UnmarshalTOMLContext`, so a long custom
decode can abort on cancellation instead of running to completion. The
context a non-cancellable entry point carries is `context.Background`, never
nil.
### Custom decoding: `encoding.TextUnmarshaler`
A destination type that implements `encoding.TextUnmarshaler` receives a TOML
@@ -752,7 +772,9 @@ See [Custom encoding](#custom-encoding-marshaler).
### `type Unmarshaler interface{ UnmarshalTOML(data any) error }`
See [Custom decoding](#custom-decoding-unmarshaler).
See [Custom decoding](#custom-decoding-unmarshaler). `UnmarshalerContext`
carries the decode's context through `UnmarshalTOMLContext(ctx, data)` and
wins when a type implements both.
### `type Number string`
+15 -2
View File
@@ -241,7 +241,9 @@ func UnmarshalContext(ctx context.Context, data []byte, v any) error {
if err != nil {
return err
}
return newDecoder().decode(tree, v)
dec := newDecoder()
dec.ctx = ctx
return dec.decode(tree, v)
}
// A Decoder decodes a TOML document into a Go value with configurable
@@ -315,6 +317,7 @@ func (d *Decoder) DecodeContext(ctx context.Context, data []byte, v any) error {
}
dec := newDecoder()
dec.disallowUnknown = d.disallowUnknown
dec.ctx = ctx
return dec.decode(tree, v)
}
@@ -336,7 +339,8 @@ type Marshaler interface {
// argument is whatever the parser produced for that key: one of string,
// bool, int64, float64, OffsetDateTime, LocalDateTime, LocalDate, LocalTime,
// []any, or map[string]any. A tree built by hand may carry a plain time.Time
// where the parser would put an OffsetDateTime.
// where the parser would put an OffsetDateTime, and a Decoder configured with
// UseNumber a Number.
//
// UnmarshalTOML may parse, inspect, or transform the value however it likes,
// then store the result by mutating its receiver through the standard
@@ -355,6 +359,15 @@ type Unmarshaler interface {
UnmarshalTOML(data any) error
}
// UnmarshalerContext is Unmarshaler with the decode's context handed in. A
// type that implements both interfaces gets UnmarshalTOMLContext, so a long
// custom decode can abort on cancellation instead of running to completion.
// The context a non-cancellable entry point carries is context.Background,
// never nil.
type UnmarshalerContext interface {
UnmarshalTOMLContext(ctx context.Context, data any) error
}
// Marshal returns the TOML encoding of v. The output is valid TOML 1.1.
//
// Marshal traverses v using reflection and applies the following rules: