4
0

encode.go 5.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213
  1. package wire
  2. import (
  3. "bytes"
  4. "encoding/binary"
  5. "errors"
  6. "fmt"
  7. "io"
  8. "reflect"
  9. "strings"
  10. )
  11. var (
  12. ErrMarshalFailure = errors.New("failed to marshal")
  13. errMarshalFailureNilSNAC = errors.New("attempting to marshal a nil SNAC")
  14. errNonOptionalPointer = errors.New("pointer fields must reference structs and have an `optional` struct tag")
  15. errOptionalNonPointer = errors.New("optional fields must be pointers")
  16. errInvalidStructTag = errors.New("invalid struct tag")
  17. )
  18. func Marshal(v any, w io.Writer) error {
  19. if err := marshal(reflect.TypeOf(v), reflect.ValueOf(v), "", w); err != nil {
  20. return fmt.Errorf("%w: %w", ErrMarshalFailure, err)
  21. }
  22. return nil
  23. }
  24. func marshal(t reflect.Type, v reflect.Value, tag reflect.StructTag, w io.Writer) error {
  25. if t == nil {
  26. return errMarshalFailureNilSNAC
  27. }
  28. oscTag, err := parseOSCARTag(tag)
  29. if err != nil {
  30. return err
  31. }
  32. if oscTag.optional {
  33. if t.Kind() != reflect.Ptr {
  34. return fmt.Errorf("%w: got %v", errOptionalNonPointer, t.Kind())
  35. }
  36. if v.IsNil() {
  37. return nil // nil value
  38. }
  39. // dereference pointer
  40. return marshalStruct(t.Elem(), v.Elem(), oscTag, w)
  41. } else if t.Kind() == reflect.Ptr {
  42. return errNonOptionalPointer
  43. }
  44. switch t.Kind() {
  45. case reflect.Slice:
  46. return marshalSlice(t, v, oscTag, w)
  47. case reflect.String:
  48. return marshalString(oscTag, v, w)
  49. case reflect.Struct:
  50. return marshalStruct(t, v, oscTag, w)
  51. case reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
  52. return binary.Write(w, binary.BigEndian, v.Interface())
  53. default:
  54. return fmt.Errorf("unsupported type %v", t.Kind())
  55. }
  56. }
  57. func marshalSlice(t reflect.Type, v reflect.Value, oscTag oscarTag, w io.Writer) error {
  58. // todo: only write to temporary buffer if len_prefix is set
  59. buf := &bytes.Buffer{}
  60. if t.Elem().Kind() == reflect.Struct {
  61. for j := 0; j < v.Len(); j++ {
  62. if err := marshalStruct(t.Elem(), v.Index(j), oscarTag{}, buf); err != nil {
  63. return err
  64. }
  65. }
  66. } else {
  67. if err := binary.Write(buf, binary.BigEndian, v.Interface()); err != nil {
  68. return fmt.Errorf("error marshalling %s", t.Elem().Kind())
  69. }
  70. }
  71. if oscTag.hasLenPrefix {
  72. if err := marshalUnsignedInt(oscTag.lenPrefix, buf.Len(), w); err != nil {
  73. return err
  74. }
  75. } else if oscTag.hasCountPrefix {
  76. if err := marshalUnsignedInt(oscTag.countPrefix, v.Len(), w); err != nil {
  77. return err
  78. }
  79. }
  80. if buf.Len() > 0 {
  81. _, err := w.Write(buf.Bytes())
  82. return err
  83. }
  84. return nil
  85. }
  86. func marshalString(oscTag oscarTag, v reflect.Value, w io.Writer) error {
  87. if oscTag.hasLenPrefix {
  88. if err := marshalUnsignedInt(oscTag.lenPrefix, len(v.String()), w); err != nil {
  89. return err
  90. }
  91. }
  92. return binary.Write(w, binary.BigEndian, []byte(v.String()))
  93. }
  94. func marshalStruct(t reflect.Type, v reflect.Value, oscTag oscarTag, w io.Writer) error {
  95. marshalEachField := func(w io.Writer) error {
  96. for i := 0; i < t.NumField(); i++ {
  97. field := t.Field(i)
  98. value := v.Field(i)
  99. if field.Type.Kind() == reflect.Ptr {
  100. if i != t.NumField()-1 {
  101. return fmt.Errorf("pointer type found at non-final field %s", field.Name)
  102. }
  103. if field.Type.Elem().Kind() != reflect.Struct {
  104. return fmt.Errorf("field %s must point to a struct, got %v instead", field.Name,
  105. field.Type.Elem().Kind())
  106. }
  107. }
  108. if err := marshal(field.Type, value, field.Tag, w); err != nil {
  109. return err
  110. }
  111. }
  112. return nil
  113. }
  114. if oscTag.hasLenPrefix {
  115. buf := &bytes.Buffer{}
  116. if err := marshalEachField(buf); err != nil {
  117. return err
  118. }
  119. // write struct length
  120. if err := marshalUnsignedInt(oscTag.lenPrefix, buf.Len(), w); err != nil {
  121. return err
  122. }
  123. // write struct bytes
  124. if buf.Len() > 0 {
  125. _, err := w.Write(buf.Bytes())
  126. return err
  127. }
  128. return nil
  129. }
  130. return marshalEachField(w)
  131. }
  132. func marshalUnsignedInt(intType reflect.Kind, intVal int, w io.Writer) error {
  133. switch intType {
  134. case reflect.Uint8:
  135. if err := binary.Write(w, binary.BigEndian, uint8(intVal)); err != nil {
  136. return err
  137. }
  138. case reflect.Uint16:
  139. if err := binary.Write(w, binary.BigEndian, uint16(intVal)); err != nil {
  140. return err
  141. }
  142. default:
  143. panic(fmt.Sprintf("unsupported type %s. allowed types: uint8, uint16", intType))
  144. }
  145. return nil
  146. }
  147. type oscarTag struct {
  148. hasCountPrefix bool
  149. countPrefix reflect.Kind
  150. hasLenPrefix bool
  151. lenPrefix reflect.Kind
  152. optional bool
  153. }
  154. func parseOSCARTag(tag reflect.StructTag) (oscarTag, error) {
  155. var oscTag oscarTag
  156. val, ok := tag.Lookup("oscar")
  157. if !ok {
  158. return oscTag, nil
  159. }
  160. for _, kv := range strings.Split(val, ",") {
  161. kvSplit := strings.SplitN(kv, "=", 2)
  162. if len(kvSplit) == 2 {
  163. switch kvSplit[0] {
  164. case "len_prefix":
  165. oscTag.hasLenPrefix = true
  166. switch kvSplit[1] {
  167. case "uint8":
  168. oscTag.lenPrefix = reflect.Uint8
  169. case "uint16":
  170. oscTag.lenPrefix = reflect.Uint16
  171. default:
  172. return oscTag, fmt.Errorf("%w: unsupported type %s. allowed types: uint8, uint16",
  173. errInvalidStructTag, kvSplit[1])
  174. }
  175. case "count_prefix":
  176. oscTag.hasCountPrefix = true
  177. switch kvSplit[1] {
  178. case "uint8":
  179. oscTag.countPrefix = reflect.Uint8
  180. case "uint16":
  181. oscTag.countPrefix = reflect.Uint16
  182. default:
  183. return oscTag, fmt.Errorf("%w: unsupported type %s. allowed types: uint8, uint16",
  184. errInvalidStructTag, kvSplit[1])
  185. }
  186. }
  187. } else {
  188. oscTag.optional = kvSplit[0] == "optional"
  189. }
  190. }
  191. var err error
  192. if oscTag.hasCountPrefix && oscTag.hasLenPrefix {
  193. err = fmt.Errorf("%w: struct elem has both len_prefix and count_prefix", errInvalidStructTag)
  194. }
  195. return oscTag, err
  196. }