decode.go 4.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171
  1. package wire
  2. import (
  3. "bytes"
  4. "encoding/binary"
  5. "errors"
  6. "fmt"
  7. "io"
  8. "reflect"
  9. )
  10. var ErrUnmarshalFailure = errors.New("failed to unmarshal")
  11. func Unmarshal(v any, r io.Reader) error {
  12. return unmarshal(reflect.TypeOf(v).Elem(), reflect.ValueOf(v).Elem(), "", r)
  13. }
  14. func unmarshal(t reflect.Type, v reflect.Value, tag reflect.StructTag, r io.Reader) error {
  15. switch v.Kind() {
  16. case reflect.Struct:
  17. for i := 0; i < v.NumField(); i++ {
  18. if err := unmarshal(t.Field(i).Type, v.Field(i), t.Field(i).Tag, r); err != nil {
  19. return err
  20. }
  21. }
  22. return nil
  23. case reflect.String:
  24. var bufLen int
  25. if lenTag, ok := tag.Lookup("len_prefix"); ok {
  26. switch lenTag {
  27. case "uint8":
  28. var l uint8
  29. if err := binary.Read(r, binary.BigEndian, &l); err != nil {
  30. return err
  31. }
  32. bufLen = int(l)
  33. case "uint16":
  34. var l uint16
  35. if err := binary.Read(r, binary.BigEndian, &l); err != nil {
  36. return err
  37. }
  38. bufLen = int(l)
  39. default:
  40. return fmt.Errorf("%w: unsupported len_prefix type %s. allowed types: uint8, uint16", ErrUnmarshalFailure, lenTag)
  41. }
  42. } else {
  43. return fmt.Errorf("%w: missing len_prefix tag", ErrUnmarshalFailure)
  44. }
  45. buf := make([]byte, bufLen)
  46. if bufLen > 0 {
  47. if _, err := io.ReadFull(r, buf); err != nil {
  48. return err
  49. }
  50. }
  51. // todo is there a more efficient way?
  52. v.SetString(string(buf))
  53. return nil
  54. case reflect.Uint8:
  55. var l uint8
  56. if err := binary.Read(r, binary.BigEndian, &l); err != nil {
  57. return err
  58. }
  59. v.Set(reflect.ValueOf(l))
  60. return nil
  61. case reflect.Uint16:
  62. var l uint16
  63. if err := binary.Read(r, binary.BigEndian, &l); err != nil {
  64. return err
  65. }
  66. v.Set(reflect.ValueOf(l))
  67. return nil
  68. case reflect.Uint32:
  69. var l uint32
  70. if err := binary.Read(r, binary.BigEndian, &l); err != nil {
  71. return err
  72. }
  73. v.Set(reflect.ValueOf(l))
  74. return nil
  75. case reflect.Uint64:
  76. var l uint64
  77. if err := binary.Read(r, binary.BigEndian, &l); err != nil {
  78. return err
  79. }
  80. v.Set(reflect.ValueOf(l))
  81. return nil
  82. case reflect.Slice:
  83. if lenTag, ok := tag.Lookup("len_prefix"); ok {
  84. var bufLen int
  85. switch lenTag {
  86. case "uint8":
  87. var l uint8
  88. if err := binary.Read(r, binary.BigEndian, &l); err != nil {
  89. return err
  90. }
  91. bufLen = int(l)
  92. case "uint16":
  93. var l uint16
  94. if err := binary.Read(r, binary.BigEndian, &l); err != nil {
  95. return err
  96. }
  97. bufLen = int(l)
  98. default:
  99. return fmt.Errorf("%w: unsupported len_prefix type %s. allowed types: uint8, uint16", ErrUnmarshalFailure, lenTag)
  100. }
  101. buf := make([]byte, bufLen)
  102. if bufLen > 0 {
  103. if _, err := io.ReadFull(r, buf); err != nil {
  104. return err
  105. }
  106. }
  107. b := bytes.NewBuffer(buf)
  108. slice := reflect.New(v.Type()).Elem()
  109. // todo: if this is a slice of scalars, there should be no need to
  110. // call Unmarshal on each element. it should be possible to just
  111. // call binary.Read(r, binary.BigEndian, []byte)
  112. for b.Len() > 0 {
  113. v1 := reflect.New(v.Type().Elem()).Interface()
  114. if err := Unmarshal(v1, b); err != nil {
  115. return err
  116. }
  117. slice = reflect.Append(slice, reflect.ValueOf(v1).Elem())
  118. }
  119. v.Set(slice)
  120. } else if countTag, ok := tag.Lookup("count_prefix"); ok {
  121. var count int
  122. switch countTag {
  123. case "uint8":
  124. var l uint8
  125. if err := binary.Read(r, binary.BigEndian, &l); err != nil {
  126. return err
  127. }
  128. count = int(l)
  129. case "uint16":
  130. var l uint16
  131. if err := binary.Read(r, binary.BigEndian, &l); err != nil {
  132. return err
  133. }
  134. count = int(l)
  135. default:
  136. return fmt.Errorf("%w: unsupported count_prefix type %s. allowed types: uint8, uint16", ErrUnmarshalFailure, lenTag)
  137. }
  138. slice := reflect.New(v.Type()).Elem()
  139. for i := 0; i < count; i++ {
  140. v1 := reflect.New(v.Type().Elem()).Interface()
  141. if err := Unmarshal(v1, r); err != nil {
  142. return err
  143. }
  144. slice = reflect.Append(slice, reflect.ValueOf(v1).Elem())
  145. }
  146. v.Set(slice)
  147. } else {
  148. slice := reflect.New(v.Type()).Elem()
  149. for {
  150. v1 := reflect.New(v.Type().Elem()).Interface()
  151. if err := Unmarshal(v1, r); err != nil {
  152. if err == io.EOF {
  153. break
  154. }
  155. return err
  156. }
  157. slice = reflect.Append(slice, reflect.ValueOf(v1).Elem())
  158. }
  159. v.Set(slice)
  160. }
  161. return nil
  162. default:
  163. return fmt.Errorf("%w: unsupported type %v", ErrUnmarshalFailure, t.Kind())
  164. }
  165. }