decode.go 5.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213
  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. // UnmarshalBE unmarshalls OSCAR protocol messages in big-endian format.
  12. func UnmarshalBE(v any, r io.Reader) error {
  13. if err := unmarshal(reflect.TypeOf(v).Elem(), reflect.ValueOf(v).Elem(), "", r, binary.BigEndian); err != nil {
  14. return fmt.Errorf("%w: %w", ErrUnmarshalFailure, err)
  15. }
  16. return nil
  17. }
  18. // UnmarshalLE unmarshalls OSCAR protocol messages in little-endian format.
  19. func UnmarshalLE(v any, r io.Reader) error {
  20. if err := unmarshal(reflect.TypeOf(v).Elem(), reflect.ValueOf(v).Elem(), "", r, binary.LittleEndian); err != nil {
  21. return fmt.Errorf("%w: %w", ErrUnmarshalFailure, err)
  22. }
  23. return nil
  24. }
  25. // MarshalLE marshals ICQ protocol messages in little-endian format.
  26. func unmarshal(t reflect.Type, v reflect.Value, tag reflect.StructTag, r io.Reader, order binary.ByteOrder) error {
  27. oscTag, err := parseOSCARTag(tag)
  28. if err != nil {
  29. return fmt.Errorf("error parsing tag: %w", err)
  30. }
  31. if oscTag.optional {
  32. v.Set(reflect.New(t.Elem()))
  33. err := unmarshalStruct(t.Elem(), v.Elem(), oscTag, r, order)
  34. if errors.Is(err, io.EOF) {
  35. // no values to read, but that's ok since this struct is optional
  36. v.Set(reflect.Zero(t))
  37. err = nil
  38. }
  39. return err
  40. } else if v.Kind() == reflect.Ptr {
  41. return errNonOptionalPointer
  42. }
  43. switch v.Kind() {
  44. case reflect.Slice:
  45. return unmarshalSlice(v, oscTag, r, order)
  46. case reflect.String:
  47. return unmarshalString(v, oscTag, r, order)
  48. case reflect.Struct:
  49. return unmarshalStruct(t, v, oscTag, r, order)
  50. case reflect.Uint8:
  51. var l uint8
  52. if err := binary.Read(r, order, &l); err != nil {
  53. return err
  54. }
  55. v.Set(reflect.ValueOf(l))
  56. return nil
  57. case reflect.Uint16:
  58. var l uint16
  59. if err := binary.Read(r, order, &l); err != nil {
  60. return err
  61. }
  62. v.Set(reflect.ValueOf(l))
  63. return nil
  64. case reflect.Uint32:
  65. var l uint32
  66. if err := binary.Read(r, order, &l); err != nil {
  67. return err
  68. }
  69. v.Set(reflect.ValueOf(l))
  70. return nil
  71. case reflect.Uint64:
  72. var l uint64
  73. if err := binary.Read(r, order, &l); err != nil {
  74. return err
  75. }
  76. v.Set(reflect.ValueOf(l))
  77. return nil
  78. default:
  79. return fmt.Errorf("unsupported type %v", t.Kind())
  80. }
  81. }
  82. func unmarshalSlice(v reflect.Value, oscTag oscarTag, r io.Reader, order binary.ByteOrder) error {
  83. slice := reflect.New(v.Type()).Elem()
  84. elemType := v.Type().Elem()
  85. if oscTag.hasLenPrefix {
  86. bufLen, err := unmarshalUnsignedInt(oscTag.lenPrefix, r, order)
  87. if err != nil {
  88. return err
  89. }
  90. b := make([]byte, bufLen)
  91. if bufLen > 0 {
  92. if _, err := io.ReadFull(r, b); err != nil {
  93. return err
  94. }
  95. }
  96. buf := bytes.NewBuffer(b)
  97. for buf.Len() > 0 {
  98. elem := reflect.New(elemType).Elem()
  99. if err := unmarshal(elemType, elem, "", buf, order); err != nil {
  100. return err
  101. }
  102. slice = reflect.Append(slice, elem)
  103. }
  104. } else if oscTag.hasCountPrefix {
  105. count, err := unmarshalUnsignedInt(oscTag.countPrefix, r, order)
  106. if err != nil {
  107. return err
  108. }
  109. for i := 0; i < count; i++ {
  110. elem := reflect.New(elemType).Elem()
  111. if err := unmarshal(elemType, elem, "", r, order); err != nil {
  112. return err
  113. }
  114. slice = reflect.Append(slice, elem)
  115. }
  116. } else {
  117. for {
  118. elem := reflect.New(elemType).Elem()
  119. if err := unmarshal(elemType, elem, "", r, order); err != nil {
  120. if errors.Is(err, io.EOF) {
  121. break
  122. }
  123. return err
  124. }
  125. slice = reflect.Append(slice, elem)
  126. }
  127. }
  128. v.Set(slice)
  129. return nil
  130. }
  131. func unmarshalString(v reflect.Value, oscTag oscarTag, r io.Reader, order binary.ByteOrder) error {
  132. if !oscTag.hasLenPrefix {
  133. return fmt.Errorf("missing len_prefix tag")
  134. }
  135. bufLen, err := unmarshalUnsignedInt(oscTag.lenPrefix, r, order)
  136. if err != nil {
  137. return err
  138. }
  139. buf := make([]byte, bufLen)
  140. if bufLen > 0 {
  141. if _, err := io.ReadFull(r, buf); err != nil {
  142. return err
  143. }
  144. }
  145. // todo is there a more efficient way?
  146. v.SetString(string(buf))
  147. return nil
  148. }
  149. func unmarshalStruct(t reflect.Type, v reflect.Value, oscTag oscarTag, r io.Reader, order binary.ByteOrder) error {
  150. if oscTag.hasLenPrefix {
  151. bufLen, err := unmarshalUnsignedInt(oscTag.lenPrefix, r, order)
  152. if err != nil {
  153. return err
  154. }
  155. b := make([]byte, bufLen)
  156. if bufLen > 0 {
  157. if _, err := io.ReadFull(r, b); err != nil {
  158. return err
  159. }
  160. }
  161. r = bytes.NewBuffer(b)
  162. }
  163. for i := 0; i < v.NumField(); i++ {
  164. field := t.Field(i)
  165. value := v.Field(i)
  166. if field.Type.Kind() == reflect.Ptr {
  167. if i != v.NumField()-1 {
  168. return fmt.Errorf("pointer type found at non-final field %s", field.Name)
  169. }
  170. if field.Type.Elem().Kind() != reflect.Struct {
  171. return fmt.Errorf("%w: field %s must point to a struct, got %v instead",
  172. errNonOptionalPointer, field.Name, field.Type.Elem().Kind())
  173. }
  174. }
  175. if err := unmarshal(field.Type, value, field.Tag, r, order); err != nil {
  176. return err
  177. }
  178. }
  179. return nil
  180. }
  181. func unmarshalUnsignedInt(intType reflect.Kind, r io.Reader, order binary.ByteOrder) (int, error) {
  182. var bufLen int
  183. switch intType {
  184. case reflect.Uint8:
  185. var l uint8
  186. if err := binary.Read(r, order, &l); err != nil {
  187. return 0, err
  188. }
  189. bufLen = int(l)
  190. case reflect.Uint16:
  191. var l uint16
  192. if err := binary.Read(r, order, &l); err != nil {
  193. return 0, err
  194. }
  195. bufLen = int(l)
  196. default:
  197. panic(fmt.Sprintf("unsupported type %s. allowed types: uint8, uint16", intType))
  198. }
  199. return bufLen, nil
  200. }