seqdec_amd64.go 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387
  1. //go:build amd64 && !appengine && !noasm && gc
  2. package zstd
  3. import (
  4. "fmt"
  5. "io"
  6. "github.com/klauspost/compress/internal/cpuinfo"
  7. )
  8. type decodeSyncAsmContext struct {
  9. llTable []decSymbol
  10. mlTable []decSymbol
  11. ofTable []decSymbol
  12. llState uint64
  13. mlState uint64
  14. ofState uint64
  15. iteration int
  16. litRemain int
  17. out []byte
  18. outPosition int
  19. literals []byte
  20. litPosition int
  21. history []byte
  22. windowSize int
  23. ll int // set on error (not for all errors, please refer to _generate/gen.go)
  24. ml int // set on error (not for all errors, please refer to _generate/gen.go)
  25. mo int // set on error (not for all errors, please refer to _generate/gen.go)
  26. }
  27. // sequenceDecs_decodeSync_amd64 implements the main loop of sequenceDecs.decodeSync in x86 asm.
  28. //
  29. // Please refer to seqdec_generic.go for the reference implementation.
  30. //
  31. //go:noescape
  32. func sequenceDecs_decodeSync_amd64(s *sequenceDecs, br *bitReader, ctx *decodeSyncAsmContext) int
  33. // sequenceDecs_decodeSync_bmi2 implements the main loop of sequenceDecs.decodeSync in x86 asm with BMI2 extensions.
  34. //
  35. //go:noescape
  36. func sequenceDecs_decodeSync_bmi2(s *sequenceDecs, br *bitReader, ctx *decodeSyncAsmContext) int
  37. // sequenceDecs_decodeSync_safe_amd64 does the same as above, but does not write more than output buffer.
  38. //
  39. //go:noescape
  40. func sequenceDecs_decodeSync_safe_amd64(s *sequenceDecs, br *bitReader, ctx *decodeSyncAsmContext) int
  41. // sequenceDecs_decodeSync_safe_bmi2 does the same as above, but does not write more than output buffer.
  42. //
  43. //go:noescape
  44. func sequenceDecs_decodeSync_safe_bmi2(s *sequenceDecs, br *bitReader, ctx *decodeSyncAsmContext) int
  45. // decode sequences from the stream with the provided history but without a dictionary.
  46. func (s *sequenceDecs) decodeSyncSimple(hist []byte) (bool, error) {
  47. if len(s.dict) > 0 {
  48. return false, nil
  49. }
  50. if s.maxSyncLen == 0 && cap(s.out)-len(s.out) < maxCompressedBlockSize {
  51. return false, nil
  52. }
  53. // FIXME: Using unsafe memory copies leads to rare, random crashes
  54. // with fuzz testing. It is therefore disabled for now.
  55. const useSafe = true
  56. /*
  57. useSafe := false
  58. if s.maxSyncLen == 0 && cap(s.out)-len(s.out) < maxCompressedBlockSizeAlloc {
  59. useSafe = true
  60. }
  61. if s.maxSyncLen > 0 && cap(s.out)-len(s.out)-compressedBlockOverAlloc < int(s.maxSyncLen) {
  62. useSafe = true
  63. }
  64. if cap(s.literals) < len(s.literals)+compressedBlockOverAlloc {
  65. useSafe = true
  66. }
  67. */
  68. br := s.br
  69. maxBlockSize := min(s.windowSize, maxCompressedBlockSize)
  70. ctx := decodeSyncAsmContext{
  71. llTable: s.litLengths.fse.dt[:maxTablesize],
  72. mlTable: s.matchLengths.fse.dt[:maxTablesize],
  73. ofTable: s.offsets.fse.dt[:maxTablesize],
  74. llState: uint64(s.litLengths.state.state),
  75. mlState: uint64(s.matchLengths.state.state),
  76. ofState: uint64(s.offsets.state.state),
  77. iteration: s.nSeqs - 1,
  78. litRemain: len(s.literals),
  79. out: s.out,
  80. outPosition: len(s.out),
  81. literals: s.literals,
  82. windowSize: s.windowSize,
  83. history: hist,
  84. }
  85. s.seqSize = 0
  86. startSize := len(s.out)
  87. var errCode int
  88. if cpuinfo.HasBMI2() {
  89. if useSafe {
  90. errCode = sequenceDecs_decodeSync_safe_bmi2(s, br, &ctx)
  91. } else {
  92. errCode = sequenceDecs_decodeSync_bmi2(s, br, &ctx)
  93. }
  94. } else {
  95. if useSafe {
  96. errCode = sequenceDecs_decodeSync_safe_amd64(s, br, &ctx)
  97. } else {
  98. errCode = sequenceDecs_decodeSync_amd64(s, br, &ctx)
  99. }
  100. }
  101. switch errCode {
  102. case noError:
  103. break
  104. case errorMatchLenOfsMismatch:
  105. return true, fmt.Errorf("zero matchoff and matchlen (%d) > 0", ctx.ml)
  106. case errorMatchLenTooBig:
  107. return true, fmt.Errorf("match len (%d) bigger than max allowed length", ctx.ml)
  108. case errorMatchOffTooBig:
  109. return true, fmt.Errorf("match offset (%d) bigger than current history (%d)",
  110. ctx.mo, ctx.outPosition+len(hist)-startSize)
  111. case errorNotEnoughLiterals:
  112. return true, fmt.Errorf("unexpected literal count, want %d bytes, but only %d is available",
  113. ctx.ll, ctx.litRemain+ctx.ll)
  114. case errorOverread:
  115. return true, io.ErrUnexpectedEOF
  116. case errorNotEnoughSpace:
  117. size := ctx.outPosition + ctx.ll + ctx.ml
  118. if debugDecoder {
  119. println("msl:", s.maxSyncLen, "cap", cap(s.out), "bef:", startSize, "sz:", size-startSize, "mbs:", maxBlockSize, "outsz:", cap(s.out)-startSize)
  120. }
  121. return true, fmt.Errorf("output bigger than max block size (%d)", maxBlockSize)
  122. default:
  123. return true, fmt.Errorf("sequenceDecs_decode returned erroneous code %d", errCode)
  124. }
  125. s.seqSize += ctx.litRemain
  126. if s.seqSize > maxBlockSize {
  127. return true, fmt.Errorf("output bigger than max block size (%d)", maxBlockSize)
  128. }
  129. err := br.close()
  130. if err != nil {
  131. printf("Closing sequences: %v, %+v\n", err, *br)
  132. return true, err
  133. }
  134. s.literals = s.literals[ctx.litPosition:]
  135. t := ctx.outPosition
  136. s.out = s.out[:t]
  137. // Add final literals
  138. s.out = append(s.out, s.literals...)
  139. if debugDecoder {
  140. t += len(s.literals)
  141. if t != len(s.out) {
  142. panic(fmt.Errorf("length mismatch, want %d, got %d", len(s.out), t))
  143. }
  144. }
  145. return true, nil
  146. }
  147. // --------------------------------------------------------------------------------
  148. type decodeAsmContext struct {
  149. llTable []decSymbol
  150. mlTable []decSymbol
  151. ofTable []decSymbol
  152. llState uint64
  153. mlState uint64
  154. ofState uint64
  155. iteration int
  156. seqs []seqVals
  157. litRemain int
  158. }
  159. const noError = 0
  160. // error reported when mo == 0 && ml > 0
  161. const errorMatchLenOfsMismatch = 1
  162. // error reported when ml > maxMatchLen
  163. const errorMatchLenTooBig = 2
  164. // error reported when mo > available history or mo > s.windowSize
  165. const errorMatchOffTooBig = 3
  166. // error reported when the sum of literal lengths exeeceds the literal buffer size
  167. const errorNotEnoughLiterals = 4
  168. // error reported when capacity of `out` is too small
  169. const errorNotEnoughSpace = 5
  170. // error reported when bits are overread.
  171. const errorOverread = 6
  172. // sequenceDecs_decode implements the main loop of sequenceDecs in x86 asm.
  173. //
  174. // Please refer to seqdec_generic.go for the reference implementation.
  175. //
  176. //go:noescape
  177. func sequenceDecs_decode_amd64(s *sequenceDecs, br *bitReader, ctx *decodeAsmContext) int
  178. // sequenceDecs_decode implements the main loop of sequenceDecs in x86 asm.
  179. //
  180. // Please refer to seqdec_generic.go for the reference implementation.
  181. //
  182. //go:noescape
  183. func sequenceDecs_decode_56_amd64(s *sequenceDecs, br *bitReader, ctx *decodeAsmContext) int
  184. // sequenceDecs_decode implements the main loop of sequenceDecs in x86 asm with BMI2 extensions.
  185. //
  186. //go:noescape
  187. func sequenceDecs_decode_bmi2(s *sequenceDecs, br *bitReader, ctx *decodeAsmContext) int
  188. // sequenceDecs_decode implements the main loop of sequenceDecs in x86 asm with BMI2 extensions.
  189. //
  190. //go:noescape
  191. func sequenceDecs_decode_56_bmi2(s *sequenceDecs, br *bitReader, ctx *decodeAsmContext) int
  192. // decode sequences from the stream without the provided history.
  193. func (s *sequenceDecs) decode(seqs []seqVals) error {
  194. br := s.br
  195. maxBlockSize := min(s.windowSize, maxCompressedBlockSize)
  196. ctx := decodeAsmContext{
  197. llTable: s.litLengths.fse.dt[:maxTablesize],
  198. mlTable: s.matchLengths.fse.dt[:maxTablesize],
  199. ofTable: s.offsets.fse.dt[:maxTablesize],
  200. llState: uint64(s.litLengths.state.state),
  201. mlState: uint64(s.matchLengths.state.state),
  202. ofState: uint64(s.offsets.state.state),
  203. seqs: seqs,
  204. iteration: len(seqs) - 1,
  205. litRemain: len(s.literals),
  206. }
  207. if debugDecoder {
  208. println("decode: decoding", len(seqs), "sequences", br.remain(), "bits remain on stream")
  209. }
  210. s.seqSize = 0
  211. lte56bits := s.maxBits+s.offsets.fse.actualTableLog+s.matchLengths.fse.actualTableLog+s.litLengths.fse.actualTableLog <= 56
  212. var errCode int
  213. if cpuinfo.HasBMI2() {
  214. if lte56bits {
  215. errCode = sequenceDecs_decode_56_bmi2(s, br, &ctx)
  216. } else {
  217. errCode = sequenceDecs_decode_bmi2(s, br, &ctx)
  218. }
  219. } else {
  220. if lte56bits {
  221. errCode = sequenceDecs_decode_56_amd64(s, br, &ctx)
  222. } else {
  223. errCode = sequenceDecs_decode_amd64(s, br, &ctx)
  224. }
  225. }
  226. if errCode != 0 {
  227. i := len(seqs) - ctx.iteration - 1
  228. switch errCode {
  229. case errorMatchLenOfsMismatch:
  230. ml := ctx.seqs[i].ml
  231. return fmt.Errorf("zero matchoff and matchlen (%d) > 0", ml)
  232. case errorMatchLenTooBig:
  233. ml := ctx.seqs[i].ml
  234. return fmt.Errorf("match len (%d) bigger than max allowed length", ml)
  235. case errorNotEnoughLiterals:
  236. ll := ctx.seqs[i].ll
  237. return fmt.Errorf("unexpected literal count, want %d bytes, but only %d is available", ll, ctx.litRemain+ll)
  238. case errorOverread:
  239. return io.ErrUnexpectedEOF
  240. }
  241. return fmt.Errorf("sequenceDecs_decode_amd64 returned erroneous code %d", errCode)
  242. }
  243. if ctx.litRemain < 0 {
  244. return fmt.Errorf("literal count is too big: total available %d, total requested %d",
  245. len(s.literals), len(s.literals)-ctx.litRemain)
  246. }
  247. s.seqSize += ctx.litRemain
  248. if s.seqSize > maxBlockSize {
  249. return fmt.Errorf("output bigger than max block size (%d)", maxBlockSize)
  250. }
  251. if debugDecoder {
  252. println("decode: ", br.remain(), "bits remain on stream. code:", errCode)
  253. }
  254. err := br.close()
  255. if err != nil {
  256. printf("Closing sequences: %v, %+v\n", err, *br)
  257. }
  258. return err
  259. }
  260. // --------------------------------------------------------------------------------
  261. type executeAsmContext struct {
  262. seqs []seqVals
  263. seqIndex int
  264. out []byte
  265. history []byte
  266. literals []byte
  267. outPosition int
  268. litPosition int
  269. windowSize int
  270. }
  271. // sequenceDecs_executeSimple_amd64 implements the main loop of sequenceDecs.executeSimple in x86 asm.
  272. //
  273. // Returns false if a match offset is too big.
  274. //
  275. // Please refer to seqdec_generic.go for the reference implementation.
  276. //
  277. //go:noescape
  278. func sequenceDecs_executeSimple_amd64(ctx *executeAsmContext) bool
  279. // Same as above, but with safe memcopies
  280. //
  281. //go:noescape
  282. func sequenceDecs_executeSimple_safe_amd64(ctx *executeAsmContext) bool
  283. // executeSimple handles cases when dictionary is not used.
  284. func (s *sequenceDecs) executeSimple(seqs []seqVals, hist []byte) error {
  285. // Ensure we have enough output size...
  286. if len(s.out)+s.seqSize+compressedBlockOverAlloc > cap(s.out) {
  287. addBytes := s.seqSize + len(s.out) + compressedBlockOverAlloc
  288. s.out = append(s.out, make([]byte, addBytes)...)
  289. s.out = s.out[:len(s.out)-addBytes]
  290. }
  291. if debugDecoder {
  292. printf("Execute %d seqs with literals: %d into %d bytes\n", len(seqs), len(s.literals), s.seqSize)
  293. }
  294. var t = len(s.out)
  295. out := s.out[:t+s.seqSize]
  296. ctx := executeAsmContext{
  297. seqs: seqs,
  298. seqIndex: 0,
  299. out: out,
  300. history: hist,
  301. outPosition: t,
  302. litPosition: 0,
  303. literals: s.literals,
  304. windowSize: s.windowSize,
  305. }
  306. var ok bool
  307. if cap(s.literals) < len(s.literals)+compressedBlockOverAlloc {
  308. ok = sequenceDecs_executeSimple_safe_amd64(&ctx)
  309. } else {
  310. ok = sequenceDecs_executeSimple_amd64(&ctx)
  311. }
  312. if !ok {
  313. return fmt.Errorf("match offset (%d) bigger than current history (%d)",
  314. seqs[ctx.seqIndex].mo, ctx.outPosition+len(hist))
  315. }
  316. s.literals = s.literals[ctx.litPosition:]
  317. t = ctx.outPosition
  318. // Add final literals
  319. copy(out[t:], s.literals)
  320. if debugDecoder {
  321. t += len(s.literals)
  322. if t != len(out) {
  323. panic(fmt.Errorf("length mismatch, want %d, got %d, ss: %d", len(out), t, s.seqSize))
  324. }
  325. }
  326. s.out = out
  327. return nil
  328. }