ProtoCore v0.0.2
Deterministic, zero-heap network stack for embedded targets
Loading...
Searching...
No Matches
inflate.cpp
Go to the documentation of this file.
1// Copyright (C) 2026 Douglas Quigg (dstroy0) <dquigg123@gmail.com>
2// SPDX-License-Identifier: AGPL-3.0-or-later
3
4/**
5 * @file inflate.cpp
6 * @brief Bounded RFC 1951 DEFLATE decompressor - implementation.
7 *
8 * A compact canonical-Huffman INFLATE (the classic count[]/symbol[] decode, as
9 * in Mark Adler's "puff" reference) written from RFC 1951. Decoding is bit by
10 * bit - small and deterministic rather than fast, which suits the small messages
11 * this serves. All state is on the stack plus a caller-supplied table scratch;
12 * the output buffer doubles as the LZ77 window (see inflate.h).
13 */
14
15#include "inflate.h"
16
17#if PC_ENABLE_WS_DEFLATE
18
19#include <string.h>
20
21namespace
22{
23#define PC_MAXBITS 15 // max bits in a Huffman code
24#define PC_MAXLCODES 288 // max literal/length codes
25#define PC_MAXDCODES 32 // max distance codes (30 used; 32 for safety)
26
27// Huffman decoding table: count[len] = #codes of that length, symbol[] = symbols
28// in canonical order. Both point into the caller's table scratch.
29struct Huffman
30{
31 short *count;
32 short *symbol;
33};
34
35// All the table memory inflate_raw() needs, laid over the caller's scratch.
36struct Tables
37{
38 short lcount[PC_MAXBITS + 1];
39 short lsym[PC_MAXLCODES];
40 short dcount[PC_MAXBITS + 1];
41 short dsym[PC_MAXDCODES];
42 short lengths[PC_MAXLCODES + PC_MAXDCODES]; // code lengths during construction
43};
44static_assert(sizeof(Tables) <= INFLATE_SCRATCH_SIZE, "bump INFLATE_SCRATCH_SIZE");
45
46// Decoder state.
47struct State
48{
49 uint8_t *out; // output buffer (also the back-reference window)
50 size_t outcap; // capacity
51 size_t outcnt; // bytes written
52 const uint8_t *in; // input
53 size_t inlen;
54 size_t incnt; // bytes consumed
55 int bitbuf; // bit accumulator (LSB first)
56 int bitcnt; // bits available in bitbuf
57 bool err; // ran out of input mid-element
58};
59
60// Length code base values and extra bits (RFC 1951 sec 3.2.5), codes 257..285.
61const short LEN_BASE[29] = {3, 4, 5, 6, 7, 8, 9, 10, 11, 13, 15, 17, 19, 23, 27,
62 31, 35, 43, 51, 59, 67, 83, 99, 115, 131, 163, 195, 227, 258};
63const short LEN_EXTRA[29] = {0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 2, 2, 2, 2, 3, 3, 3, 3, 4, 4, 4, 4, 5, 5, 5, 5, 0};
64
65// Distance code base values and extra bits, codes 0..29.
66const short DIST_BASE[30] = {1, 2, 3, 4, 5, 7, 9, 13, 17, 25, 33, 49, 65, 97, 129,
67 193, 257, 385, 513, 769, 1025, 1537, 2049, 3073, 4097, 6145, 8193, 12289, 16385, 24577};
68const short DIST_EXTRA[30] = {0, 0, 0, 0, 1, 1, 2, 2, 3, 3, 4, 4, 5, 5, 6,
69 6, 7, 7, 8, 8, 9, 9, 10, 10, 11, 11, 12, 12, 13, 13};
70
71// Pull @p need bits (LSB first). On end-of-input sets s->err and returns 0.
72int bits(State *s, int need)
73{
74 long val = s->bitbuf;
75 while (s->bitcnt < need)
76 {
77 if (s->incnt >= s->inlen)
78 {
79 s->err = true;
80 return 0;
81 }
82 val |= (long)(s->in[s->incnt++]) << s->bitcnt;
83 s->bitcnt += 8;
84 }
85 s->bitbuf = (int)(val >> need);
86 s->bitcnt -= need;
87 return (int)(val & ((1L << need) - 1));
88}
89
90// Decode one symbol using the canonical Huffman table. Returns the symbol, or -1
91// on end-of-input / an invalid code.
92int decode(State *s, const Huffman *h)
93{
94 int code = 0;
95 int first = 0;
96 int index = 0;
97 for (int len = 1; len <= PC_MAXBITS; len++)
98 {
99 code |= bits(s, 1);
100 if (s->err)
101 {
102 return -1;
103 }
104 int count = h->count[len];
105 if (code - count < first) // length len codes start at 'first'
106 {
107 return h->symbol[index + (code - first)];
108 }
109 index += count;
110 first += count;
111 first <<= 1;
112 code <<= 1;
113 }
114 return -1; // ran past PC_MAXBITS without a match
115}
116
117// Build a Huffman table from code lengths. Returns 0 if complete, >0 if
118// incomplete (left-over codes), <0 if over-subscribed.
119int construct(Huffman *h, const short *lengths, int n)
120{
121 for (int len = 0; len <= PC_MAXBITS; len++)
122 {
123 h->count[len] = 0;
124 }
125 for (int sym = 0; sym < n; sym++)
126 {
127 h->count[lengths[sym]]++;
128 }
129 if (h->count[0] == n)
130 {
131 return 0; // no codes at all -> complete (empty)
132 }
133
134 int left = 1;
135 for (int len = 1; len <= PC_MAXBITS; len++)
136 {
137 left <<= 1;
138 left -= h->count[len];
139 if (left < 0)
140 {
141 return left; // over-subscribed
142 }
143 }
144
145 short offs[PC_MAXBITS + 1];
146 offs[1] = 0;
147 for (int len = 1; len < PC_MAXBITS; len++)
148 {
149 offs[len + 1] = offs[len] + h->count[len];
150 }
151 for (int sym = 0; sym < n; sym++)
152 {
153 if (lengths[sym] != 0)
154 {
155 h->symbol[offs[lengths[sym]]++] = (short)sym;
156 }
157 }
158
159 return left; // 0 = complete, >0 = incomplete
160}
161
162// Decode literal/length + distance codes into the output. Returns 0 on the
163// end-of-block symbol, InflateResult::INFLATE_ERR_MALFORMED, or InflateResult::INFLATE_ERR_OVERFLOW.
164InflateResult codes(State *s, const Huffman *lencode, const Huffman *distcode)
165{
166 int symbol;
167 do
168 {
169 symbol = decode(s, lencode);
170 if (symbol < 0)
171 {
172 return InflateResult::INFLATE_ERR_MALFORMED;
173 }
174 if (symbol < 256)
175 {
176 if (s->outcnt >= s->outcap)
177 {
178 return InflateResult::INFLATE_ERR_OVERFLOW;
179 }
180 s->out[s->outcnt++] = (uint8_t)symbol;
181 }
182 else if (symbol > 256)
183 {
184 symbol -= 257;
185 if (symbol >= 29)
186 {
187 return InflateResult::INFLATE_ERR_MALFORMED; // invalid length code (286/287)
188 }
189 int len = LEN_BASE[symbol] + bits(s, LEN_EXTRA[symbol]);
190 if (s->err)
191 {
192 return InflateResult::INFLATE_ERR_MALFORMED;
193 }
194
195 symbol = decode(s, distcode);
196 if (symbol < 0 || symbol >= 30)
197 {
198 return InflateResult::INFLATE_ERR_MALFORMED;
199 }
200 size_t dist = (size_t)(DIST_BASE[symbol] + bits(s, DIST_EXTRA[symbol]));
201 if (s->err)
202 {
203 return InflateResult::INFLATE_ERR_MALFORMED;
204 }
205 if (dist > s->outcnt)
206 {
207 return InflateResult::INFLATE_ERR_MALFORMED; // reference before start of output
208 }
209 if (len > (int)(s->outcap - s->outcnt))
210 {
211 return InflateResult::INFLATE_ERR_OVERFLOW;
212 }
213 for (int k = 0; k < len; k++)
214 {
215 s->out[s->outcnt] = s->out[s->outcnt - dist];
216 s->outcnt++;
217 }
218 }
219 } while (symbol != 256); // 256 = end of block
220 return InflateResult::INFLATE_OK;
221}
222
223// Uncompressed (stored) block: byte-align, read LEN/NLEN, copy LEN bytes.
224InflateResult stored(State *s)
225{
226 s->bitbuf = 0;
227 s->bitcnt = 0; // discard bits to the next byte boundary
228 if (s->incnt + 4 > s->inlen)
229 {
230 return InflateResult::INFLATE_ERR_MALFORMED;
231 }
232 int len = s->in[s->incnt] | (s->in[s->incnt + 1] << 8);
233 int nlen = s->in[s->incnt + 2] | (s->in[s->incnt + 3] << 8);
234 s->incnt += 4;
235 if ((len ^ nlen) != 0xFFFF)
236 {
237 return InflateResult::INFLATE_ERR_MALFORMED; // NLEN must be ones-complement of LEN
238 }
239 if (s->incnt + (size_t)len > s->inlen)
240 {
241 return InflateResult::INFLATE_ERR_MALFORMED;
242 }
243 if ((size_t)len > s->outcap - s->outcnt)
244 {
245 return InflateResult::INFLATE_ERR_OVERFLOW;
246 }
247 memcpy(s->out + s->outcnt, s->in + s->incnt, (size_t)len);
248 s->incnt += (size_t)len;
249 s->outcnt += (size_t)len;
250 return InflateResult::INFLATE_OK;
251}
252
253// Fixed-Huffman block (RFC 1951 sec 3.2.6).
254InflateResult fixed(State *s, Huffman *lencode, Huffman *distcode, short *lengths)
255{
256 int sym = 0;
257 for (; sym < 144; sym++)
258 {
259 lengths[sym] = 8;
260 }
261 for (; sym < 256; sym++)
262 {
263 lengths[sym] = 9;
264 }
265 for (; sym < 280; sym++)
266 {
267 lengths[sym] = 7;
268 }
269 for (; sym < 288; sym++)
270 {
271 lengths[sym] = 8;
272 }
273 construct(lencode, lengths, 288);
274 for (sym = 0; sym < 30; sym++)
275 {
276 lengths[sym] = 5;
277 }
278 construct(distcode, lengths, 30);
279 return codes(s, lencode, distcode);
280}
281
282// Dynamic-Huffman block (RFC 1951 sec 3.2.7).
283InflateResult dynamic(State *s, Huffman *lencode, Huffman *distcode, short *lengths)
284{
285 static const short ORDER[19] = {16, 17, 18, 0, 8, 7, 9, 6, 10, 5, 11, 4, 12, 3, 13, 2, 14, 1, 15};
286
287 int nlen = bits(s, 5) + 257;
288 int ndist = bits(s, 5) + 1;
289 int ncode = bits(s, 4) + 4;
290 if (s->err)
291 {
292 return InflateResult::INFLATE_ERR_MALFORMED;
293 }
294 // GCOVR_EXCL_START unreachable: HLIT/HDIST are 5-bit fields, so nlen = bits(s,5)+257 <= 288 ==
295 // PC_MAXLCODES and ndist = bits(s,5)+1 <= 32 == PC_MAXDCODES always; neither comparison can ever be true.
296 // Kept as the bound that would otherwise let a future encoding change overrun the tables.
297 if (nlen > PC_MAXLCODES || ndist > PC_MAXDCODES)
298 {
299 return InflateResult::INFLATE_ERR_MALFORMED;
300 }
301 // GCOVR_EXCL_STOP
302
303 // Code-length code lengths, in the permuted order.
304 int index;
305 for (index = 0; index < ncode; index++)
306 {
307 lengths[ORDER[index]] = (short)bits(s, 3);
308 }
309 if (s->err)
310 {
311 return InflateResult::INFLATE_ERR_MALFORMED;
312 }
313 for (; index < 19; index++)
314 {
315 lengths[ORDER[index]] = 0;
316 }
317
318 // Build the code-length code (reuse lencode temporarily); it must be complete.
319 if (construct(lencode, lengths, 19) != 0)
320 {
321 return InflateResult::INFLATE_ERR_MALFORMED;
322 }
323
324 // Read the literal/length and distance code lengths.
325 index = 0;
326 while (index < nlen + ndist)
327 {
328 int symbol = decode(s, lencode);
329 if (symbol < 0)
330 {
331 return InflateResult::INFLATE_ERR_MALFORMED;
332 }
333 if (symbol < 16)
334 {
335 lengths[index++] = (short)symbol;
336 continue;
337 }
338 int repeat_len = 0;
339 int repeat;
340 if (symbol == 16)
341 {
342 if (index == 0)
343 {
344 return InflateResult::INFLATE_ERR_MALFORMED; // no previous length to repeat
345 }
346 repeat_len = lengths[index - 1];
347 repeat = 3 + bits(s, 2);
348 }
349 else if (symbol == 17)
350 {
351 repeat = 3 + bits(s, 3);
352 }
353 else // symbol == 18
354 {
355 repeat = 11 + bits(s, 7);
356 }
357 if (s->err)
358 {
359 return InflateResult::INFLATE_ERR_MALFORMED;
360 }
361 if (index + repeat > nlen + ndist)
362 {
363 return InflateResult::INFLATE_ERR_MALFORMED; // repeat past the end
364 }
365 while (repeat--)
366 {
367 lengths[index++] = (short)repeat_len;
368 }
369 }
370
371 if (lengths[256] == 0)
372 {
373 return InflateResult::INFLATE_ERR_MALFORMED; // no end-of-block code
374 }
375
376 // Build the literal/length and distance tables.
377 int err = construct(lencode, lengths, nlen);
378 if (err && (err < 0 || nlen != lencode->count[0] + lencode->count[1]))
379 {
380 return InflateResult::INFLATE_ERR_MALFORMED;
381 }
382 err = construct(distcode, lengths + nlen, ndist);
383 if (err && (err < 0 || ndist != distcode->count[0] + distcode->count[1]))
384 {
385 return InflateResult::INFLATE_ERR_MALFORMED; // incomplete distance code (ok only for 0/1 codes)
386 }
387
388 return codes(s, lencode, distcode);
389}
390} // namespace
391
392InflateResult inflate_raw(const uint8_t *src, size_t src_len, uint8_t *dst, size_t dst_cap, size_t *out_len,
393 void *scratch, size_t scratch_len)
394{
395 if (scratch_len < INFLATE_SCRATCH_SIZE)
396 {
397 return InflateResult::INFLATE_ERR_SCRATCH;
398 }
399
400 Tables *t = (Tables *)scratch;
401 Huffman lencode = {t->lcount, t->lsym};
402 Huffman distcode = {t->dcount, t->dsym};
403
404 State s;
405 s.out = dst;
406 s.outcap = dst_cap;
407 s.outcnt = 0;
408 s.in = src;
409 s.inlen = src_len;
410 s.incnt = 0;
411 s.bitbuf = 0;
412 s.bitcnt = 0;
413 s.err = false;
414
415 int last = 0;
416 do
417 {
418 // Clean end-of-input at a block boundary: permessage-deflate streams have
419 // no final block, so this (not BFINAL) is the normal termination.
420 if (s.incnt >= s.inlen && s.bitcnt == 0)
421 {
422 break;
423 }
424
425 last = bits(&s, 1);
426 int type = bits(&s, 2);
427 if (s.err)
428 {
429 return InflateResult::INFLATE_ERR_MALFORMED;
430 }
431
432 InflateResult rc;
433 if (type == 0)
434 {
435 rc = stored(&s);
436 }
437 else if (type == 1)
438 {
439 rc = fixed(&s, &lencode, &distcode, t->lengths);
440 }
441 else if (type == 2)
442 {
443 rc = dynamic(&s, &lencode, &distcode, t->lengths);
444 }
445 else
446 {
447 return InflateResult::INFLATE_ERR_MALFORMED; // type 3 is reserved
448 }
449
450 if (rc != InflateResult::INFLATE_OK)
451 {
452 return rc;
453 }
454 } while (!last);
455
456 *out_len = s.outcnt;
457 return InflateResult::INFLATE_OK;
458}
459
460#endif // PC_ENABLE_WS_DEFLATE
Bounded RFC 1951 DEFLATE decompressor (INFLATE) - no heap.