encode.go 6.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238
  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 && str != "" {
  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. if str == "" {
  105. return nil
  106. }
  107. return binary.Write(w, order, []byte(str))
  108. }
  109. func marshalStruct(t reflect.Type, v reflect.Value, oscTag oscarTag, w io.Writer, order binary.ByteOrder) error {
  110. marshalEachField := func(w io.Writer) error {
  111. for i := 0; i < t.NumField(); i++ {
  112. field := t.Field(i)
  113. value := v.Field(i)
  114. if field.Type.Kind() == reflect.Ptr {
  115. if i != t.NumField()-1 {
  116. return fmt.Errorf("pointer type found at non-final field %s", field.Name)
  117. }
  118. if field.Type.Elem().Kind() != reflect.Struct {
  119. return fmt.Errorf("field %s must point to a struct, got %v instead", field.Name,
  120. field.Type.Elem().Kind())
  121. }
  122. }
  123. if err := marshal(field.Type, value, field.Tag, w, order); err != nil {
  124. return err
  125. }
  126. }
  127. return nil
  128. }
  129. if oscTag.hasLenPrefix {
  130. buf := &bytes.Buffer{}
  131. if err := marshalEachField(buf); err != nil {
  132. return err
  133. }
  134. // write struct length
  135. if err := marshalUnsignedInt(oscTag.lenPrefix, buf.Len(), w, order); err != nil {
  136. return err
  137. }
  138. // write struct bytes
  139. if buf.Len() > 0 {
  140. _, err := w.Write(buf.Bytes())
  141. return err
  142. }
  143. return nil
  144. }
  145. return marshalEachField(w)
  146. }
  147. func marshalUnsignedInt(intType reflect.Kind, intVal int, w io.Writer, order binary.ByteOrder) error {
  148. switch intType {
  149. case reflect.Uint8:
  150. if err := binary.Write(w, order, uint8(intVal)); err != nil {
  151. return err
  152. }
  153. case reflect.Uint16:
  154. if err := binary.Write(w, order, uint16(intVal)); err != nil {
  155. return err
  156. }
  157. default:
  158. panic(fmt.Sprintf("unsupported type %s. allowed types: uint8, uint16", intType))
  159. }
  160. return nil
  161. }
  162. type oscarTag struct {
  163. hasCountPrefix bool
  164. countPrefix reflect.Kind
  165. hasLenPrefix bool
  166. lenPrefix reflect.Kind
  167. optional bool
  168. nullTerminated bool
  169. }
  170. func parseOSCARTag(tag reflect.StructTag) (oscarTag, error) {
  171. var oscTag oscarTag
  172. val, ok := tag.Lookup("oscar")
  173. if !ok {
  174. return oscTag, nil
  175. }
  176. for _, kv := range strings.Split(val, ",") {
  177. kvSplit := strings.SplitN(kv, "=", 2)
  178. if len(kvSplit) == 2 {
  179. switch kvSplit[0] {
  180. case "len_prefix":
  181. oscTag.hasLenPrefix = true
  182. switch kvSplit[1] {
  183. case "uint8":
  184. oscTag.lenPrefix = reflect.Uint8
  185. case "uint16":
  186. oscTag.lenPrefix = reflect.Uint16
  187. default:
  188. return oscTag, fmt.Errorf("%w: unsupported type %s. allowed types: uint8, uint16",
  189. errInvalidStructTag, kvSplit[1])
  190. }
  191. case "count_prefix":
  192. oscTag.hasCountPrefix = true
  193. switch kvSplit[1] {
  194. case "uint8":
  195. oscTag.countPrefix = reflect.Uint8
  196. case "uint16":
  197. oscTag.countPrefix = reflect.Uint16
  198. default:
  199. return oscTag, fmt.Errorf("%w: unsupported type %s. allowed types: uint8, uint16",
  200. errInvalidStructTag, kvSplit[1])
  201. }
  202. }
  203. } else {
  204. switch kvSplit[0] {
  205. case "optional":
  206. oscTag.optional = true
  207. case "nullterm":
  208. oscTag.nullTerminated = true
  209. default:
  210. return oscTag, fmt.Errorf("%w: unsupported struct tag %s",
  211. errInvalidStructTag, kvSplit[0])
  212. }
  213. }
  214. }
  215. var err error
  216. if oscTag.hasCountPrefix && oscTag.hasLenPrefix {
  217. err = fmt.Errorf("%w: struct elem has both len_prefix and count_prefix", errInvalidStructTag)
  218. }
  219. return oscTag, err
  220. }