encode.go 3.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114
  1. package wire
  2. import (
  3. "bytes"
  4. "encoding/binary"
  5. "errors"
  6. "fmt"
  7. "io"
  8. "reflect"
  9. )
  10. var ErrMarshalFailure = errors.New("failed to marshal")
  11. var ErrMarshalFailureNilSNAC = errors.New("attempting to marshal a nil SNAC")
  12. func Marshal(v any, w io.Writer) error {
  13. return marshal(reflect.TypeOf(v), reflect.ValueOf(v), "", w)
  14. }
  15. func marshal(t reflect.Type, v reflect.Value, tag reflect.StructTag, w io.Writer) error {
  16. if t == nil {
  17. return ErrMarshalFailureNilSNAC
  18. }
  19. switch t.Kind() {
  20. case reflect.Struct:
  21. marshalEachField := func(w io.Writer) error {
  22. for i := 0; i < t.NumField(); i++ {
  23. if err := marshal(t.Field(i).Type, v.Field(i), t.Field(i).Tag, w); err != nil {
  24. return err
  25. }
  26. }
  27. return nil
  28. }
  29. if lenTag, ok := tag.Lookup("len_prefix"); ok {
  30. buf := &bytes.Buffer{}
  31. if err := marshalEachField(buf); err != nil {
  32. return err
  33. }
  34. // write struct length
  35. if err := writeUnsignedInt(lenTag, buf.Len(), w); err != nil {
  36. return err
  37. }
  38. // write struct bytes
  39. if buf.Len() > 0 {
  40. _, err := w.Write(buf.Bytes())
  41. return err
  42. }
  43. return nil
  44. }
  45. return marshalEachField(w)
  46. case reflect.String:
  47. if lenTag, ok := tag.Lookup("len_prefix"); ok {
  48. if err := writeUnsignedInt(lenTag, len(v.String()), w); err != nil {
  49. return err
  50. }
  51. }
  52. return binary.Write(w, binary.BigEndian, []byte(v.String()))
  53. case reflect.Slice:
  54. // todo: only write to temporary buffer if len_prefix is set
  55. buf := &bytes.Buffer{}
  56. if t.Elem().Kind() == reflect.Struct {
  57. for j := 0; j < v.Len(); j++ {
  58. element := v.Index(j)
  59. if err := Marshal(element.Interface(), buf); err != nil {
  60. return err
  61. }
  62. }
  63. } else {
  64. if err := binary.Write(buf, binary.BigEndian, v.Interface()); err != nil {
  65. return fmt.Errorf("%w: error marshalling %s", ErrMarshalFailure, t.Elem().Kind())
  66. }
  67. }
  68. var hasLenPrefix bool
  69. if l, ok := tag.Lookup("len_prefix"); ok {
  70. hasLenPrefix = true
  71. if err := writeUnsignedInt(l, buf.Len(), w); err != nil {
  72. return err
  73. }
  74. }
  75. if l, ok := tag.Lookup("count_prefix"); ok {
  76. if hasLenPrefix {
  77. return fmt.Errorf("%w: struct elem has both len_prefix and count_prefix: ", ErrMarshalFailure)
  78. }
  79. if err := writeUnsignedInt(l, v.Len(), w); err != nil {
  80. return err
  81. }
  82. }
  83. if buf.Len() > 0 {
  84. _, err := w.Write(buf.Bytes())
  85. return err
  86. }
  87. return nil
  88. case reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
  89. return binary.Write(w, binary.BigEndian, v.Interface())
  90. default:
  91. return fmt.Errorf("%w: unsupported type %v", ErrMarshalFailure, t.Kind())
  92. }
  93. }
  94. func writeUnsignedInt(intType string, intVal int, w io.Writer) error {
  95. switch intType {
  96. case "uint8":
  97. if err := binary.Write(w, binary.BigEndian, uint8(intVal)); err != nil {
  98. return err
  99. }
  100. case "uint16":
  101. if err := binary.Write(w, binary.BigEndian, uint16(intVal)); err != nil {
  102. return err
  103. }
  104. default:
  105. return fmt.Errorf("%w: unsupported type %s. allowed types: uint8, uint16", ErrMarshalFailure, intType)
  106. }
  107. return nil
  108. }