encode.go 6.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254
  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. // MarshalBE marshals OSCAR protocol messages in big-endian format.
  19. func MarshalBE(v any, w io.Writer) error {
  20. if err := marshal(reflect.TypeOf(v), reflect.ValueOf(v), "", w, binary.BigEndian); err != nil {
  21. return fmt.Errorf("%w: %w", ErrMarshalFailure, err)
  22. }
  23. return nil
  24. }
  25. // MarshalLE marshals ICQ protocol messages in little-endian format.
  26. func MarshalLE(v any, w io.Writer) error {
  27. if err := marshal(reflect.TypeOf(v), reflect.ValueOf(v), "", w, binary.LittleEndian); err != nil {
  28. return fmt.Errorf("%w: %w", ErrMarshalFailure, err)
  29. }
  30. return nil
  31. }
  32. func marshal(t reflect.Type, v reflect.Value, tag reflect.StructTag, w io.Writer, order binary.ByteOrder) error {
  33. if t == nil {
  34. return errMarshalFailureNilSNAC
  35. }
  36. oscTag, err := parseOSCARTag(tag)
  37. if err != nil {
  38. return err
  39. }
  40. if oscTag.optional {
  41. if t.Kind() != reflect.Ptr {
  42. return fmt.Errorf("%w: got %v", errOptionalNonPointer, t.Kind())
  43. }
  44. if v.IsNil() {
  45. return nil // nil value
  46. }
  47. // dereference pointer
  48. return marshalStruct(t.Elem(), v.Elem(), oscTag, w, order)
  49. } else if t.Kind() == reflect.Ptr {
  50. return errNonOptionalPointer
  51. }
  52. switch t.Kind() {
  53. case reflect.Slice:
  54. return marshalSlice(t, v, oscTag, w, order)
  55. case reflect.String:
  56. return marshalString(oscTag, v, w, order)
  57. case reflect.Struct:
  58. return marshalStruct(t, v, oscTag, w, order)
  59. case reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
  60. return binary.Write(w, order, v.Interface())
  61. case reflect.Interface:
  62. return marshalInterface(v, w, oscTag, order)
  63. default:
  64. return fmt.Errorf("unsupported type %v", t.Kind())
  65. }
  66. }
  67. func marshalInterface(v reflect.Value, w io.Writer, tag oscarTag, order binary.ByteOrder) error {
  68. elem := v.Elem()
  69. if elem.Kind() != reflect.Struct {
  70. return fmt.Errorf("interface underlying type must be a struct, got %v instead", elem.Kind())
  71. }
  72. return marshalStruct(elem.Type(), elem, tag, w, order)
  73. }
  74. func marshalSlice(t reflect.Type, v reflect.Value, oscTag oscarTag, w io.Writer, order binary.ByteOrder) error {
  75. // todo: only write to temporary buffer if len_prefix is set
  76. buf := &bytes.Buffer{}
  77. if t.Elem().Kind() == reflect.Struct {
  78. for j := 0; j < v.Len(); j++ {
  79. if err := marshalStruct(t.Elem(), v.Index(j), oscarTag{}, buf, order); err != nil {
  80. return err
  81. }
  82. }
  83. } else {
  84. if err := binary.Write(buf, order, v.Interface()); err != nil {
  85. return fmt.Errorf("error marshalling %s", t.Elem().Kind())
  86. }
  87. }
  88. if oscTag.hasLenPrefix {
  89. if err := marshalUnsignedInt(oscTag.lenPrefix, buf.Len(), w, order); err != nil {
  90. return err
  91. }
  92. } else if oscTag.hasCountPrefix {
  93. if err := marshalUnsignedInt(oscTag.countPrefix, v.Len(), w, order); err != nil {
  94. return err
  95. }
  96. }
  97. if buf.Len() > 0 {
  98. _, err := w.Write(buf.Bytes())
  99. return err
  100. }
  101. return nil
  102. }
  103. func marshalString(oscTag oscarTag, v reflect.Value, w io.Writer, order binary.ByteOrder) error {
  104. str := v.String()
  105. if oscTag.nullTerminated && str != "" {
  106. str = str + "\x00"
  107. }
  108. if oscTag.hasLenPrefix {
  109. if err := marshalUnsignedInt(oscTag.lenPrefix, len(str), w, order); err != nil {
  110. return err
  111. }
  112. }
  113. if str == "" {
  114. return nil
  115. }
  116. return binary.Write(w, order, []byte(str))
  117. }
  118. func marshalStruct(t reflect.Type, v reflect.Value, oscTag oscarTag, w io.Writer, order binary.ByteOrder) error {
  119. // marshal ICQ messages in little endian order
  120. if t.Name() == "ICQMessageReplyEnvelope" {
  121. order = binary.LittleEndian
  122. }
  123. marshalEachField := func(w io.Writer) error {
  124. for i := 0; i < t.NumField(); i++ {
  125. field := t.Field(i)
  126. value := v.Field(i)
  127. if field.Type.Kind() == reflect.Ptr {
  128. if i != t.NumField()-1 {
  129. return fmt.Errorf("pointer type found at non-final field %s", field.Name)
  130. }
  131. if field.Type.Elem().Kind() != reflect.Struct {
  132. return fmt.Errorf("field %s must point to a struct, got %v instead", field.Name,
  133. field.Type.Elem().Kind())
  134. }
  135. }
  136. if err := marshal(field.Type, value, field.Tag, w, order); err != nil {
  137. return err
  138. }
  139. }
  140. return nil
  141. }
  142. if oscTag.hasLenPrefix {
  143. buf := &bytes.Buffer{}
  144. if err := marshalEachField(buf); err != nil {
  145. return err
  146. }
  147. // write struct length
  148. if err := marshalUnsignedInt(oscTag.lenPrefix, buf.Len(), w, order); err != nil {
  149. return err
  150. }
  151. // write struct bytes
  152. if buf.Len() > 0 {
  153. _, err := w.Write(buf.Bytes())
  154. return err
  155. }
  156. return nil
  157. }
  158. return marshalEachField(w)
  159. }
  160. func marshalUnsignedInt(intType reflect.Kind, intVal int, w io.Writer, order binary.ByteOrder) error {
  161. switch intType {
  162. case reflect.Uint8:
  163. if err := binary.Write(w, order, uint8(intVal)); err != nil {
  164. return err
  165. }
  166. case reflect.Uint16:
  167. if err := binary.Write(w, order, uint16(intVal)); err != nil {
  168. return err
  169. }
  170. default:
  171. panic(fmt.Sprintf("unsupported type %s. allowed types: uint8, uint16", intType))
  172. }
  173. return nil
  174. }
  175. type oscarTag struct {
  176. hasCountPrefix bool
  177. countPrefix reflect.Kind
  178. hasLenPrefix bool
  179. lenPrefix reflect.Kind
  180. optional bool
  181. nullTerminated bool
  182. }
  183. func parseOSCARTag(tag reflect.StructTag) (oscarTag, error) {
  184. var oscTag oscarTag
  185. val, ok := tag.Lookup("oscar")
  186. if !ok {
  187. return oscTag, nil
  188. }
  189. for _, kv := range strings.Split(val, ",") {
  190. kvSplit := strings.SplitN(kv, "=", 2)
  191. if len(kvSplit) == 2 {
  192. switch kvSplit[0] {
  193. case "len_prefix":
  194. oscTag.hasLenPrefix = true
  195. switch kvSplit[1] {
  196. case "uint8":
  197. oscTag.lenPrefix = reflect.Uint8
  198. case "uint16":
  199. oscTag.lenPrefix = reflect.Uint16
  200. default:
  201. return oscTag, fmt.Errorf("%w: unsupported type %s. allowed types: uint8, uint16",
  202. errInvalidStructTag, kvSplit[1])
  203. }
  204. case "count_prefix":
  205. oscTag.hasCountPrefix = true
  206. switch kvSplit[1] {
  207. case "uint8":
  208. oscTag.countPrefix = reflect.Uint8
  209. case "uint16":
  210. oscTag.countPrefix = reflect.Uint16
  211. default:
  212. return oscTag, fmt.Errorf("%w: unsupported type %s. allowed types: uint8, uint16",
  213. errInvalidStructTag, kvSplit[1])
  214. }
  215. }
  216. } else {
  217. switch kvSplit[0] {
  218. case "optional":
  219. oscTag.optional = true
  220. case "nullterm":
  221. oscTag.nullTerminated = true
  222. default:
  223. return oscTag, fmt.Errorf("%w: unsupported struct tag %s",
  224. errInvalidStructTag, kvSplit[0])
  225. }
  226. }
  227. }
  228. var err error
  229. if oscTag.hasCountPrefix && oscTag.hasLenPrefix {
  230. err = fmt.Errorf("%w: struct elem has both len_prefix and count_prefix", errInvalidStructTag)
  231. }
  232. return oscTag, err
  233. }