1
2
3
4
5 package flate
6
7 import (
8 "math"
9 "math/bits"
10 "slices"
11 "sync"
12 )
13
14 const (
15 maxBitsLimit = 16
16
17 literalCount = 286
18 )
19
20
21 type hcode uint32
22
23
24 func (h hcode) len() uint8 {
25 return uint8(h)
26 }
27
28
29 func (h hcode) code64() uint64 {
30 return uint64(h >> 8)
31 }
32
33
34 func (h hcode) zero() bool {
35 return h == 0
36 }
37
38
39 func (h *hcode) set(code uint16, length uint8) {
40 *h = newhcode(code, length)
41 }
42
43
44 func newhcode(code uint16, length uint8) hcode {
45 return hcode(length) | (hcode(code) << 8)
46 }
47
48
49
50
51 type huffmanEncoder struct {
52 codes []hcode
53 bitCount [17]int32
54
55
56
57
58 freqcache [literalCount + 1]literalNode
59 }
60
61
62 func newHuffmanEncoder(size int) *huffmanEncoder {
63
64 c := uint(bits.Len32(uint32(size - 1)))
65 return &huffmanEncoder{codes: make([]hcode, size, 1<<c)}
66 }
67
68
69 type literalNode struct {
70 literal uint16
71 freq uint16
72 }
73
74
75 func maxNode() literalNode { return literalNode{math.MaxUint16, math.MaxUint16} }
76
77
78 type levelInfo struct {
79
80 level int32
81
82
83 lastFreq int32
84
85
86 nextCharFreq int32
87
88
89
90 nextPairFreq int32
91
92
93
94 needed int32
95 }
96
97
98
99 func reverseBits(x uint16, b byte) uint16 {
100 return bits.Reverse16(x << ((16 - b) & 15))
101 }
102
103
104 func generateFixedLiteralEncoding() *huffmanEncoder {
105 h := newHuffmanEncoder(literalCount)
106 codes := h.codes
107 var ch uint16
108 for ch = range uint16(literalCount) {
109 var bits uint16
110 var size uint8
111 switch {
112 case ch < 144:
113
114 bits = ch + 48
115 size = 8
116 case ch < 256:
117
118 bits = ch + 400 - 144
119 size = 9
120 case ch < 280:
121
122 bits = ch - 256
123 size = 7
124 default:
125
126 bits = ch + 192 - 280
127 size = 8
128 }
129 codes[ch] = newhcode(reverseBits(bits, size), size)
130 }
131 return h
132 }
133
134 func generateFixedOffsetEncoding() *huffmanEncoder {
135 h := newHuffmanEncoder(30)
136 codes := h.codes
137 for ch := range codes {
138 codes[ch] = newhcode(reverseBits(uint16(ch), 5), 5)
139 }
140 return h
141 }
142
143 var (
144 fixedLiteralEncoding = sync.OnceValue(generateFixedLiteralEncoding)
145 fixedOffsetEncoding = sync.OnceValue(generateFixedOffsetEncoding)
146 )
147
148
149 func (h *huffmanEncoder) bitLength(freq []uint16) int {
150 var total int
151 for i, f := range freq {
152 if f != 0 {
153 total += int(f) * int(h.codes[i].len())
154 }
155 }
156 return total
157 }
158
159
160
161 func (h *huffmanEncoder) bitLengthRaw(b []byte) int {
162 var total int
163 for _, f := range b {
164 total += max(1, int(h.codes[f].len()))
165 }
166 return total
167 }
168
169
170
171 func (h *huffmanEncoder) canEncodeLen(freq []uint16) int {
172 var total int
173 for i, f := range freq {
174 if f != 0 {
175 code := h.codes[i]
176 if code.zero() {
177 return math.MaxInt32
178 }
179 total += int(f) * int(code.len())
180 }
181 }
182 return total
183 }
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198 func (h *huffmanEncoder) bitCounts(list []literalNode, maxBits int32) []int32 {
199 if maxBits >= maxBitsLimit {
200 panic("flate: maxBits too large")
201 }
202 n := int32(len(list))
203 list = list[0 : n+1]
204 list[n] = maxNode()
205
206
207
208 if maxBits > n-1 {
209 maxBits = n - 1
210 }
211
212
213
214
215
216 var levels [maxBitsLimit]levelInfo
217
218
219
220
221 var leafCounts [maxBitsLimit][maxBitsLimit]int32
222
223 _ = list[2]
224 for level := int32(1); level <= maxBits; level++ {
225
226
227 levels[level] = levelInfo{
228 level: level,
229 lastFreq: int32(list[1].freq),
230 nextCharFreq: int32(list[2].freq),
231 nextPairFreq: int32(list[0].freq) + int32(list[1].freq),
232 }
233 leafCounts[level][level] = 2
234 if level == 1 {
235 levels[level].nextPairFreq = math.MaxInt32
236 }
237 }
238
239
240 levels[maxBits].needed = 2*n - 4
241
242 level := uint32(maxBits)
243 for level < 16 {
244 l := &levels[level]
245 if l.nextPairFreq == math.MaxInt32 && l.nextCharFreq == math.MaxInt32 {
246
247
248
249
250 l.needed = 0
251 levels[level+1].nextPairFreq = math.MaxInt32
252 level++
253 continue
254 }
255
256 prevFreq := l.lastFreq
257 if l.nextCharFreq < l.nextPairFreq {
258
259 n := leafCounts[level][level] + 1
260 l.lastFreq = l.nextCharFreq
261
262 leafCounts[level][level] = n
263 e := list[n]
264 if e.literal < math.MaxUint16 {
265 l.nextCharFreq = int32(e.freq)
266 } else {
267 l.nextCharFreq = math.MaxInt32
268 }
269 } else {
270
271
272
273 l.lastFreq = l.nextPairFreq
274
275 save := leafCounts[level][level]
276 leafCounts[level] = leafCounts[level-1]
277 leafCounts[level][level] = save
278 levels[l.level-1].needed = 2
279 }
280
281 if l.needed--; l.needed == 0 {
282
283
284
285
286 if l.level == maxBits {
287
288 break
289 }
290 levels[l.level+1].nextPairFreq = prevFreq + l.lastFreq
291 level++
292 } else {
293
294 for levels[level-1].needed > 0 {
295 level--
296 }
297 }
298 }
299
300
301
302 if leafCounts[maxBits][maxBits] != n {
303 panic("leafCounts[maxBits][maxBits] != n")
304 }
305
306 bitCount := h.bitCount[:maxBits+1]
307 bits := 1
308 counts := &leafCounts[maxBits]
309 for level := maxBits; level > 0; level-- {
310
311
312 bitCount[bits] = counts[level] - counts[level-1]
313 bits++
314 }
315 return bitCount
316 }
317
318
319
320 func (h *huffmanEncoder) assignEncodingAndSize(bitCount []int32, list []literalNode) {
321 code := uint16(0)
322 for n, bits := range bitCount {
323 code <<= 1
324 if n == 0 || bits == 0 {
325 continue
326 }
327
328
329
330
331 chunk := list[len(list)-int(bits):]
332
333 slices.SortFunc(chunk, func(a, b literalNode) int {
334 return int(a.literal) - int(b.literal)
335 })
336 for _, node := range chunk {
337 h.codes[node.literal] = newhcode(reverseBits(code, uint8(n)), uint8(n))
338 code++
339 }
340 list = list[0 : len(list)-int(bits)]
341 }
342 }
343
344
345
346
347 func (h *huffmanEncoder) generate(freq []uint16, maxBits int32) {
348 list := h.freqcache[:len(freq)+1]
349 codes := h.codes[:len(freq)]
350
351 count := 0
352
353 for i, f := range freq {
354 if f != 0 {
355 list[count] = literalNode{uint16(i), f}
356 count++
357 } else {
358 codes[i] = 0
359 }
360 }
361 list[count] = literalNode{}
362
363 list = list[:count]
364 if count <= 2 {
365
366
367 for i, node := range list {
368
369 h.codes[node.literal].set(uint16(i), 1)
370 }
371 return
372 }
373 slices.SortFunc(list, func(a, b literalNode) int {
374
375 return (int(a.freq)<<10 + int(a.literal)) - (int(b.freq)<<10 + int(b.literal))
376 })
377
378
379 bitCount := h.bitCounts(list, maxBits)
380
381 h.assignEncodingAndSize(bitCount, list)
382 }
383
384 func histogram(b []byte, h []uint16) {
385 if len(b) >= 8<<10 {
386 histogramSplit(b, h)
387 return
388 }
389 h = h[:256]
390 for _, t := range b {
391 h[t]++
392 }
393 }
394
395 func histogramSplit(b []byte, h []uint16) {
396
397
398 h = h[:256]
399
400 for len(b)&3 != 0 {
401 h[b[0]]++
402 b = b[1:]
403 }
404 n := len(b) / 4
405 x, y, z, w := b[:n], b[n:], b[n+n:], b[n+n+n:]
406 y, z, w = y[:len(x)], z[:len(x)], w[:len(x)]
407 for i, t := range x {
408 v0 := &h[t]
409 v1 := &h[y[i]]
410 v2 := &h[z[i]]
411 v3 := &h[w[i]]
412 *v0++
413 *v1++
414 *v2++
415 *v3++
416 }
417 }
418
View as plain text