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/extension/extension_test.go

479 lines
18 KiB

package extension_test
import (
"bytes"
"encoding/hex"
"errors"
"fmt"
"slices"
"strings"
"testing"
"time"
datekeys "g.activething.com/go/DateKeys"
"g.activething.com/go/DateKeys/codec"
"g.activething.com/go/DateKeys/extension"
"g.activething.com/go/DateKeys/internal/cbortest"
)
func ext(t *testing.T, id string, v uint64, data []byte) extension.Extension {
t.Helper()
if data == nil {
return extension.Extension{ID: id, Version: v}
}
e, err := extension.New(id, v, data)
if err != nil {
t.Fatal(err)
}
return e
}
func TestNew(t *testing.T) {
data := []byte("public label")
e, err := extension.New("org.example.label", 1, data)
if err != nil || e.ID != "org.example.label" || e.Version != 1 || !bytes.Equal(e.Data, data) {
t.Fatalf("%+v %v", e, err)
}
data[0] = 'P'
if e.Data[0] != 'p' {
t.Fatal("New does not copy data")
}
for name, tc := range map[string]struct {
id string
version uint64
data []byte
}{
// Spec §54, §58.1: an extension without data omits key 2; it is built
// as a literal, never through New.
"nil data": {"org.a", 1, nil},
"empty data": {"org.a", 1, []byte{}},
"empty id": {"", 1, []byte{1}},
"invalid UTF-8 id": {"org.\xff", 1, []byte{1}},
"id too long": {strings.Repeat("a", extension.MaxIDLen+1), 1, []byte{1}},
"version above 2^32": {"org.a", extension.MaxVersion + 1, []byte{1}},
} {
if _, err := extension.New(tc.id, tc.version, tc.data); !errors.Is(err, datekeys.ErrNonCanonicalCBOR) {
t.Errorf("%s: %v", name, err)
}
}
if _, err := extension.New("org.a", extension.MaxVersion, []byte{0}); err != nil {
t.Fatalf("version 2^32-1 rejected: %v", err)
}
}
// decodeArray decodes b, one extension array, as the containing objects do:
// DecodeArray and then the re-encoding check.
func decodeArray(b []byte) ([]extension.Extension, error) {
var exts []extension.Extension
decode := func(d *codec.Decoder) (err error) { exts, err = extension.DecodeArray(d); return err }
encode := func(e *codec.Encoder) { extension.EncodeArray(e, exts) }
if err := codec.Unmarshal(b, decode, encode); err != nil {
return nil, err
}
return exts, nil
}
func encodeArray(t *testing.T, exts []extension.Extension) []byte {
t.Helper()
var e codec.Encoder
extension.EncodeArray(&e, exts)
b, err := e.Out()
if err != nil {
t.Fatal(err)
}
return b
}
func TestCanonicalSorts(t *testing.T) {
in := []extension.Extension{ext(t, "org.b", 1, nil), ext(t, "org.a", 2, []byte("x")), ext(t, "Z", 9, nil), ext(t, "org.aa", 1, []byte{7})}
w, err := extension.Canonical(in)
if err != nil {
t.Fatal(err)
}
var order []string
for _, e := range w {
order = append(order, e.ID)
}
// Bytewise UTF-8 order: uppercase before lowercase, prefixes first.
if got := []string{"Z", "org.a", "org.aa", "org.b"}; !slices.Equal(order, got) {
t.Fatalf("order %v, want %v", order, got)
}
if in[0].ID != "org.b" {
t.Fatal("Canonical reordered its input")
}
if w, _ := extension.Canonical(nil); w != nil {
t.Fatal("an empty array must be nil so that the key is omitted")
}
back, err := decodeArray(encodeArray(t, w))
if err != nil || len(back) != 4 || back[1].ID != "org.a" || string(back[1].Data) != "x" || back[0].Data != nil {
t.Fatalf("decode: %+v %v", back, err)
}
// The wire form: data is a byte string, and key 2 is omitted without data.
if got, want := hex.EncodeToString(encodeArray(t, w[:2])), "82"+"a200615a0109"+"a300656f72672e6101020241"+"78"; got != want {
t.Fatalf("wire %s, want %s", got, want)
}
}
func many(n int) []extension.Extension {
out := make([]extension.Extension, n)
for i := range out {
out[i] = extension.Extension{ID: fmt.Sprintf("org.example.%03d", i), Version: 1}
}
return out
}
func TestCanonicalRejects(t *testing.T) {
for name, in := range map[string][]extension.Extension{
"same id twice": {ext(t, "org.a", 1, nil), ext(t, "org.a", 2, nil)},
"empty id": {{ID: "", Version: 1}},
"invalid UTF-8 id": {{ID: "org.\xff", Version: 1}},
"present but empty": {{ID: "org.a", Version: 1, Data: []byte{}}},
"version above 2^32": {{ID: "org.a", Version: extension.MaxVersion + 1}},
"65 extensions (§64)": many(extension.MaxExtensions + 1),
"duplicate among many": append(many(3), extension.Extension{ID: "org.example.001", Version: 2}),
} {
if _, err := extension.Canonical(in); !errors.Is(err, datekeys.ErrNonCanonicalCBOR) {
t.Errorf("%s: %v", name, err)
}
}
if w, err := extension.Canonical(many(extension.MaxExtensions)); err != nil || len(w) != extension.MaxExtensions {
t.Fatalf("64 extensions rejected: %v", err)
}
}
// EncodeArray never writes an array that DecodeArray rejects: the Encoder
// records the error and Out returns it.
func TestEncodeArrayRejects(t *testing.T) {
for name, in := range map[string][]extension.Extension{
"nil": nil,
"empty": {},
"65 extensions": many(extension.MaxExtensions + 1),
"present but empty": {{ID: "org.a", Version: 1, Data: []byte{}}},
"out of order": {{ID: "org.b", Version: 1}, {ID: "org.a", Version: 1}},
"same id twice": {{ID: "org.a", Version: 1}, {ID: "org.a", Version: 2}},
"version above 2^32": {{ID: "org.a", Version: extension.MaxVersion + 1}},
"empty id": {{ID: "", Version: 1}},
"id too long": {{ID: strings.Repeat("a", extension.MaxIDLen+1), Version: 1}},
"invalid UTF-8 id": {{ID: "org.\xff", Version: 1}},
} {
var e codec.Encoder
extension.EncodeArray(&e, in)
if b, err := e.Out(); b != nil || !errors.Is(err, datekeys.ErrNonCanonicalCBOR) {
t.Errorf("%s: %x %v", name, b, err)
}
}
if b := encodeArray(t, many(extension.MaxExtensions)); b[0] != 0x98 || b[1] != extension.MaxExtensions {
t.Fatalf("64 extensions: %x", b[:2])
}
}
// entries returns the wire maps of exts, without data.
func entries(exts []extension.Extension) []any {
out := make([]any, len(exts))
for i, e := range exts {
out[i] = map[uint64]any{0: e.ID, 1: e.Version}
}
return out
}
func TestDecodeArrayRejects(t *testing.T) {
entry := func(id any, v uint64) map[uint64]any { return map[uint64]any{0: id, 1: v} }
for name, in := range map[string]any{
"out of order": []any{entry("org.b", 1), entry("org.a", 1)},
"repeated id": []any{entry("org.a", 1), entry("org.a", 2)},
"present but empty": []any{map[uint64]any{0: "org.a", 1: uint64(1), 2: []byte{}}},
"version above 2^32": []any{entry("org.a", extension.MaxVersion+1)},
"empty id": []any{entry("", 1)},
"id too long": []any{entry(strings.Repeat("a", extension.MaxIDLen+1), 1)},
"invalid UTF-8 id": []any{entry("org.\xff", 1)},
"id not a string": []any{entry(uint64(1), 1)},
"65 extensions": entries(many(extension.MaxExtensions + 1)),
"empty array": []any{},
"not an array": entry("org.a", 1),
"entry not a map": []any{"org.a"},
"four keys": []any{map[uint64]any{0: "org.a", 1: uint64(1), 2: []byte{1}, 3: uint64(0)}},
"unknown key": []any{map[uint64]any{0: "org.a", 1: uint64(1), 3: uint64(0)}},
"without version": []any{map[uint64]any{0: "org.a"}},
"without id": []any{map[uint64]any{1: uint64(1)}},
"text key": cbortest.Raw{0x81, 0xa2, 0x61, 0x61, 0x00, 0x01, 0x01},
"truncated": cbortest.Raw{0x82, 0xa2, 0x00, 0x61, 0x61, 0x01, 0x01},
} {
b, err := cbortest.Marshal(in)
if err != nil {
t.Fatal(err)
}
if _, err := decodeArray(b); !errors.Is(err, datekeys.ErrNonCanonicalCBOR) {
t.Errorf("%s: %v", name, err)
}
}
b, _ := cbortest.Marshal(entries(many(extension.MaxExtensions)))
if got, err := decodeArray(b); err != nil || len(got) != extension.MaxExtensions {
t.Fatalf("64 extensions rejected: %v", err)
}
}
// Spec §54: key 2 is absent or a non-empty byte string, whatever it contains.
// Every other form is rejected by the extension map itself.
func TestData(t *testing.T) {
const head = "81" + "a3006161" + "0101" + "02" // [{0: "a", 1: 1, 2: ...}]
valid := []struct{ name, item, data string }{
{"one byte", "4100", "00"},
{"bytes that are not CBOR", "44ff1c00f7", "ff1c00f7"},
{"CBOR that the base protocol never decodes", "49a2f97e0000f97e0001", "a2f97e0000f97e0001"},
{"14 nested arrays", "4f" + strings.Repeat("81", 14) + "00", strings.Repeat("81", 14) + "00"},
{"24 bytes, one-byte length", "5818" + strings.Repeat("ab", 24), strings.Repeat("ab", 24)},
}
for _, tc := range valid {
b, _ := hex.DecodeString(head + tc.item)
if w, err := decodeArray(b); err != nil || hex.EncodeToString(w[0].Data) != tc.data {
t.Errorf("%s: %v", tc.name, err)
}
}
if none, err := decodeArray([]byte{0x81, 0xa2, 0x00, 0x61, 0x61, 0x01, 0x01}); err != nil || none[0].Data != nil || none[0].ID != "a" {
t.Fatalf("extension without data: %+v %v", none, err)
}
for _, tc := range []struct{ name, item string }{
{"empty byte string h''", "40"},
{"text string", "6161"},
{"unsigned integer", "01"},
{"negative integer", "20"},
{"array", "820102"},
{"empty array", "80"},
{"map", "a10001"},
{"true", "f5"},
{"null", "f6"},
{"undefined", "f7"},
{"float", "f93c00"},
{"tag", "c24101"},
{"length not in shortest form", "5801" + "00"},
{"indefinite-length byte string", "5f4100ff"},
{"truncated byte string", "42" + "00"},
} {
b, _ := hex.DecodeString(head + tc.item)
if w, err := decodeArray(b); !errors.Is(err, datekeys.ErrNonCanonicalCBOR) {
t.Errorf("%s: %+v %v", tc.name, w, err)
}
}
}
func TestCheckDisjoint(t *testing.T) {
crit := []extension.Extension{ext(t, "org.a", 1, nil)}
non := []extension.Extension{ext(t, "org.a", 2, nil)}
if err := extension.CheckDisjoint(crit, non); !errors.Is(err, datekeys.ErrNonCanonicalCBOR) {
t.Fatalf("id in both arrays: %v", err)
}
a := []extension.Extension{{ID: "a"}, {ID: "c"}, {ID: "e"}}
b := []extension.Extension{{ID: "b"}, {ID: "d"}, {ID: "f"}}
if err := extension.CheckDisjoint(a, b); err != nil {
t.Fatal(err)
}
// Input that is not in canonical order is sorted before the merge.
unsorted := []extension.Extension{{ID: "z"}, {ID: "e"}}
if err := extension.CheckDisjoint(a, unsorted); !errors.Is(err, datekeys.ErrNonCanonicalCBOR) {
t.Fatalf("unsorted input: %v", err)
}
if !slices.IsSortedFunc(unsorted, func(x, y extension.Extension) int { return strings.Compare(y.ID, x.ID) }) {
t.Fatal("CheckDisjoint reordered its input")
}
if err := extension.CheckDisjoint(nil, b); err != nil {
t.Fatal(err)
}
}
// The merge is linear: 200 000 + 200 000 identifiers take milliseconds, where
// the former pairwise comparison took hours (spec §76, case 6).
func TestCheckDisjointIsLinear(t *testing.T) {
const n = 200_000
crit, non := make([]extension.Extension, n), make([]extension.Extension, n)
for i := range n {
crit[i].ID = fmt.Sprintf("a.%07d", i)
non[i].ID = fmt.Sprintf("b.%07d", i)
}
start := time.Now()
if err := extension.CheckDisjoint(crit, non); err != nil {
t.Fatal(err)
}
if d := time.Since(start); d > 5*time.Second {
t.Fatalf("CheckDisjoint took %s", d)
}
}
// validator knows org.a v1 and org.b v1, and accepts only the data "ok".
type validator struct{}
func (validator) Known(id string, v uint64) bool { return (id == "org.a" || id == "org.b") && v == 1 }
func (validator) ValidateData(e extension.Extension) error {
if string(e.Data) != "ok" {
// A validator may report a normative code of its own; the result
// carries ErrExtensionDataInvalid only.
return fmt.Errorf("want \"ok\": %w", datekeys.ErrNonCanonicalCBOR)
}
return nil
}
func TestCheckCritical(t *testing.T) {
crit := []extension.Extension{ext(t, "org.a", 1, nil)}
if err := extension.CheckCritical(crit, nil); !errors.Is(err, datekeys.ErrExtensionCriticalUnknown) {
t.Fatalf("unknown critical accepted by the base protocol: %v", err)
}
if err := extension.CheckCritical(crit, extension.Set{"org.a": {2}}); !errors.Is(err, datekeys.ErrExtensionCriticalUnknown) {
t.Fatalf("other version accepted: %v", err)
}
if err := extension.CheckCritical(crit, extension.Set{"org.a": {1}}); err != nil {
t.Fatalf("known critical rejected: %v", err)
}
if err := extension.CheckCritical(nil, nil); err != nil {
t.Fatal(err)
}
// Spec §54: a known critical extension with invalid data.
good, bad := ext(t, "org.a", 1, []byte("ok")), ext(t, "org.b", 1, []byte("ko"))
if err := extension.CheckCritical([]extension.Extension{good}, validator{}); err != nil {
t.Fatal(err)
}
err := extension.CheckCritical([]extension.Extension{good, bad}, validator{})
if datekeys.Code(err) != "ERR_EXTENSION_DATA_INVALID" || errors.Is(err, datekeys.ErrNonCanonicalCBOR) {
t.Fatalf("invalid data: %v", err)
}
// An unknown critical extension takes precedence over invalid data.
unknown := ext(t, "org.c", 1, nil)
if err := extension.CheckCritical([]extension.Extension{bad, unknown}, validator{}); !errors.Is(err, datekeys.ErrExtensionCriticalUnknown) {
t.Fatalf("unknown and invalid: %v", err)
}
}
func TestCheckNoncritical(t *testing.T) {
good, bad, unknown := ext(t, "org.a", 1, []byte("ok")), ext(t, "org.b", 1, nil), ext(t, "org.c", 1, []byte("ko"))
all := []extension.Extension{good, bad, unknown}
if u := extension.CheckNoncritical(all, nil); u != nil {
t.Fatalf("base protocol: %v", u)
}
if u := extension.CheckNoncritical(all, extension.Set{"org.b": {1}}); u != nil {
t.Fatalf("a Set does not validate data: %v", u)
}
u := extension.CheckNoncritical(all, validator{})
if len(u) != 1 || u[0].ID != "org.b" || u[0].Version != 1 || !errors.Is(u[0].Err, datekeys.ErrExtensionDataInvalid) {
t.Fatalf("unusable: %+v", u)
}
}
// placed is validator with the placement of its registrations (spec §72):
// org.a only in the critical_extensions of CONTROL_CBOR, org.b only in the
// noncritical_extensions of PUBLIC_HEADER and of the .dkk.
type placed struct{ validator }
func (placed) RegisteredIn(id string, v uint64, obj extension.Object, arr extension.Array) bool {
switch id {
case "org.a":
return obj == extension.Control && arr == extension.Critical
case "org.b":
return arr == extension.Noncritical && obj != extension.Control
}
return false
}
// Spec §54, §72: a known extension that appears in an object or array it is
// not registered for is treated there as unknown. A Registry that is not a
// Placement knows its extensions everywhere, and CheckCritical and
// CheckNoncritical, which do not know the object, consult no Placement.
func TestPlacement(t *testing.T) {
objects := []extension.Object{extension.PublicHeader, extension.Control, extension.AccessKey}
arrays := []extension.Array{extension.Critical, extension.Noncritical}
for _, obj := range objects {
for _, arr := range arrays {
if !extension.KnownIn(validator{}, "org.a", 1, obj, arr) || extension.KnownIn(validator{}, "org.c", 1, obj, arr) || extension.KnownIn(nil, "org.a", 1, obj, arr) {
t.Fatalf("without a Placement, in the %s of %s", arr, obj)
}
if got, want := extension.KnownIn(placed{}, "org.a", 1, obj, arr), obj == extension.Control && arr == extension.Critical; got != want {
t.Errorf("org.a in the %s of %s: known %v", arr, obj, got)
}
if extension.KnownIn(placed{}, "org.a", 2, obj, arr) {
t.Errorf("org.a v2, which the registry does not know, in the %s of %s", arr, obj)
}
}
}
// A critical extension outside its registration is unknown there, with
// the object in the message.
crit := []extension.Extension{ext(t, "org.a", 1, []byte("ok"))}
if err := extension.CheckCriticalIn(extension.Control, crit, placed{}); err != nil {
t.Fatalf("where it is registered: %v", err)
}
for _, obj := range []extension.Object{extension.PublicHeader, extension.AccessKey} {
err := extension.CheckCriticalIn(obj, crit, placed{})
if !errors.Is(err, datekeys.ErrExtensionCriticalUnknown) || !strings.Contains(err.Error(), obj.String()) {
t.Errorf("CONTROL_CBOR extension in the critical_extensions of %s: %v", obj, err)
}
if err := extension.CheckCriticalIn(obj, crit, validator{}); err != nil {
t.Errorf("a Registry that is not a Placement, in %s: %v", obj, err)
}
}
if err := extension.CheckCritical(crit, placed{}); err != nil {
t.Errorf("CheckCritical consulted the Placement: %v", err)
}
// Unknown there comes before the invalid data of another extension.
mixed := []extension.Extension{ext(t, "org.a", 1, []byte("ko")), ext(t, "org.b", 1, []byte("ok"))}
if err := extension.CheckCriticalIn(extension.Control, mixed, placed{}); !errors.Is(err, datekeys.ErrExtensionCriticalUnknown) {
t.Errorf("noncritical-only extension in a critical array, after invalid data: %v", err)
}
// A noncritical extension outside its registration is ignored, its data
// unchecked; where it is registered, invalid data makes it unusable.
non := []extension.Extension{ext(t, "org.b", 1, []byte("ko"))}
for _, obj := range objects {
u := extension.CheckNoncriticalIn(obj, non, placed{})
if obj == extension.Control && u != nil || obj != extension.Control && (len(u) != 1 || u[0].ID != "org.b") {
t.Errorf("org.b with invalid data in the noncritical_extensions of %s: %+v", obj, u)
}
if u := extension.CheckNoncriticalIn(obj, non, validator{}); len(u) != 1 {
t.Errorf("a Registry that is not a Placement, in %s: %+v", obj, u)
}
}
if u := extension.CheckNoncritical(non, placed{}); len(u) != 1 {
t.Errorf("CheckNoncritical consulted the Placement: %+v", u)
}
for v, want := range map[fmt.Stringer]string{
extension.PublicHeader: "PUBLIC_HEADER", extension.Control: "CONTROL_CBOR", extension.AccessKey: ".dkk", extension.Object(0): "Object(0)",
extension.Critical: "critical_extensions", extension.Noncritical: "noncritical_extensions", extension.Array(3): "Array(3)",
} {
if v.String() != want {
t.Errorf("%s, want %s", v, want)
}
}
}
// FuzzDecodeArray: whatever DecodeArray accepts is in canonical order and
// re-encodes to its input with EncodeArray, and every error carries
// ErrNonCanonicalCBOR.
func FuzzDecodeArray(f *testing.F) {
for _, h := range []string{
"81a2006161" + "0101",
"82a2006161" + "0101" + "a3006162" + "0102" + "02" + "4100",
"81a3006161010102" + "40",
"80",
} {
b, _ := hex.DecodeString(h)
f.Add(b)
}
same := func(a, b extension.Extension) bool {
return a.ID == b.ID && a.Version == b.Version && bytes.Equal(a.Data, b.Data) && (a.Data == nil) == (b.Data == nil)
}
f.Fuzz(func(t *testing.T, b []byte) {
exts, err := decodeArray(b)
if err != nil {
if !errors.Is(err, datekeys.ErrNonCanonicalCBOR) {
t.Fatalf("error without ErrNonCanonicalCBOR: %v", err)
}
return
}
if c, err := extension.Canonical(exts); err != nil || !slices.EqualFunc(c, exts, same) {
t.Fatalf("decoded array is not canonical: %v", err)
}
if re := encodeArray(t, exts); !bytes.Equal(re, b) {
t.Fatalf("%x re-encodes to %x", b, re)
}
})
}

Powered by TurnKey Linux.