24#define PC_MAXLCODES 288
25#define PC_MAXDCODES 32
33const short LEN_BASE[29] = {3, 4, 5, 6, 7, 8, 9, 10, 11, 13, 15, 17, 19, 23, 27,
34 31, 35, 43, 51, 59, 67, 83, 99, 115, 131, 163, 195, 227, 258};
35const 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};
36const short DIST_BASE[30] = {1, 2, 3, 4, 5, 7, 9, 13, 17, 25, 33, 49, 65, 97, 129,
37 193, 257, 385, 513, 769, 1025, 1537, 2049, 3073, 4097, 6145, 8193, 12289, 16385, 24577};
38const short DIST_EXTRA[30] = {0, 0, 0, 0, 1, 1, 2, 2, 3, 3, 4, 4, 5, 5, 6,
39 6, 7, 7, 8, 8, 9, 9, 10, 10, 11, 11, 12, 12, 13, 13};
45 short lcount[PC_MAXBITS + 1];
46 short lsym[PC_MAXLCODES];
47 short dcount[PC_MAXBITS + 1];
48 short dsym[PC_MAXDCODES];
49 short lengths[PC_MAXLCODES + PC_MAXDCODES];
70inline uint8_t logical_byte(
const uint8_t *seg0,
size_t n0,
const uint8_t *seg1,
size_t i)
72 return i < n0 ? seg0[i] : seg1[i - n0];
76int getbits(BitIn *b,
int need)
79 for (
int i = 0; i < need; i++)
81 if (b->bitpos >= b->nbits)
86 size_t byte = b->bitpos >> 3;
87 int bit = (int)(b->bitpos & 7u);
88 v |= ((logical_byte(b->seg0, b->n0, b->seg1,
byte) >> bit) & 1) << i;
95int hdecode(BitIn *b,
const Huffman *h)
100 for (
int len = 1; len <= PC_MAXBITS; len++)
102 code |= getbits(b, 1);
107 int count = h->count[len];
108 if (code - count < first)
110 return h->symbol[index + (code - first)];
121int construct(Huffman *h,
const short *lengths,
int n)
123 for (
int len = 0; len <= PC_MAXBITS; len++)
127 for (
int sym = 0; sym < n; sym++)
129 h->count[lengths[sym]]++;
131 if (h->count[0] == n)
136 for (
int len = 1; len <= PC_MAXBITS; len++)
139 left -= h->count[len];
145 short offs[PC_MAXBITS + 1];
147 for (
int len = 1; len < PC_MAXBITS; len++)
149 offs[len + 1] = offs[len] + h->count[len];
151 for (
int sym = 0; sym < n; sym++)
153 if (lengths[sym] != 0)
155 h->symbol[offs[lengths[sym]]++] = (short)sym;
171void put_byte(OutCtx *o, uint8_t
byte)
173 if (o->cnt >= o->cap)
178 o->dst[o->cnt++] = byte;
179 o->z->window[o->z->wpos] = byte;
180 o->z->wpos = (o->z->wpos + 1u) & (SSH_INFLATE_WINDOW - 1u);
181 if (o->z->whist < SSH_INFLATE_WINDOW)
188int do_codes(BitIn *b, OutCtx *o,
const Huffman *lc,
const Huffman *dc)
192 int sym = hdecode(b, lc);
207 put_byte(o, (uint8_t)sym);
219 int len = LEN_BASE[sym] + getbits(b, LEN_EXTRA[sym]);
224 int dsym = hdecode(b, dc);
229 if (dsym < 0 || dsym >= 30)
233 size_t dist = (size_t)(DIST_BASE[dsym] + getbits(b, DIST_EXTRA[dsym]));
238 if (dist == 0 || dist > o->z->whist)
242 for (
int k = 0; k < len; k++)
244 uint32_t idx = (o->z->wpos - (uint32_t)dist) & (SSH_INFLATE_WINDOW - 1u);
245 put_byte(o, o->z->window[idx]);
255int do_stored(BitIn *b, OutCtx *o)
259 b->bitpos = (b->bitpos + 7u) & ~(
size_t)7u;
261 if (b->bitpos + 32u > b->nbits)
265 int len = getbits(b, 16);
266 int nlen = getbits(b, 16);
267 if ((len ^ nlen) != 0xFFFF)
271 if (b->bitpos + (
size_t)len * 8u > b->nbits)
275 for (
int k = 0; k < len; k++)
277 put_byte(o, (uint8_t)getbits(b, 8));
287int do_fixed(BitIn *b, OutCtx *o, Tables *t)
289 Huffman lc = {t->lcount, t->lsym};
290 Huffman dc = {t->dcount, t->dsym};
292 for (; sym < 144; sym++)
296 for (; sym < 256; sym++)
300 for (; sym < 280; sym++)
304 for (; sym < 288; sym++)
308 construct(&lc, t->lengths, 288);
309 for (sym = 0; sym < 30; sym++)
313 construct(&dc, t->lengths, 30);
314 return do_codes(b, o, &lc, &dc);
319int do_dynamic(BitIn *b, OutCtx *o, Tables *t)
321 static const short ORDER[19] = {16, 17, 18, 0, 8, 7, 9, 6, 10, 5, 11, 4, 12, 3, 13, 2, 14, 1, 15};
322 Huffman lc = {t->lcount, t->lsym};
323 Huffman dc = {t->dcount, t->dsym};
325 int nlen = getbits(b, 5) + 257;
326 int ndist = getbits(b, 5) + 1;
327 int ncode = getbits(b, 4) + 4;
332 if (nlen > PC_MAXLCODES || ndist > PC_MAXDCODES)
338 for (index = 0; index < ncode; index++)
340 t->lengths[ORDER[index]] = (short)getbits(b, 3);
346 for (; index < 19; index++)
348 t->lengths[ORDER[index]] = 0;
350 if (construct(&lc, t->lengths, 19) != 0)
356 while (index < nlen + ndist)
358 int symbol = hdecode(b, &lc);
369 t->lengths[index++] = (short)symbol;
380 repeat_len = t->lengths[index - 1];
381 repeat = 3 + getbits(b, 2);
383 else if (symbol == 17)
385 repeat = 3 + getbits(b, 3);
389 repeat = 11 + getbits(b, 7);
395 if (index + repeat > nlen + ndist)
401 t->lengths[index++] = (short)repeat_len;
404 if (t->lengths[256] == 0)
409 int err = construct(&lc, t->lengths, nlen);
410 if (err && (err < 0 || nlen != lc.count[0] + lc.count[1]))
414 err = construct(&dc, t->lengths + nlen, ndist);
415 if (err && (err < 0 || ndist != dc.count[0] + dc.count[1]))
419 return do_codes(b, o, &lc, &dc);
423int do_block(BitIn *b, OutCtx *o, Tables *t)
426 int type = getbits(b, 2);
433 return do_stored(b, o);
437 return do_fixed(b, o, t);
441 return do_dynamic(b, o, t);
447void ssh_inflate_init(SshInflate *z, uint8_t *window)
454 z->header_seen =
false;
457int ssh_inflate_packet(SshInflate *z,
const uint8_t *src,
size_t src_len, uint8_t *dst,
size_t dst_cap,
size_t *out_len)
459 if (!z || (src_len && !src) || !out_len)
469 b.nbits = ((size_t)z->carry_len + src_len) * 8u;
470 b.bitpos = z->bit_off;
484 if (b.nbits - b.bitpos < 16u)
487 size_t rem = (size_t)z->carry_len + src_len;
488 if (rem > SSH_INFLATE_CARRY)
492 uint8_t tmp[SSH_INFLATE_CARRY];
493 for (
size_t i = 0; i < rem; i++)
495 tmp[i] = logical_byte(z->carry, z->carry_len, src, i);
497 memcpy(z->carry, tmp, rem);
498 z->carry_len = (uint8_t)rem;
503 int cmf = getbits(&b, 8);
504 int flg = getbits(&b, 8);
505 if ((cmf & 0x0F) != 8)
509 if ((((
unsigned)cmf << 8) | (
unsigned)flg) % 31u != 0u)
517 z->header_seen =
true;
521 size_t boundary = b.bitpos;
525 if (b.bitpos >= b.nbits)
529 size_t cp_bit = b.bitpos;
530 size_t cp_cnt = o.cnt;
531 uint32_t cp_wpos = z->wpos;
532 uint32_t cp_whist = z->whist;
534 int st = do_block(&b, &o, &tbl);
540 if (st == PC_BLK_NEED)
554 size_t bstart_byte = boundary >> 3;
555 size_t total_bytes = (size_t)z->carry_len + src_len;
556 size_t rem = total_bytes - bstart_byte;
557 if (rem > SSH_INFLATE_CARRY)
561 uint8_t tmp[SSH_INFLATE_CARRY];
562 for (
size_t i = 0; i < rem; i++)
564 tmp[i] = logical_byte(z->carry, z->carry_len, src, bstart_byte + i);
566 memcpy(z->carry, tmp, rem);
567 z->carry_len = (uint8_t)rem;
568 z->bit_off = (uint8_t)(boundary & 7u);
SSH client-to-server decompression: a resumable, context-takeover INFLATE (no heap).