From 0ba145ba0c703456c8cca4c7071d29692bc2886a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Petr=20Balv=C3=ADn?= Date: Tue, 22 Sep 2026 00:04:16 +0200 Subject: [PATCH] feat: add the required tag option and UnmarshalerContext Assisted-by: GLM 5.3 Flash --- CHANGELOG.md | 10 +++++ decode.go | 108 +++++++++++++++++++++++++++++++++++++++++++------ decode_test.go | 107 ++++++++++++++++++++++++++++++++++++++++++++++++ docs/API.md | 24 ++++++++++- interpres.go | 17 +++++++- 5 files changed, 250 insertions(+), 16 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index f3c2018..496dd53 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -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. diff --git a/decode.go b/decode.go index e7d39e1..512b9cd 100644 --- a/decode.go +++ b/decode.go @@ -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 } diff --git a/decode_test.go b/decode_test.go index b06dc73..c75f6fe 100644 --- a/decode_test.go +++ b/decode_test.go @@ -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 +} diff --git a/docs/API.md b/docs/API.md index 6ef5154..d6fe2be 100644 --- a/docs/API.md +++ b/docs/API.md @@ -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` diff --git a/interpres.go b/interpres.go index 98962d1..7db463d 100644 --- a/interpres.go +++ b/interpres.go @@ -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: