encode.go 6.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235
  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. default:
  62. return fmt.Errorf("unsupported type %v", t.Kind())
  63. }
  64. }
  65. func marshalSlice(t reflect.Type, v reflect.Value, oscTag oscarTag, w io.Writer, order binary.ByteOrder) error {
  66. // todo: only write to temporary buffer if len_prefix is set
  67. buf := &bytes.Buffer{}
  68. if t.Elem().Kind() == reflect.Struct {
  69. for j := 0; j < v.Len(); j++ {
  70. if err := marshalStruct(t.Elem(), v.Index(j), oscarTag{}, buf, order); err != nil {
  71. return err
  72. }
  73. }
  74. } else {
  75. if err := binary.Write(buf, order, v.Interface()); err != nil {
  76. return fmt.Errorf("error marshalling %s", t.Elem().Kind())
  77. }
  78. }
  79. if oscTag.hasLenPrefix {
  80. if err := marshalUnsignedInt(oscTag.lenPrefix, buf.Len(), w, order); err != nil {
  81. return err
  82. }
  83. } else if oscTag.hasCountPrefix {
  84. if err := marshalUnsignedInt(oscTag.countPrefix, v.Len(), w, order); err != nil {
  85. return err
  86. }
  87. }
  88. if buf.Len() > 0 {
  89. _, err := w.Write(buf.Bytes())
  90. return err
  91. }
  92. return nil
  93. }
  94. func marshalString(oscTag oscarTag, v reflect.Value, w io.Writer, order binary.ByteOrder) error {
  95. str := v.String()
  96. if oscTag.nullTerminated {
  97. str = str + "\x00"
  98. }
  99. if oscTag.hasLenPrefix {
  100. if err := marshalUnsignedInt(oscTag.lenPrefix, len(str), w, order); err != nil {
  101. return err
  102. }
  103. }
  104. return binary.Write(w, order, []byte(str))
  105. }
  106. func marshalStruct(t reflect.Type, v reflect.Value, oscTag oscarTag, w io.Writer, order binary.ByteOrder) error {
  107. marshalEachField := func(w io.Writer) error {
  108. for i := 0; i < t.NumField(); i++ {
  109. field := t.Field(i)
  110. value := v.Field(i)
  111. if field.Type.Kind() == reflect.Ptr {
  112. if i != t.NumField()-1 {
  113. return fmt.Errorf("pointer type found at non-final field %s", field.Name)
  114. }
  115. if field.Type.Elem().Kind() != reflect.Struct {
  116. return fmt.Errorf("field %s must point to a struct, got %v instead", field.Name,
  117. field.Type.Elem().Kind())
  118. }
  119. }
  120. if err := marshal(field.Type, value, field.Tag, w, order); err != nil {
  121. return err
  122. }
  123. }
  124. return nil
  125. }
  126. if oscTag.hasLenPrefix {
  127. buf := &bytes.Buffer{}
  128. if err := marshalEachField(buf); err != nil {
  129. return err
  130. }
  131. // write struct length
  132. if err := marshalUnsignedInt(oscTag.lenPrefix, buf.Len(), w, order); err != nil {
  133. return err
  134. }
  135. // write struct bytes
  136. if buf.Len() > 0 {
  137. _, err := w.Write(buf.Bytes())
  138. return err
  139. }
  140. return nil
  141. }
  142. return marshalEachField(w)
  143. }
  144. func marshalUnsignedInt(intType reflect.Kind, intVal int, w io.Writer, order binary.ByteOrder) error {
  145. switch intType {
  146. case reflect.Uint8:
  147. if err := binary.Write(w, order, uint8(intVal)); err != nil {
  148. return err
  149. }
  150. case reflect.Uint16:
  151. if err := binary.Write(w, order, uint16(intVal)); err != nil {
  152. return err
  153. }
  154. default:
  155. panic(fmt.Sprintf("unsupported type %s. allowed types: uint8, uint16", intType))
  156. }
  157. return nil
  158. }
  159. type oscarTag struct {
  160. hasCountPrefix bool
  161. countPrefix reflect.Kind
  162. hasLenPrefix bool
  163. lenPrefix reflect.Kind
  164. optional bool
  165. nullTerminated bool
  166. }
  167. func parseOSCARTag(tag reflect.StructTag) (oscarTag, error) {
  168. var oscTag oscarTag
  169. val, ok := tag.Lookup("oscar")
  170. if !ok {
  171. return oscTag, nil
  172. }
  173. for _, kv := range strings.Split(val, ",") {
  174. kvSplit := strings.SplitN(kv, "=", 2)
  175. if len(kvSplit) == 2 {
  176. switch kvSplit[0] {
  177. case "len_prefix":
  178. oscTag.hasLenPrefix = true
  179. switch kvSplit[1] {
  180. case "uint8":
  181. oscTag.lenPrefix = reflect.Uint8
  182. case "uint16":
  183. oscTag.lenPrefix = reflect.Uint16
  184. default:
  185. return oscTag, fmt.Errorf("%w: unsupported type %s. allowed types: uint8, uint16",
  186. errInvalidStructTag, kvSplit[1])
  187. }
  188. case "count_prefix":
  189. oscTag.hasCountPrefix = true
  190. switch kvSplit[1] {
  191. case "uint8":
  192. oscTag.countPrefix = reflect.Uint8
  193. case "uint16":
  194. oscTag.countPrefix = reflect.Uint16
  195. default:
  196. return oscTag, fmt.Errorf("%w: unsupported type %s. allowed types: uint8, uint16",
  197. errInvalidStructTag, kvSplit[1])
  198. }
  199. }
  200. } else {
  201. switch kvSplit[0] {
  202. case "optional":
  203. oscTag.optional = true
  204. case "nullterm":
  205. oscTag.nullTerminated = true
  206. default:
  207. return oscTag, fmt.Errorf("%w: unsupported struct tag %s",
  208. errInvalidStructTag, kvSplit[0])
  209. }
  210. }
  211. }
  212. var err error
  213. if oscTag.hasCountPrefix && oscTag.hasLenPrefix {
  214. err = fmt.Errorf("%w: struct elem has both len_prefix and count_prefix", errInvalidStructTag)
  215. }
  216. return oscTag, err
  217. }