You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
DateKeys/internal/cbortest/cbortest.go

253 lines
6.2 KiB

// Package cbortest encodes and decodes generic CBOR values for tests.
//
// It is written independently of package codec, on purpose: tests use it to
// build inputs that codec cannot write, such as null, negative integers or a
// non-canonical head inside an otherwise valid object, and as a second
// reading of the CBOR profile of spec §58 to check codec against.
//
// It is internal and exists for tests only.
package cbortest
import (
"encoding/binary"
"errors"
"fmt"
"slices"
"unicode/utf8"
)
// Raw is an encoded data item that Marshal writes verbatim.
type Raw []byte
// Pairs is a map whose entries Marshal writes in the given order, keys and
// values alternating. Its keys may be of any type Marshal accepts, so it
// builds maps with keys out of order, repeated or not unsigned integers.
type Pairs []any
// Marshal returns the deterministic encoding of v, which is one of:
// uint64, uint, int or int64 (a negative value as major type 1), string,
// []byte, bool, nil (null), []any, map[uint64]any (keys in ascending order),
// Pairs or Raw. Strings are written as they are, even if they are not valid
// UTF-8.
func Marshal(v any) ([]byte, error) { return appendValue(nil, v) }
func appendHead(b []byte, major byte, arg uint64) []byte {
m := major << 5
switch {
case arg < 24:
return append(b, m|byte(arg))
case arg < 1<<8:
return append(b, m|24, byte(arg))
case arg < 1<<16:
return binary.BigEndian.AppendUint16(append(b, m|25), uint16(arg))
case arg < 1<<32:
return binary.BigEndian.AppendUint32(append(b, m|26), uint32(arg))
}
return binary.BigEndian.AppendUint64(append(b, m|27), arg)
}
func appendInt(b []byte, v int64) []byte {
if v < 0 {
return appendHead(b, 1, uint64(-(v + 1)))
}
return appendHead(b, 0, uint64(v))
}
func appendValue(b []byte, v any) ([]byte, error) {
switch v := v.(type) {
case nil:
return append(b, 0xf6), nil
case bool:
if v {
return append(b, 0xf5), nil
}
return append(b, 0xf4), nil
case uint64:
return appendHead(b, 0, v), nil
case uint:
return appendHead(b, 0, uint64(v)), nil
case int:
return appendInt(b, int64(v)), nil
case int64:
return appendInt(b, v), nil
case []byte:
return append(appendHead(b, 2, uint64(len(v))), v...), nil
case string:
return append(appendHead(b, 3, uint64(len(v))), v...), nil
case Raw:
return append(b, v...), nil
case []any:
b = appendHead(b, 4, uint64(len(v)))
return appendAll(b, v)
case Pairs:
if len(v)%2 != 0 {
return nil, fmt.Errorf("cbortest: Pairs of odd length %d", len(v))
}
b = appendHead(b, 5, uint64(len(v)/2))
return appendAll(b, v)
case map[uint64]any:
keys := make([]uint64, 0, len(v))
for k := range v {
keys = append(keys, k)
}
slices.Sort(keys)
b = appendHead(b, 5, uint64(len(v)))
for _, k := range keys {
b = appendHead(b, 0, k)
var err error
if b, err = appendValue(b, v[k]); err != nil {
return nil, err
}
}
return b, nil
}
return nil, fmt.Errorf("cbortest: cannot encode %T", v)
}
func appendAll(b []byte, vs []any) ([]byte, error) {
for _, x := range vs {
var err error
if b, err = appendValue(b, x); err != nil {
return nil, err
}
}
return b, nil
}
// MaxDepth bounds the nesting that Unmarshal follows.
const MaxDepth = 1000
// Unmarshal decodes exactly one data item of the CBOR profile of spec §58
// into uint64, []byte, string, []any and map[uint64]any. It rejects every
// other major type, indefinite lengths, heads not in their shortest form,
// map keys that are not unsigned integers in strictly ascending order,
// invalid UTF-8, truncation, trailing bytes and nesting deeper than MaxDepth.
func Unmarshal(b []byte) (any, error) {
r := &reader{b: b}
v, err := r.value(0)
if err != nil {
return nil, err
}
if r.off != len(b) {
return nil, fmt.Errorf("cbortest: %d trailing bytes", len(b)-r.off)
}
return v, nil
}
// UnmarshalMap is Unmarshal for an input that must hold a map.
func UnmarshalMap(b []byte) (map[uint64]any, error) {
v, err := Unmarshal(b)
if err != nil {
return nil, err
}
m, ok := v.(map[uint64]any)
if !ok {
return nil, fmt.Errorf("cbortest: %T, not a map", v)
}
return m, nil
}
var errTruncated = errors.New("cbortest: truncated")
type reader struct {
b []byte
off int
}
func (r *reader) head() (major byte, arg uint64, err error) {
if r.off >= len(r.b) {
return 0, 0, errTruncated
}
ib := r.b[r.off]
r.off++
major, info := ib>>5, ib&0x1f
if major == 1 || major >= 6 {
return 0, 0, fmt.Errorf("cbortest: major type %d", major)
}
if info < 24 {
return major, uint64(info), nil
}
if info > 27 {
return 0, 0, fmt.Errorf("cbortest: additional information %d", info)
}
n := 1 << (info - 24)
if len(r.b)-r.off < n {
return 0, 0, errTruncated
}
for _, c := range r.b[r.off : r.off+n] {
arg = arg<<8 | uint64(c)
}
r.off += n
if (n == 1 && arg < 24) || (n > 1 && arg>>(4*n) == 0) {
return 0, 0, fmt.Errorf("cbortest: %d not in its shortest form", arg)
}
return major, arg, nil
}
func (r *reader) bytes(n uint64) ([]byte, error) {
if n > uint64(len(r.b)-r.off) {
return nil, errTruncated
}
s := r.b[r.off : r.off+int(n)]
r.off += int(n)
return append([]byte{}, s...), nil
}
func (r *reader) value(depth int) (any, error) {
major, arg, err := r.head()
if err != nil {
return nil, err
}
switch major {
case 0:
return arg, nil
case 2:
return r.bytes(arg)
case 3:
s, err := r.bytes(arg)
if err != nil {
return nil, err
}
if !utf8.Valid(s) {
return nil, errors.New("cbortest: invalid UTF-8")
}
return string(s), nil
}
if depth >= MaxDepth {
return nil, errors.New("cbortest: nested too deep")
}
if arg > uint64(len(r.b)-r.off) {
return nil, errTruncated
}
if major == 4 {
out := make([]any, 0, arg)
for range arg {
v, err := r.value(depth + 1)
if err != nil {
return nil, err
}
out = append(out, v)
}
return out, nil
}
out := make(map[uint64]any, arg)
var last uint64
for i := range arg {
km, k, err := r.head()
if err != nil {
return nil, err
}
if km != 0 {
return nil, fmt.Errorf("cbortest: map key of major type %d", km)
}
if i > 0 && k <= last {
return nil, fmt.Errorf("cbortest: map key %d after %d", k, last)
}
last = k
if out[k], err = r.value(depth + 1); err != nil {
return nil, err
}
}
return out, nil
}

Powered by TurnKey Linux.