encoder.go 16 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671
  1. // Copyright 2019+ Klaus Post. All rights reserved.
  2. // License information can be found in the LICENSE file.
  3. // Based on work by Yann Collet, released under BSD License.
  4. package zstd
  5. import (
  6. "crypto/rand"
  7. "errors"
  8. "fmt"
  9. "io"
  10. "math"
  11. rdebug "runtime/debug"
  12. "sync"
  13. "github.com/klauspost/compress/zstd/internal/xxhash"
  14. )
  15. // Encoder provides encoding to Zstandard.
  16. // An Encoder can be used for either compressing a stream via the
  17. // io.WriteCloser interface supported by the Encoder or as multiple independent
  18. // tasks via the EncodeAll function.
  19. // Smaller encodes are encouraged to use the EncodeAll function.
  20. // Use NewWriter to create a new instance.
  21. type Encoder struct {
  22. o encoderOptions
  23. encoders chan encoder
  24. state encoderState
  25. init sync.Once
  26. }
  27. type encoder interface {
  28. Encode(blk *blockEnc, src []byte)
  29. EncodeNoHist(blk *blockEnc, src []byte)
  30. Block() *blockEnc
  31. CRC() *xxhash.Digest
  32. AppendCRC([]byte) []byte
  33. WindowSize(size int64) int32
  34. UseBlock(*blockEnc)
  35. Reset(d *dict, singleBlock bool)
  36. }
  37. type encoderState struct {
  38. w io.Writer
  39. filling []byte
  40. current []byte
  41. previous []byte
  42. encoder encoder
  43. writing *blockEnc
  44. err error
  45. writeErr error
  46. nWritten int64
  47. nInput int64
  48. frameContentSize int64
  49. headerWritten bool
  50. eofWritten bool
  51. fullFrameWritten bool
  52. // This waitgroup indicates an encode is running.
  53. wg sync.WaitGroup
  54. // This waitgroup indicates we have a block encoding/writing.
  55. wWg sync.WaitGroup
  56. }
  57. // NewWriter will create a new Zstandard encoder.
  58. // If the encoder will be used for encoding blocks a nil writer can be used.
  59. func NewWriter(w io.Writer, opts ...EOption) (*Encoder, error) {
  60. initPredefined()
  61. var e Encoder
  62. e.o.setDefault()
  63. for _, o := range opts {
  64. err := o(&e.o)
  65. if err != nil {
  66. return nil, err
  67. }
  68. }
  69. if w != nil {
  70. e.Reset(w)
  71. }
  72. return &e, nil
  73. }
  74. func (e *Encoder) initialize() {
  75. if e.o.concurrent == 0 {
  76. e.o.setDefault()
  77. }
  78. e.encoders = make(chan encoder, e.o.concurrent)
  79. for i := 0; i < e.o.concurrent; i++ {
  80. enc := e.o.encoder()
  81. e.encoders <- enc
  82. }
  83. }
  84. // Reset will re-initialize the writer and new writes will encode to the supplied writer
  85. // as a new, independent stream.
  86. func (e *Encoder) Reset(w io.Writer) {
  87. s := &e.state
  88. s.wg.Wait()
  89. s.wWg.Wait()
  90. if cap(s.filling) == 0 {
  91. s.filling = make([]byte, 0, e.o.blockSize)
  92. }
  93. if e.o.concurrent > 1 {
  94. if cap(s.current) == 0 {
  95. s.current = make([]byte, 0, e.o.blockSize)
  96. }
  97. if cap(s.previous) == 0 {
  98. s.previous = make([]byte, 0, e.o.blockSize)
  99. }
  100. s.current = s.current[:0]
  101. s.previous = s.previous[:0]
  102. if s.writing == nil {
  103. s.writing = &blockEnc{lowMem: e.o.lowMem}
  104. s.writing.init()
  105. }
  106. s.writing.initNewEncode()
  107. }
  108. if s.encoder == nil {
  109. s.encoder = e.o.encoder()
  110. }
  111. s.filling = s.filling[:0]
  112. s.encoder.Reset(e.o.dict, false)
  113. s.headerWritten = false
  114. s.eofWritten = false
  115. s.fullFrameWritten = false
  116. s.w = w
  117. s.err = nil
  118. s.nWritten = 0
  119. s.nInput = 0
  120. s.writeErr = nil
  121. s.frameContentSize = 0
  122. }
  123. // ResetWithOptions will re-initialize the writer and apply the given options
  124. // as a new, independent stream.
  125. // Options are applied on top of the existing options.
  126. // Some options cannot be changed on reset and will return an error.
  127. func (e *Encoder) ResetWithOptions(w io.Writer, opts ...EOption) error {
  128. e.o.resetOpt = true
  129. defer func() { e.o.resetOpt = false }()
  130. hadDict := e.o.dict != nil
  131. for _, o := range opts {
  132. if err := o(&e.o); err != nil {
  133. return err
  134. }
  135. }
  136. hasDict := e.o.dict != nil
  137. if hadDict != hasDict {
  138. // Dict presence changed — encoder type must be recreated.
  139. e.state.encoder = nil
  140. e.init = sync.Once{}
  141. }
  142. e.Reset(w)
  143. return nil
  144. }
  145. // ResetContentSize will reset and set a content size for the next stream.
  146. // If the bytes written does not match the size given an error will be returned
  147. // when calling Close().
  148. // This is removed when Reset is called.
  149. // Sizes <= 0 results in no content size set.
  150. func (e *Encoder) ResetContentSize(w io.Writer, size int64) {
  151. e.Reset(w)
  152. if size >= 0 {
  153. e.state.frameContentSize = size
  154. }
  155. }
  156. // Write data to the encoder.
  157. // Input data will be buffered and as the buffer fills up
  158. // content will be compressed and written to the output.
  159. // When done writing, use Close to flush the remaining output
  160. // and write CRC if requested.
  161. func (e *Encoder) Write(p []byte) (n int, err error) {
  162. s := &e.state
  163. if s.eofWritten {
  164. return 0, ErrEncoderClosed
  165. }
  166. for len(p) > 0 {
  167. if len(p)+len(s.filling) < e.o.blockSize {
  168. if e.o.crc {
  169. _, _ = s.encoder.CRC().Write(p)
  170. }
  171. s.filling = append(s.filling, p...)
  172. return n + len(p), nil
  173. }
  174. add := p
  175. if len(p)+len(s.filling) > e.o.blockSize {
  176. add = add[:e.o.blockSize-len(s.filling)]
  177. }
  178. if e.o.crc {
  179. _, _ = s.encoder.CRC().Write(add)
  180. }
  181. s.filling = append(s.filling, add...)
  182. p = p[len(add):]
  183. n += len(add)
  184. if len(s.filling) < e.o.blockSize {
  185. return n, nil
  186. }
  187. err := e.nextBlock(false)
  188. if err != nil {
  189. return n, err
  190. }
  191. if debugAsserts && len(s.filling) > 0 {
  192. panic(len(s.filling))
  193. }
  194. }
  195. return n, nil
  196. }
  197. // nextBlock will synchronize and start compressing input in e.state.filling.
  198. // If an error has occurred during encoding it will be returned.
  199. func (e *Encoder) nextBlock(final bool) error {
  200. s := &e.state
  201. // Wait for current block.
  202. s.wg.Wait()
  203. if s.err != nil {
  204. return s.err
  205. }
  206. if len(s.filling) > e.o.blockSize {
  207. return fmt.Errorf("block > maxStoreBlockSize")
  208. }
  209. if !s.headerWritten {
  210. // If we have a single block encode, do a sync compression.
  211. if final && len(s.filling) == 0 && !e.o.fullZero {
  212. s.headerWritten = true
  213. s.fullFrameWritten = true
  214. s.eofWritten = true
  215. return nil
  216. }
  217. if final && len(s.filling) > 0 {
  218. s.current = e.encodeAll(s.encoder, s.filling, s.current[:0])
  219. var n2 int
  220. n2, s.err = s.w.Write(s.current)
  221. if s.err != nil {
  222. return s.err
  223. }
  224. s.nWritten += int64(n2)
  225. s.nInput += int64(len(s.filling))
  226. s.current = s.current[:0]
  227. s.filling = s.filling[:0]
  228. s.headerWritten = true
  229. s.fullFrameWritten = true
  230. s.eofWritten = true
  231. return nil
  232. }
  233. var tmp [maxHeaderSize]byte
  234. fh := frameHeader{
  235. ContentSize: uint64(s.frameContentSize),
  236. WindowSize: uint32(s.encoder.WindowSize(s.frameContentSize)),
  237. SingleSegment: false,
  238. Checksum: e.o.crc,
  239. DictID: e.o.dict.ID(),
  240. }
  241. dst := fh.appendTo(tmp[:0])
  242. s.headerWritten = true
  243. s.wWg.Wait()
  244. var n2 int
  245. n2, s.err = s.w.Write(dst)
  246. if s.err != nil {
  247. return s.err
  248. }
  249. s.nWritten += int64(n2)
  250. }
  251. if s.eofWritten {
  252. // Ensure we only write it once.
  253. final = false
  254. }
  255. if len(s.filling) == 0 {
  256. // Final block, but no data.
  257. if final {
  258. enc := s.encoder
  259. blk := enc.Block()
  260. blk.reset(nil)
  261. blk.last = true
  262. blk.encodeRaw(nil)
  263. s.wWg.Wait()
  264. _, s.err = s.w.Write(blk.output)
  265. s.nWritten += int64(len(blk.output))
  266. s.eofWritten = true
  267. }
  268. return s.err
  269. }
  270. // SYNC:
  271. if e.o.concurrent == 1 {
  272. src := s.filling
  273. s.nInput += int64(len(s.filling))
  274. if debugEncoder {
  275. println("Adding sync block,", len(src), "bytes, final:", final)
  276. }
  277. enc := s.encoder
  278. blk := enc.Block()
  279. blk.reset(nil)
  280. enc.Encode(blk, src)
  281. blk.last = final
  282. if final {
  283. s.eofWritten = true
  284. }
  285. s.err = blk.encode(src, e.o.noEntropy, !e.o.allLitEntropy)
  286. if s.err != nil {
  287. return s.err
  288. }
  289. _, s.err = s.w.Write(blk.output)
  290. s.nWritten += int64(len(blk.output))
  291. s.filling = s.filling[:0]
  292. return s.err
  293. }
  294. // Move blocks forward.
  295. s.filling, s.current, s.previous = s.previous[:0], s.filling, s.current
  296. s.nInput += int64(len(s.current))
  297. s.wg.Add(1)
  298. if final {
  299. s.eofWritten = true
  300. }
  301. go func(src []byte) {
  302. if debugEncoder {
  303. println("Adding block,", len(src), "bytes, final:", final)
  304. }
  305. defer func() {
  306. if r := recover(); r != nil {
  307. s.err = fmt.Errorf("panic while encoding: %v", r)
  308. rdebug.PrintStack()
  309. }
  310. s.wg.Done()
  311. }()
  312. enc := s.encoder
  313. blk := enc.Block()
  314. enc.Encode(blk, src)
  315. blk.last = final
  316. // Wait for pending writes.
  317. s.wWg.Wait()
  318. if s.writeErr != nil {
  319. s.err = s.writeErr
  320. return
  321. }
  322. // Transfer encoders from previous write block.
  323. blk.swapEncoders(s.writing)
  324. // Transfer recent offsets to next.
  325. enc.UseBlock(s.writing)
  326. s.writing = blk
  327. s.wWg.Add(1)
  328. go func() {
  329. defer func() {
  330. if r := recover(); r != nil {
  331. s.writeErr = fmt.Errorf("panic while encoding/writing: %v", r)
  332. rdebug.PrintStack()
  333. }
  334. s.wWg.Done()
  335. }()
  336. s.writeErr = blk.encode(src, e.o.noEntropy, !e.o.allLitEntropy)
  337. if s.writeErr != nil {
  338. return
  339. }
  340. _, s.writeErr = s.w.Write(blk.output)
  341. s.nWritten += int64(len(blk.output))
  342. }()
  343. }(s.current)
  344. return nil
  345. }
  346. // ReadFrom reads data from r until EOF or error.
  347. // The return value n is the number of bytes read.
  348. // Any error except io.EOF encountered during the read is also returned.
  349. //
  350. // The Copy function uses ReaderFrom if available.
  351. func (e *Encoder) ReadFrom(r io.Reader) (n int64, err error) {
  352. if debugEncoder {
  353. println("Using ReadFrom")
  354. }
  355. // Flush any current writes.
  356. if len(e.state.filling) > 0 {
  357. if err := e.nextBlock(false); err != nil {
  358. return 0, err
  359. }
  360. }
  361. e.state.filling = e.state.filling[:e.o.blockSize]
  362. src := e.state.filling
  363. for {
  364. n2, err := r.Read(src)
  365. if e.o.crc {
  366. _, _ = e.state.encoder.CRC().Write(src[:n2])
  367. }
  368. // src is now the unfilled part...
  369. src = src[n2:]
  370. n += int64(n2)
  371. switch err {
  372. case io.EOF:
  373. e.state.filling = e.state.filling[:len(e.state.filling)-len(src)]
  374. if debugEncoder {
  375. println("ReadFrom: got EOF final block:", len(e.state.filling))
  376. }
  377. return n, nil
  378. case nil:
  379. default:
  380. if debugEncoder {
  381. println("ReadFrom: got error:", err)
  382. }
  383. e.state.err = err
  384. return n, err
  385. }
  386. if len(src) > 0 {
  387. if debugEncoder {
  388. println("ReadFrom: got space left in source:", len(src))
  389. }
  390. continue
  391. }
  392. err = e.nextBlock(false)
  393. if err != nil {
  394. return n, err
  395. }
  396. e.state.filling = e.state.filling[:e.o.blockSize]
  397. src = e.state.filling
  398. }
  399. }
  400. // Flush will send the currently written data to output
  401. // and block until everything has been written.
  402. // This should only be used on rare occasions where pushing the currently queued data is critical.
  403. func (e *Encoder) Flush() error {
  404. s := &e.state
  405. if len(s.filling) > 0 {
  406. err := e.nextBlock(false)
  407. if err != nil {
  408. // Ignore Flush after Close.
  409. if errors.Is(s.err, ErrEncoderClosed) {
  410. return nil
  411. }
  412. return err
  413. }
  414. }
  415. s.wg.Wait()
  416. s.wWg.Wait()
  417. if s.err != nil {
  418. // Ignore Flush after Close.
  419. if errors.Is(s.err, ErrEncoderClosed) {
  420. return nil
  421. }
  422. return s.err
  423. }
  424. return s.writeErr
  425. }
  426. // Close will flush the final output and close the stream.
  427. // The function will block until everything has been written.
  428. // The Encoder can still be re-used after calling this.
  429. func (e *Encoder) Close() error {
  430. s := &e.state
  431. if s.encoder == nil {
  432. return nil
  433. }
  434. if s.w == nil {
  435. if len(s.filling) == 0 && !s.headerWritten && !s.eofWritten && s.nInput == 0 {
  436. return nil
  437. }
  438. return errors.New("zstd: encoder has no writer")
  439. }
  440. err := e.nextBlock(true)
  441. if err != nil {
  442. if errors.Is(s.err, ErrEncoderClosed) {
  443. return nil
  444. }
  445. return err
  446. }
  447. if s.frameContentSize > 0 {
  448. if s.nInput != s.frameContentSize {
  449. return fmt.Errorf("frame content size %d given, but %d bytes was written", s.frameContentSize, s.nInput)
  450. }
  451. }
  452. if e.state.fullFrameWritten {
  453. return s.err
  454. }
  455. s.wg.Wait()
  456. s.wWg.Wait()
  457. if s.err != nil {
  458. return s.err
  459. }
  460. if s.writeErr != nil {
  461. return s.writeErr
  462. }
  463. // Write CRC
  464. if e.o.crc && s.err == nil {
  465. // heap alloc.
  466. var tmp [4]byte
  467. _, s.err = s.w.Write(s.encoder.AppendCRC(tmp[:0]))
  468. s.nWritten += 4
  469. }
  470. // Add padding with content from crypto/rand.Reader
  471. if s.err == nil && e.o.pad > 0 {
  472. add := calcSkippableFrame(s.nWritten, int64(e.o.pad))
  473. frame, err := skippableFrame(s.filling[:0], add, rand.Reader)
  474. if err != nil {
  475. return err
  476. }
  477. _, s.err = s.w.Write(frame)
  478. }
  479. if s.err == nil {
  480. s.err = ErrEncoderClosed
  481. return nil
  482. }
  483. return s.err
  484. }
  485. // EncodeAll will encode all input in src and append it to dst.
  486. // This function can be called concurrently, but each call will only run on a single goroutine.
  487. // If empty input is given, nothing is returned, unless WithZeroFrames is specified.
  488. // Encoded blocks can be concatenated and the result will be the combined input stream.
  489. // Data compressed with EncodeAll can be decoded with the Decoder,
  490. // using either a stream or DecodeAll.
  491. func (e *Encoder) EncodeAll(src, dst []byte) []byte {
  492. e.init.Do(e.initialize)
  493. enc := <-e.encoders
  494. defer func() {
  495. e.encoders <- enc
  496. }()
  497. return e.encodeAll(enc, src, dst)
  498. }
  499. func (e *Encoder) encodeAll(enc encoder, src, dst []byte) []byte {
  500. if len(src) == 0 {
  501. if e.o.fullZero {
  502. // Add frame header.
  503. fh := frameHeader{
  504. ContentSize: 0,
  505. WindowSize: MinWindowSize,
  506. SingleSegment: true,
  507. // Adding a checksum would be a waste of space.
  508. Checksum: false,
  509. DictID: 0,
  510. }
  511. dst = fh.appendTo(dst)
  512. // Write raw block as last one only.
  513. var blk blockHeader
  514. blk.setSize(0)
  515. blk.setType(blockTypeRaw)
  516. blk.setLast(true)
  517. dst = blk.appendTo(dst)
  518. }
  519. return dst
  520. }
  521. // Use single segments when above minimum window and below window size.
  522. single := len(src) <= e.o.windowSize && len(src) > MinWindowSize
  523. if e.o.single != nil {
  524. single = *e.o.single
  525. }
  526. fh := frameHeader{
  527. ContentSize: uint64(len(src)),
  528. WindowSize: uint32(enc.WindowSize(int64(len(src)))),
  529. SingleSegment: single,
  530. Checksum: e.o.crc,
  531. DictID: e.o.dict.ID(),
  532. }
  533. // If less than 1MB, allocate a buffer up front.
  534. if len(dst) == 0 && cap(dst) == 0 && len(src) < 1<<20 && !e.o.lowMem {
  535. dst = make([]byte, 0, len(src))
  536. }
  537. dst = fh.appendTo(dst)
  538. // If we can do everything in one block, prefer that.
  539. if len(src) <= e.o.blockSize {
  540. enc.Reset(e.o.dict, true)
  541. // Slightly faster with no history and everything in one block.
  542. if e.o.crc {
  543. _, _ = enc.CRC().Write(src)
  544. }
  545. blk := enc.Block()
  546. blk.last = true
  547. if e.o.dict == nil {
  548. enc.EncodeNoHist(blk, src)
  549. } else {
  550. enc.Encode(blk, src)
  551. }
  552. // If we got the exact same number of literals as input,
  553. // assume the literals cannot be compressed.
  554. oldout := blk.output
  555. // Output directly to dst
  556. blk.output = dst
  557. err := blk.encode(src, e.o.noEntropy, !e.o.allLitEntropy)
  558. if err != nil {
  559. panic(err)
  560. }
  561. dst = blk.output
  562. blk.output = oldout
  563. } else {
  564. enc.Reset(e.o.dict, false)
  565. blk := enc.Block()
  566. for len(src) > 0 {
  567. todo := src
  568. if len(todo) > e.o.blockSize {
  569. todo = todo[:e.o.blockSize]
  570. }
  571. src = src[len(todo):]
  572. if e.o.crc {
  573. _, _ = enc.CRC().Write(todo)
  574. }
  575. blk.pushOffsets()
  576. enc.Encode(blk, todo)
  577. if len(src) == 0 {
  578. blk.last = true
  579. }
  580. err := blk.encode(todo, e.o.noEntropy, !e.o.allLitEntropy)
  581. if err != nil {
  582. panic(err)
  583. }
  584. dst = append(dst, blk.output...)
  585. blk.reset(nil)
  586. }
  587. }
  588. if e.o.crc {
  589. dst = enc.AppendCRC(dst)
  590. }
  591. // Add padding with content from crypto/rand.Reader
  592. if e.o.pad > 0 {
  593. add := calcSkippableFrame(int64(len(dst)), int64(e.o.pad))
  594. var err error
  595. dst, err = skippableFrame(dst, add, rand.Reader)
  596. if err != nil {
  597. panic(err)
  598. }
  599. }
  600. return dst
  601. }
  602. // MaxEncodedSize returns the expected maximum
  603. // size of an encoded block or stream.
  604. func (e *Encoder) MaxEncodedSize(size int) int {
  605. frameHeader := 4 + 2 // magic + frame header & window descriptor
  606. if e.o.dict != nil {
  607. frameHeader += 4
  608. }
  609. // Frame content size:
  610. if size < 256 {
  611. frameHeader++
  612. } else if size < 65536+256 {
  613. frameHeader += 2
  614. } else if size < math.MaxInt32 {
  615. frameHeader += 4
  616. } else {
  617. frameHeader += 8
  618. }
  619. // Final crc
  620. if e.o.crc {
  621. frameHeader += 4
  622. }
  623. // Max overhead is 3 bytes/block.
  624. // There cannot be 0 blocks.
  625. blocks := (size + e.o.blockSize) / e.o.blockSize
  626. // Combine, add padding.
  627. maxSz := frameHeader + 3*blocks + size
  628. if e.o.pad > 1 {
  629. maxSz += calcSkippableFrame(int64(maxSz), int64(e.o.pad))
  630. }
  631. return maxSz
  632. }