encode.go 6.0 KB

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