diff --git a/encode.go b/encode.go index ba43a58..56fb517 100644 --- a/encode.go +++ b/encode.go @@ -132,6 +132,10 @@ type encoder struct { // their indentation. inlineDepth int + // valueDepth is the nesting level of boxed containers the value writer + // walks, the bound a cyclic map or slice hits instead of the stack. + valueDepth int + // limit is the column at which an inline table is broken; only a // measuring encoder raises it. limit int @@ -950,9 +954,11 @@ func addArrayValue(doc *tomlDoc, name string, v reflect.Value, path encPath, for return nil } -// writeArrayValue writes a value array from its reflect value, its elements -// written one by one, each falling back to the boxed path only where the -// boxed rules rewrite it. +// writeArrayValue writes a value array from its reflect value. Only the +// plain scalar kinds reach it: addArrayValue's direct path takes nothing but +// arrays whose elements are scalar kinds without methods, so one writer per +// element needs no boxing branch and no depth walk of its own; the slices +// and maps the boxed rules re-write travel the normaliseValue route. func (e *encoder) writeArrayValue(v reflect.Value, depth int) error { if atDepthLimit(depth) { return errDepthLimit() @@ -962,7 +968,7 @@ func (e *encoder) writeArrayValue(v reflect.Value, depth int) error { if i > 0 { e.buf.WriteString(", ") } - if err := e.writeArrayElem(v.Index(i), depth+1); err != nil { + if err := e.writeScalarValue(v.Index(i)); err != nil { return err } } @@ -970,129 +976,6 @@ func (e *encoder) writeArrayValue(v reflect.Value, depth int) error { return nil } -// writeArrayElem writes one element of a value array. The kinds the boxed -// rules rewrite are handed to normaliseValue and the boxed writer; the rest -// write directly, including nested arrays and inline tables. -func (e *encoder) writeArrayElem(v reflect.Value, depth int) error { - if v.Kind() == reflect.Interface { - if v.IsNil() { - return fmt.Errorf("interpres: cannot encode nil value") - } - v = v.Elem() - } - switch v.Kind() { - case reflect.Slice, reflect.Array: - return e.writeArrayValue(v, depth) - case reflect.Map: - return e.writeInlineMapFromReflect(v, depth) - } - if _, isMarshaler := marshalerOf(v); isMarshaler { - return e.writeNormalisedElem(v, depth) - } - if _, isText, err := textValue(v); err != nil || isText { - if err != nil { - return err - } - return e.writeNormalisedElem(v, depth) - } - switch v.Kind() { - case reflect.String, reflect.Bool, reflect.Int, reflect.Int8, reflect.Int16, - reflect.Int32, reflect.Int64, reflect.Uint, reflect.Uint8, reflect.Uint16, - reflect.Uint32, reflect.Uint64, reflect.Float32, reflect.Float64: - if t := v.Type(); t != durationType && t != numberType { - return e.writeScalarValue(v) - } - } - return e.writeNormalisedElem(v, depth) -} - -// writeNormalisedElem normalises one element through the boxed rules and -// writes the result. -func (e *encoder) writeNormalisedElem(v reflect.Value, depth int) error { - val, err := normaliseValueAt(v, depth) - if err != nil { - return err - } - return e.writeValue(val) -} - -// writeInlineMapFromReflect renders a map from its reflect value as an inline -// table, the shape the boxed path gives the table elements of a value array: -// sorted keys, single line when it fits, across lines when it does not. -func (e *encoder) writeInlineMapFromReflect(v reflect.Value, depth int) error { - if e.limit >= noInlineBreak { - return e.writeInlineMapFlatReflect(v, depth) - } - flat := e.flat() - err := flat.writeInlineMapFlatReflect(v, depth) - if err != nil { - flat.release() - return err - } - fits := e.column()+flat.buf.Len() <= e.limit - if fits { - e.buf.Write(flat.buf.Bytes()) - } - flat.release() - if fits { - return nil - } - return e.writeInlineMapMultilineReflect(v, depth) -} - -// writeInlineMapFlatReflect renders the single-line form. -func (e *encoder) writeInlineMapFlatReflect(v reflect.Value, depth int) error { - if v.Type().Key().Kind() != reflect.String { - return fmt.Errorf("map key must be string, got %s", v.Type().Key()) - } - keys := make([]string, 0, v.Len()) - for _, k := range v.MapKeys() { - keys = append(keys, k.String()) - } - slices.Sort(keys) - e.buf.WriteByte('{') - for i, k := range keys { - if i > 0 { - e.buf.WriteString(", ") - } - if err := e.writeKey(k); err != nil { - return err - } - e.buf.WriteString(" = ") - if err := e.writeArrayElem(v.MapIndex(reflect.ValueOf(k)), depth+1); err != nil { - return err - } - } - e.buf.WriteByte('}') - return nil -} - -// writeInlineMapMultilineReflect renders the across-lines form. -func (e *encoder) writeInlineMapMultilineReflect(v reflect.Value, depth int) error { - keys := make([]string, 0, v.Len()) - for _, k := range v.MapKeys() { - keys = append(keys, k.String()) - } - slices.Sort(keys) - e.buf.WriteString("{\n") - e.inlineDepth++ - for _, k := range keys { - e.writeInlineIndent() - if err := e.writeKey(k); err != nil { - return err - } - e.buf.WriteString(" = ") - if err := e.writeArrayElem(v.MapIndex(reflect.ValueOf(k)), depth+1); err != nil { - return err - } - e.buf.WriteString(",\n") - } - e.inlineDepth-- - e.writeInlineIndent() - e.buf.WriteByte('}') - return nil -} - // marshalerOf finds the Marshaler a value carries: on the value itself, or on // its address, so a pointer-receiver MarshalTOML is found on an addressable // struct field or slice element, exactly as textMarshalerOf finds MarshalText. @@ -1364,7 +1247,13 @@ func textMarshalerOf(v reflect.Value) (encoding.TextMarshaler, bool) { return nil, false } +// isTableElementType reports whether a slice of t is an array of tables. A +// pointer element is looked through, so []*T behaves as []T: the empty-slice +// decision and the header form agree on the same element type. func isTableElementType(t reflect.Type) bool { + for t.Kind() == reflect.Pointer { + t = t.Elem() + } switch t.Kind() { case reflect.Struct: return !isScalarStruct(t) && !isTextMarshalerType(t) @@ -1395,12 +1284,24 @@ func (e *encoder) writeBlankLine() { } func (e *encoder) emitDoc(doc *tomlDoc, prefix []string) error { + // The walk that built the document checked the context on its own + // cadence; the emission of a large document is long enough to need the + // same checks, or a cancellation that lands after the walk would wait a + // full document out. + if err := e.checkCtx(); err != nil { + return err + } if e.opts.layout == LayoutKindGrouped { // Scalars first, then inline sub-tables as value lines, then the // remaining tables as headers, then arrays of tables. Each pass walks // the entries in place; grouping copies of them cost the encoder a // third of its allocations for nothing. for i := range doc.entries { + if i%ctxCheckInterval == 0 { + if err := e.checkCtx(); err != nil { + return err + } + } kv := &doc.entries[i] if kv.kind != entryScalar { continue @@ -1413,6 +1314,11 @@ func (e *encoder) emitDoc(doc *tomlDoc, prefix []string) error { // header of this document: a line written after a [header] would be // read back as part of that table. for i := range doc.entries { + if i%ctxCheckInterval == 0 { + if err := e.checkCtx(); err != nil { + return err + } + } t := &doc.entries[i] if t.kind != entryTable { continue @@ -1424,6 +1330,11 @@ func (e *encoder) emitDoc(doc *tomlDoc, prefix []string) error { t.emitted = inlined } for i := range doc.entries { + if i%ctxCheckInterval == 0 { + if err := e.checkCtx(); err != nil { + return err + } + } t := &doc.entries[i] if t.kind != entryTable || t.emitted { continue @@ -1441,12 +1352,22 @@ func (e *encoder) emitDoc(doc *tomlDoc, prefix []string) error { } } for i := range doc.entries { + if i%ctxCheckInterval == 0 { + if err := e.checkCtx(); err != nil { + return err + } + } a := &doc.entries[i] if a.kind != entryArray { continue } path := append(append([]string{}, prefix...), a.key) for j, sub := range a.docs { + if j%ctxCheckInterval == 0 { + if err := e.checkCtx(); err != nil { + return err + } + } e.writeBlankLine() if j == 0 { e.writeComments(a.comments) @@ -1469,7 +1390,12 @@ func (e *encoder) emitDoc(doc *tomlDoc, prefix []string) error { // own section content; the emitter still writes sub-documents as separate // nested blocks, so a "" sub-keyed scalar following a header for the same // section is impossible in practice (struct fields are visited in order). - for _, ent := range doc.entries { + for i, ent := range doc.entries { + if i%ctxCheckInterval == 0 { + if err := e.checkCtx(); err != nil { + return err + } + } switch ent.kind { case entryScalar: if err := e.writeKV(&ent); err != nil { @@ -1699,9 +1625,15 @@ func (e *encoder) writeValue(val any) error { case float64: return e.writeFloat(v) case time.Time: + if err := wholeMinuteOffset(v); err != nil { + return err + } e.buf.WriteString(offsetString(v)) return nil case OffsetDateTime: + if err := wholeMinuteOffset(v.Time); err != nil { + return err + } e.buf.WriteString(v.String()) return nil case LocalDateTime: @@ -1714,19 +1646,19 @@ func (e *encoder) writeValue(val any) error { e.buf.WriteString(v.String()) return nil case []any: - e.buf.WriteByte('[') - for i, item := range v { - if i > 0 { - e.buf.WriteString(", ") - } - if err := e.writeValue(item); err != nil { - return err - } + if err := e.enterValueDepth(); err != nil { + return err } - e.buf.WriteByte(']') - return nil + err := e.writeValueArray(v) + e.valueDepth-- + return err case map[string]any: - return e.writeInlineMap(v) + if err := e.enterValueDepth(); err != nil { + return err + } + err := e.writeInlineMap(v) + e.valueDepth-- + return err case nil: return fmt.Errorf("interpres: cannot encode nil value") default: @@ -1734,6 +1666,32 @@ func (e *encoder) writeValue(val any) error { } } +// enterValueDepth counts one level of boxed container nesting. The +// reflection walk has its own bound, but a tree built by hand and written +// through the boxed path carries no walk, so the writer bounds itself. +func (e *encoder) enterValueDepth() error { + e.valueDepth++ + if e.valueDepth > maxEncodeDepth { + return errDepthLimit() + } + return nil +} + +// writeValueArray writes a boxed value array, one writeValue per element. +func (e *encoder) writeValueArray(v []any) error { + e.buf.WriteByte('[') + for i, item := range v { + if i > 0 { + e.buf.WriteString(", ") + } + if err := e.writeValue(item); err != nil { + return err + } + } + e.buf.WriteByte(']') + return nil +} + // writeInlineMap renders m as a TOML inline table, on one line when it fits // there and across lines when it does not. func (e *encoder) writeInlineMap(m map[string]any) error { diff --git a/encode_test.go b/encode_test.go index 756ba07..6f5c55d 100644 --- a/encode_test.go +++ b/encode_test.go @@ -67,7 +67,7 @@ func TestMarshalFloatSpecials(t *testing.T) { } } -func TestMarshalFloatNormalizesNegativeZero(t *testing.T) { +func TestMarshalFloatNormalisesNegativeZero(t *testing.T) { // The output contract normalises negative zero to "0.0". type Cfg struct { Z float64 `toml:"z"` @@ -2284,3 +2284,115 @@ func TestEmitFieldComments(t *testing.T) { } }) } + +// errWriter fails every write with a fixed error. +type errWriter struct{ err error } + +func (w errWriter) Write([]byte) (int, error) { return 0, w.err } + +// TestMarshalWrite covers the streaming entry: the happy path with options +// and a failing writer. +func TestMarshalWrite(t *testing.T) { + var buf bytes.Buffer + err := MarshalWrite(&buf, map[string]any{"b": 2, "a": 1}) + if err != nil { + t.Fatalf("MarshalWrite: %v", err) + } + // A map carries no order, so the writer uses the sorted one. + if buf.String() != "a = 1\nb = 2\n" { + t.Errorf("output = %q", buf.String()) + } + writeErr := errors.New("boom") + if err := MarshalWrite(errWriter{writeErr}, map[string]any{"a": 1}); !errors.Is(err, writeErr) { + t.Errorf("err = %v, want the write error wrapped", err) + } +} + +// TestMarshalRejectsUnsupportedKinds pins the clear error a field of a kind +// TOML cannot carry raises, through the struct walk. +func TestMarshalRejectsUnsupportedKinds(t *testing.T) { + tests := []struct { + name string + value any + }{ + {"func", struct { + F func() `toml:"f"` + }{}}, + {"chan", struct { + C chan int `toml:"c"` + }{}}, + {"complex", struct { + Z complex128 `toml:"z"` + }{}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + _, err := Marshal(tt.value) + if err == nil { + t.Fatalf("Marshal accepted %#v", tt.value) + } + if !strings.Contains(err.Error(), "cannot encode") { + t.Errorf("err = %v, want the cannot-encode complaint", err) + } + }) + } +} + +// TestMarshalOmitsEmptyPointerTableSlice pins that an empty slice of pointer +// tables is omitted, the rule its non-pointer form already follows. +func TestMarshalOmitsEmptyPointerTableSlice(t *testing.T) { + type item struct { + N int `toml:"n"` + } + out, err := Marshal(struct { + Items []*item `toml:"items"` + }{}) + if err != nil { + t.Fatalf("Marshal: %v", err) + } + if len(out) != 0 { + t.Errorf("output = %q, want the empty array of tables omitted", out) + } +} + +// TestMarshalRejectsNonWholeMinuteOffset pins that a zone offset carrying +// seconds is refused instead of silently losing them. +func TestMarshalRejectsNonWholeMinuteOffset(t *testing.T) { + z := time.FixedZone("", 57*60+44) + _, err := Marshal(struct { + Stamp time.Time `toml:"stamp"` + }{Stamp: time.Date(1890, 1, 1, 12, 0, 0, 0, z)}) + if err == nil || !strings.Contains(err.Error(), "not a whole number of minutes") { + t.Errorf("err = %v, want the whole-minute offset complaint", err) + } + _, err = Marshal(struct { + Stamp OffsetDateTime `toml:"stamp"` + }{Stamp: OffsetDateTime{time.Date(1890, 1, 1, 12, 0, 0, 0, z)}}) + if err == nil || !strings.Contains(err.Error(), "not a whole number of minutes") { + t.Errorf("err = %v, want the whole-minute offset complaint for the wrapper", err) + } +} + +// cancelOnMarshal cancels the context the encode runs under, the moment its +// method is called, so the emission that follows is already past the walk's +// own checks. +type cancelOnMarshal struct { + cancel context.CancelFunc +} + +func (c cancelOnMarshal) MarshalTOML() (any, error) { + c.cancel() + return int64(1), nil +} + +// TestMarshalContextCancelsDuringEmission pins that a context cancelled +// between the walk and the emission stops the encode instead of writing the +// whole document out. +func TestMarshalContextCancelsDuringEmission(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + value := map[string]any{"k": cancelOnMarshal{cancel}} + if _, err := MarshalContext(ctx, value); !errors.Is(err, context.Canceled) { + t.Errorf("err = %v, want the cancellation", err) + } +}