decode.go 5.6 KB

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