encode.go 7.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290
  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.Array:
  54. return marshalArray(t, v, w, order)
  55. case reflect.Slice:
  56. return marshalSlice(t, v, oscTag, w, order)
  57. case reflect.String:
  58. return marshalString(oscTag, v, w, order)
  59. case reflect.Struct:
  60. return marshalStruct(t, v, oscTag, w, order)
  61. case reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
  62. return binary.Write(w, order, v.Interface())
  63. case reflect.Interface:
  64. return marshalInterface(v, w, oscTag, order)
  65. default:
  66. return fmt.Errorf("unsupported type %v", t.Kind())
  67. }
  68. }
  69. func marshalInterface(v reflect.Value, w io.Writer, tag oscarTag, order binary.ByteOrder) error {
  70. elem := v.Elem()
  71. if elem.Kind() != reflect.Struct {
  72. return fmt.Errorf("interface underlying type must be a struct, got %v instead", elem.Kind())
  73. }
  74. return marshalStruct(elem.Type(), elem, tag, w, order)
  75. }
  76. func marshalArray(t reflect.Type, v reflect.Value, w io.Writer, order binary.ByteOrder) error {
  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{}, w, order); err != nil {
  80. return fmt.Errorf("error marshalling %s: %w", t.Elem().Kind(), err)
  81. }
  82. }
  83. } else {
  84. if err := binary.Write(w, order, v.Interface()); err != nil {
  85. return fmt.Errorf("error marshalling %s: %w", t.Elem().Kind(), err)
  86. }
  87. }
  88. return nil
  89. }
  90. func marshalSlice(t reflect.Type, v reflect.Value, oscTag oscarTag, w io.Writer, order binary.ByteOrder) error {
  91. // todo: only write to temporary buffer if len_prefix is set
  92. buf := &bytes.Buffer{}
  93. if t.Elem().Kind() == reflect.Struct {
  94. for j := 0; j < v.Len(); j++ {
  95. if err := marshalStruct(t.Elem(), v.Index(j), oscarTag{}, buf, order); err != nil {
  96. return err
  97. }
  98. }
  99. } else {
  100. if err := binary.Write(buf, order, v.Interface()); err != nil {
  101. return fmt.Errorf("error marshalling %s", t.Elem().Kind())
  102. }
  103. }
  104. if oscTag.hasLenPrefix {
  105. if err := marshalUnsignedInt(oscTag.lenPrefix, buf.Len(), w, order); err != nil {
  106. return err
  107. }
  108. } else if oscTag.hasCountPrefix {
  109. if err := marshalUnsignedInt(oscTag.countPrefix, v.Len(), w, order); err != nil {
  110. return err
  111. }
  112. }
  113. if buf.Len() > 0 {
  114. _, err := w.Write(buf.Bytes())
  115. return err
  116. }
  117. return nil
  118. }
  119. func marshalString(oscTag oscarTag, v reflect.Value, w io.Writer, order binary.ByteOrder) error {
  120. str := v.String()
  121. if oscTag.nullTerminated && str != "" {
  122. str = str + "\x00"
  123. }
  124. if oscTag.hasLenPrefix {
  125. if err := marshalUnsignedInt(oscTag.lenPrefix, len(str), w, order); err != nil {
  126. return err
  127. }
  128. }
  129. if str == "" {
  130. return nil
  131. }
  132. return binary.Write(w, order, []byte(str))
  133. }
  134. func marshalStruct(t reflect.Type, v reflect.Value, oscTag oscarTag, w io.Writer, order binary.ByteOrder) error {
  135. // marshal ICQ messages in little endian order
  136. if t.Name() == "ICQMessageReplyEnvelope" {
  137. order = binary.LittleEndian
  138. }
  139. marshalEachField := func(w io.Writer) error {
  140. for i := 0; i < t.NumField(); i++ {
  141. field := t.Field(i)
  142. value := v.Field(i)
  143. if field.Type.Kind() == reflect.Ptr {
  144. if i != t.NumField()-1 {
  145. return fmt.Errorf("pointer type found at non-final field %s", field.Name)
  146. }
  147. if field.Type.Elem().Kind() != reflect.Struct {
  148. return fmt.Errorf("field %s must point to a struct, got %v instead", field.Name,
  149. field.Type.Elem().Kind())
  150. }
  151. }
  152. if err := marshal(field.Type, value, field.Tag, w, order); err != nil {
  153. return err
  154. }
  155. }
  156. return nil
  157. }
  158. if oscTag.hasLenPrefix {
  159. buf := &bytes.Buffer{}
  160. if err := marshalEachField(buf); err != nil {
  161. return err
  162. }
  163. // write struct length
  164. if err := marshalUnsignedInt(oscTag.lenPrefix, buf.Len(), w, order); err != nil {
  165. return err
  166. }
  167. // write struct bytes
  168. if buf.Len() > 0 {
  169. _, err := w.Write(buf.Bytes())
  170. return err
  171. }
  172. return nil
  173. }
  174. return marshalEachField(w)
  175. }
  176. func marshalUnsignedInt(intType reflect.Kind, intVal int, w io.Writer, order binary.ByteOrder) error {
  177. switch intType {
  178. case reflect.Uint8:
  179. if err := binary.Write(w, order, uint8(intVal)); err != nil {
  180. return err
  181. }
  182. case reflect.Uint16:
  183. if err := binary.Write(w, order, uint16(intVal)); err != nil {
  184. return err
  185. }
  186. default:
  187. panic(fmt.Sprintf("unsupported type %s. allowed types: uint8, uint16", intType))
  188. }
  189. return nil
  190. }
  191. type oscarTag struct {
  192. hasCountPrefix bool
  193. countPrefix reflect.Kind
  194. hasLenPrefix bool
  195. lenPrefix reflect.Kind
  196. optional bool
  197. nullTerminated bool
  198. // quirk names an unmarshal-only decode exception
  199. quirk string
  200. }
  201. func parseOSCARTag(tag reflect.StructTag) (oscarTag, error) {
  202. var oscTag oscarTag
  203. val, ok := tag.Lookup("oscar")
  204. if !ok {
  205. return oscTag, nil
  206. }
  207. for _, kv := range strings.Split(val, ",") {
  208. kv = strings.TrimSpace(kv)
  209. if kv == "" {
  210. continue
  211. }
  212. kvSplit := strings.SplitN(kv, "=", 2)
  213. key := strings.TrimSpace(kvSplit[0])
  214. if len(kvSplit) == 2 {
  215. valPart := strings.TrimSpace(kvSplit[1])
  216. switch key {
  217. case "len_prefix":
  218. oscTag.hasLenPrefix = true
  219. switch valPart {
  220. case "uint8":
  221. oscTag.lenPrefix = reflect.Uint8
  222. case "uint16":
  223. oscTag.lenPrefix = reflect.Uint16
  224. default:
  225. return oscTag, fmt.Errorf("%w: unsupported type %s. allowed types: uint8, uint16",
  226. errInvalidStructTag, valPart)
  227. }
  228. case "count_prefix":
  229. oscTag.hasCountPrefix = true
  230. switch valPart {
  231. case "uint8":
  232. oscTag.countPrefix = reflect.Uint8
  233. case "uint16":
  234. oscTag.countPrefix = reflect.Uint16
  235. default:
  236. return oscTag, fmt.Errorf("%w: unsupported type %s. allowed types: uint8, uint16",
  237. errInvalidStructTag, valPart)
  238. }
  239. case "quirk":
  240. if valPart == "" {
  241. return oscTag, fmt.Errorf("%w: empty quirk value", errInvalidStructTag)
  242. }
  243. if oscTag.quirk != "" {
  244. return oscTag, fmt.Errorf("%w: duplicate quirk key", errInvalidStructTag)
  245. }
  246. oscTag.quirk = valPart
  247. default:
  248. return oscTag, fmt.Errorf("%w: unsupported struct tag %s",
  249. errInvalidStructTag, key)
  250. }
  251. } else {
  252. switch key {
  253. case "optional":
  254. oscTag.optional = true
  255. case "nullterm":
  256. oscTag.nullTerminated = true
  257. default:
  258. return oscTag, fmt.Errorf("%w: unsupported struct tag %s",
  259. errInvalidStructTag, key)
  260. }
  261. }
  262. }
  263. var err error
  264. if oscTag.hasCountPrefix && oscTag.hasLenPrefix {
  265. err = fmt.Errorf("%w: struct elem has both len_prefix and count_prefix", errInvalidStructTag)
  266. }
  267. return oscTag, err
  268. }