decode.go 5.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240
  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.Array:
  48. return unmarshalArray(v, r, order)
  49. case reflect.Slice:
  50. return unmarshalSlice(v, oscTag, r, order)
  51. case reflect.String:
  52. return unmarshalString(v, oscTag, r, order)
  53. case reflect.Struct:
  54. return unmarshalStruct(t, v, oscTag, r, order)
  55. case reflect.Uint8:
  56. var l uint8
  57. if err := binary.Read(r, order, &l); err != nil {
  58. return err
  59. }
  60. v.Set(reflect.ValueOf(l))
  61. return nil
  62. case reflect.Uint16:
  63. var l uint16
  64. if err := binary.Read(r, order, &l); err != nil {
  65. return err
  66. }
  67. v.Set(reflect.ValueOf(l))
  68. return nil
  69. case reflect.Uint32:
  70. var l uint32
  71. if err := binary.Read(r, order, &l); err != nil {
  72. return err
  73. }
  74. v.Set(reflect.ValueOf(l))
  75. return nil
  76. case reflect.Uint64:
  77. var l uint64
  78. if err := binary.Read(r, order, &l); err != nil {
  79. return err
  80. }
  81. v.Set(reflect.ValueOf(l))
  82. return nil
  83. default:
  84. return fmt.Errorf("unsupported type %v", t.Kind())
  85. }
  86. }
  87. func unmarshalArray(v reflect.Value, r io.Reader, order binary.ByteOrder) error {
  88. arrLen := v.Len()
  89. arrType := v.Type().Elem()
  90. for i := 0; i < arrLen; i++ {
  91. elem := reflect.New(arrType).Elem()
  92. if err := unmarshal(arrType, elem, "", r, order); err != nil {
  93. return err
  94. }
  95. v.Index(i).Set(elem)
  96. }
  97. return nil
  98. }
  99. func unmarshalSlice(v reflect.Value, oscTag oscarTag, r io.Reader, order binary.ByteOrder) error {
  100. slice := reflect.New(v.Type()).Elem()
  101. elemType := v.Type().Elem()
  102. if oscTag.hasLenPrefix {
  103. bufLen, err := unmarshalUnsignedInt(oscTag.lenPrefix, r, order)
  104. if err != nil {
  105. return err
  106. }
  107. b := make([]byte, bufLen)
  108. if bufLen > 0 {
  109. if _, err := io.ReadFull(r, b); err != nil {
  110. return err
  111. }
  112. }
  113. buf := bytes.NewBuffer(b)
  114. for buf.Len() > 0 {
  115. elem := reflect.New(elemType).Elem()
  116. if err := unmarshal(elemType, elem, "", buf, order); err != nil {
  117. return err
  118. }
  119. slice = reflect.Append(slice, elem)
  120. }
  121. } else if oscTag.hasCountPrefix {
  122. count, err := unmarshalUnsignedInt(oscTag.countPrefix, r, order)
  123. if err != nil {
  124. return err
  125. }
  126. for i := 0; i < count; i++ {
  127. elem := reflect.New(elemType).Elem()
  128. if err := unmarshal(elemType, elem, "", r, order); err != nil {
  129. return err
  130. }
  131. slice = reflect.Append(slice, elem)
  132. }
  133. } else {
  134. for {
  135. elem := reflect.New(elemType).Elem()
  136. if err := unmarshal(elemType, elem, "", r, order); err != nil {
  137. if errors.Is(err, io.EOF) {
  138. break
  139. }
  140. return err
  141. }
  142. slice = reflect.Append(slice, elem)
  143. }
  144. }
  145. v.Set(slice)
  146. return nil
  147. }
  148. func unmarshalString(v reflect.Value, oscTag oscarTag, r io.Reader, order binary.ByteOrder) error {
  149. if !oscTag.hasLenPrefix {
  150. return fmt.Errorf("missing len_prefix tag")
  151. }
  152. bufLen, err := unmarshalUnsignedInt(oscTag.lenPrefix, r, order)
  153. if err != nil {
  154. return err
  155. }
  156. buf := make([]byte, bufLen)
  157. if bufLen > 0 {
  158. if _, err := io.ReadFull(r, buf); err != nil {
  159. return err
  160. }
  161. if oscTag.nullTerminated {
  162. if buf[len(buf)-1] != 0x00 {
  163. return errNotNullTerminated
  164. }
  165. buf = buf[0 : len(buf)-1] // remove null terminator
  166. }
  167. }
  168. // todo is there a more efficient way?
  169. v.SetString(string(buf))
  170. return nil
  171. }
  172. func unmarshalStruct(t reflect.Type, v reflect.Value, oscTag oscarTag, r io.Reader, order binary.ByteOrder) error {
  173. if oscTag.hasLenPrefix {
  174. bufLen, err := unmarshalUnsignedInt(oscTag.lenPrefix, r, order)
  175. if err != nil {
  176. return err
  177. }
  178. b := make([]byte, bufLen)
  179. if bufLen > 0 {
  180. if _, err := io.ReadFull(r, b); err != nil {
  181. return err
  182. }
  183. }
  184. r = bytes.NewBuffer(b)
  185. }
  186. for i := 0; i < v.NumField(); i++ {
  187. field := t.Field(i)
  188. value := v.Field(i)
  189. if field.Type.Kind() == reflect.Ptr {
  190. if i != v.NumField()-1 {
  191. return fmt.Errorf("pointer type found at non-final field %s", field.Name)
  192. }
  193. if field.Type.Elem().Kind() != reflect.Struct {
  194. return fmt.Errorf("%w: field %s must point to a struct, got %v instead",
  195. errNonOptionalPointer, field.Name, field.Type.Elem().Kind())
  196. }
  197. }
  198. if err := unmarshal(field.Type, value, field.Tag, r, order); err != nil {
  199. return err
  200. }
  201. }
  202. return nil
  203. }
  204. func unmarshalUnsignedInt(intType reflect.Kind, r io.Reader, order binary.ByteOrder) (int, error) {
  205. var bufLen int
  206. switch intType {
  207. case reflect.Uint8:
  208. var l uint8
  209. if err := binary.Read(r, order, &l); err != nil {
  210. return 0, err
  211. }
  212. bufLen = int(l)
  213. case reflect.Uint16:
  214. var l uint16
  215. if err := binary.Read(r, order, &l); err != nil {
  216. return 0, err
  217. }
  218. bufLen = int(l)
  219. default:
  220. panic(fmt.Sprintf("unsupported type %s. allowed types: uint8, uint16", intType))
  221. }
  222. return bufLen, nil
  223. }