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 // CheckWrite, the rule of encoders of spec §72: what the Registry knows goes // only where it is registered, with valid data; the rest is not checked. func TestCheckWrite(t *testing.T) { ok := []extension.Extension{ext(t, "org.a", 1, []byte("ok"))} if err := extension.CheckWrite(placed{}, extension.Control, extension.Critical, ok); err != nil { t.Errorf("where it is registered: %v", err) } if err := extension.CheckWrite(placed{}, extension.PublicHeader, extension.Critical, ok); err == nil || !strings.Contains(err.Error(), "not registered") { t.Errorf("where it is not registered: %v", err) } if err := extension.CheckWrite(placed{}, extension.Control, extension.Critical, []extension.Extension{ext(t, "org.a", 1, []byte("ko"))}); err == nil { t.Error("invalid data was written") } other := []extension.Extension{ext(t, "org.z", 1, []byte("anything"))} if extension.CheckWrite(placed{}, extension.PublicHeader, extension.Critical, other) != nil || extension.CheckWrite(nil, extension.PublicHeader, extension.Critical, ok) != nil { t.Error("an extension that the Registry does not know was checked") } var std extension.Standard note := []extension.Extension{{ID: extension.NoteID, Version: 1, Data: []byte("Cartas")}} if extension.CheckWrite(std, extension.PublicHeader, extension.Noncritical, note) != nil || extension.CheckWrite(std, extension.AccessKey, extension.Noncritical, note) == nil || extension.CheckWrite(std, extension.PublicHeader, extension.Noncritical, []extension.Extension{{ID: extension.NoteID, Version: 1, Data: []byte(" a")}}) == nil { t.Error("datekeys.note") } } // 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) } }) }