gzip/deflatelua.lua

1
--[[
2
LICENSE
3
 
4
Copyright (C) 2008, David Manura.
5
Modifications (C) 2010, Matthew Wild <mwild1@gmail.com>
6
 
7
Permission is hereby granted, free of charge, to any person obtaining a copy
8
of this software and associated documentation files (the "Software"), to deal
9
in the Software without restriction, including without limitation the rights
10
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
11
copies of the Software, and to permit persons to whom the Software is
12
furnished to do so, subject to the following conditions:
13
 
14
The above copyright notice and this permission notice shall be included in
15
all copies or substantial portions of the Software.
16
 
17
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
18
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
19
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT.  IN NO EVENT SHALL THE
20
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
21
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
22
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN
23
THE SOFTWARE.
24
 
25
(end license)
26
--]]
27
 
28
-- dmlib.deflate
29
-- deflate (and gunzip) implemented in Lua.
30
--
31
-- Note: only supports decompression.
32
-- Compression not implemented.
33
--
34
-- References
35
-- [1] DEFLATE Compressed Data Format Specification version 1.3
36
--     http://tools.ietf.org/html/rfc1951
37
-- [2] GZIP file format specification version 4.3
38
--     http://tools.ietf.org/html/rfc1952
39
-- [3] http://en.wikipedia.org/wiki/DEFLATE
40
-- [4] pyflate, by Paul Sladen
41
--     http://www.paul.sladen.org/projects/pyflate/
42
-- [5] Compress::Zlib::Perl - partial pure Perl implementation of
43
--     Compress::Zlib
44
--     http://search.cpan.org/~nwclark/Compress-Zlib-Perl/Perl.pm
45
--
46
-- (c) 2008 David Manura.  Licensed under the same terms as Lua (MIT).
47
 
48
local assert, error, ipairs, pairs, tostring, type, setmetatable, io, math, table_sort, 
49
	math_max, string_char, io_open, _G =
50
      assert, error, ipairs, pairs, tostring, type, setmetatable, io, math, table.sort, 
51
	math.max, string.char, io.open, _G;
52
 
53
local function memoize(f)
54
  local mt = {};
55
  local t = setmetatable({}, mt)
56
  function mt:__index(k)
57
    local v = f(k); t[k] = v
58
    return v
59
  end
60
  return t
61
end
62
 
63
local function runtime_error(s, level)
64
  level = level or 1
65
  error({s}, level+1)
66
end
67
 
68
 
69
local function make_os(outbs)
70
  local os = {}
71
  os.outbs = outbs
72
  os.wnd = {}
73
  os.wnd_pos = 1
74
  return os
75
end
76
 
77
 
78
local function output(os, byte)
79
  -- debug('OUTPUT:', s)
80
  local wnd_pos = os.wnd_pos
81
  os.outbs(byte)
82
  os.wnd[wnd_pos] = byte
83
  os.wnd_pos = wnd_pos % 32768 + 1  -- 32K
84
end
85
 
86
 
87
local function noeof(val)
88
  return assert(val, 'unexpected end of file')
89
end
90
 
91
 
92
local function hasbit(bits, bit)
93
  return bits % (bit + bit) >= bit
94
end
95
 
96
 
97
-- small optimization (lookup table for powers of 2)
98
local pow2 = memoize(function(n) return 2^n end)
99
--local tbits = memoize(
100
--  function(bits)
101
--    return memoize( function(bit) return getbit(bits, bit) end )
102
--  end )
103
 
104
 
105
-- weak metatable marking objects as bitstream type
106
local is_bitstream = setmetatable({}, {__mode='k'})
107
 
108
 
109
 
110
local function bytestream_from_string(s)
111
  local i = 1
112
  local o = {}
113
  function o:read()
114
    local by
115
    if i <= #s then
116
      by = s:byte(i)
117
      i = i + 1
118
    end
119
    return by
120
  end
121
  return o
122
end
123
 
124
local left
125
local function bitstream_from_bytestream(bys)
126
  local buf_byte, buf_nbit, o = 0, 0, {};
127
 
128
  function o:nbits_left_in_byte()
129
    return buf_nbit
130
  end
131
 
132
  function o:read(nbits)
133
    nbits = nbits or 1
134
    while buf_nbit < nbits do
135
      local byte = bys:read()
136
      if not byte then return end  -- note: more calls also return nil
137
      buf_byte = buf_byte + pow2[buf_nbit] * byte
138
      buf_nbit = buf_nbit + 8
139
    end
140
    local m = pow2[nbits]
141
    local bits = buf_byte % m
142
    buf_byte = (buf_byte - bits) / m
143
    buf_nbit = buf_nbit - nbits
144
    return bits
145
  end
146
 
147
  is_bitstream[o] = true
148
 
149
  return o
150
end
151
 
152
 
153
local function get_bitstream(o)
154
    return is_bitstream[o] and o or bitstream_from_bytestream(bytestream_from_string(o))
155
end
156
 
157
 
158
local function get_obytestream(o)
159
  local bs
160
  if io.type(o) == 'file' then
161
    bs = function(sbyte) o:write(string_char(sbyte)) end
162
  elseif type(o) == 'function' then
163
    bs = o
164
  end
165
  return bs
166
end
167
 
168
 
169
local function HuffmanTable(init, is_full)
170
  local t = {}
171
  if is_full then
172
    for val,nbits in pairs(init) do
173
      if nbits ~= 0 then
174
        t[#t+1] = {val=val, nbits=nbits}
175
        --debug('*',val,nbits)
176
      end
177
    end
178
  else
179
    for i=1,#init-2,2 do
180
      local firstval, nbits, nextval = init[i], init[i+1], init[i+2]
181
      --debug(val, nextval, nbits)
182
      if nbits ~= 0 then
183
        for val=firstval,nextval-1 do
184
          t[#t+1] = {val=val, nbits=nbits}
185
        end
186
      end
187
    end
188
  end
189
  table_sort(t, function(a,b)
190
    return a.nbits == b.nbits and a.val < b.val or a.nbits < b.nbits
191
  end)
192
 
193
  -- assign codes
194
  local code = 1  -- leading 1 marker
195
  local nbits = 0
196
  for i,s in ipairs(t) do
197
    if s.nbits ~= nbits then
198
      code = code * pow2[s.nbits - nbits]
199
      nbits = s.nbits
200
    end
201
    s.code = code
202
    --debug('huffman code:', i, s.nbits, s.val, code, bits_tostring(code))
203
    code = code + 1
204
  end
205
 
206
  local minbits = math.huge
207
  local look = {}
208
  for i,s in ipairs(t) do
209
    minbits = math.min(minbits, s.nbits)
210
    look[s.code] = s.val
211
  end
212
 
213
  --for _,o in ipairs(t) do
214
  --  debug(':', o.nbits, o.val)
215
  --end
216
 
217
  -- function t:lookup(bits) return look[bits] end
218
 
219
  local function msb(bits, nbits)
220
    local res = 0
221
    for i=1,nbits do
222
      local b = bits % 2
223
      bits = (bits - b) / 2
224
      res = res * 2 + b
225
    end
226
    return res
227
  end
228
  local tfirstcode = memoize(
229
    function(bits) return pow2[minbits] + msb(bits, minbits) end)
230
 
231
  function t:read(bs)
232
    local code, nbits = 1, 0 -- leading 1 marker
233
    while 1 do
234
      if nbits == 0 then  -- small optimization (optional)
235
        code = tfirstcode[noeof(bs:read(minbits))]
236
        nbits = nbits + minbits
237
      else
238
        local b = noeof(bs:read())
239
        nbits = nbits + 1
240
        --debug('b',b)
241
        code = code * 2 + b   -- MSB first
242
      end
243
      --debug('code?', code, bits_tostring(code))
244
      local val = look[code]
245
      if val then
246
        --debug('FOUND', val)
247
        return val
248
      end
249
    end
250
  end
251
 
252
  return t
253
end
254
 
255
 
256
local function parse_gzip_header(bs)
257
  -- local FLG_FTEXT = 2^0
258
  local FLG_FHCRC = 2^1
259
  local FLG_FEXTRA = 2^2
260
  local FLG_FNAME = 2^3
261
  local FLG_FCOMMENT = 2^4
262
 
263
  local id1 = bs:read(8)
264
  local id2 = bs:read(8)
265
  local cm = bs:read(8)  -- compression method
266
  local flg = bs:read(8) -- FLaGs
267
  local mtime = bs:read(32) -- Modification TIME
268
  local xfl = bs:read(8) -- eXtra FLags
269
  local os = bs:read(8) -- Operating System
270
 
271
  if hasbit(flg, FLG_FEXTRA) then
272
    local xlen = bs:read(16)
273
    local extra = 0
274
    for i=1,xlen do
275
      extra = bs:read(8)
276
    end
277
  end
278
 
279
  if hasbit(flg, FLG_FNAME) then
280
      while bs:read(8) ~= 0 do end
281
  end
282
 
283
  if hasbit(flg, FLG_FCOMMENT) then
284
      while bs:read(8) ~= 0 do end
285
  end
286
  if hasbit(flg, FLG_FHCRC) then
287
    bs:read(16)
288
  end
289
end
290
 
291
 
292
local function parse_huffmantables(bs)
293
    local hlit = bs:read(5)  -- # of literal/length codes - 257
294
    local hdist = bs:read(5) -- # of distance codes - 1
295
    local hclen = noeof(bs:read(4)) -- # of code length codes - 4
296
 
297
    local ncodelen_codes = hclen + 4
298
    local codelen_init = {}
299
    local codelen_vals = {
300
      16, 17, 18, 0, 8, 7, 9, 6, 10, 5, 11, 4, 12, 3, 13, 2, 14, 1, 15}
301
    for i=1,ncodelen_codes do
302
      local nbits = bs:read(3)
303
      local val = codelen_vals[i]
304
      codelen_init[val] = nbits
305
    end
306
    local codelentable = HuffmanTable(codelen_init, true)
307
 
308
    local function decode(ncodes)
309
      local init = {}
310
      local nbits
311
      local val = 0
312
      while val < ncodes do
313
        local codelen = codelentable:read(bs)
314
        --FIX:check nil?
315
        local nrepeat
316
        if codelen <= 15 then
317
          nrepeat = 1
318
          nbits = codelen
319
          --debug('w', nbits)
320
        elseif codelen == 16 then
321
          nrepeat = 3 + noeof(bs:read(2))
322
          -- nbits unchanged
323
        elseif codelen == 17 then
324
          nrepeat = 3 + noeof(bs:read(3))
325
          nbits = 0
326
        elseif codelen == 18 then
327
          nrepeat = 11 + noeof(bs:read(7))
328
          nbits = 0
329
        else
330
          error 'ASSERT'
331
        end
332
        for i=1,nrepeat do
333
          init[val] = nbits
334
          val = val + 1
335
        end
336
      end
337
      local huffmantable = HuffmanTable(init, true)
338
      return huffmantable
339
    end
340
 
341
    local nlit_codes = hlit + 257
342
    local ndist_codes = hdist + 1
343
 
344
    local littable = decode(nlit_codes)
345
    local disttable = decode(ndist_codes)
346
 
347
    return littable, disttable
348
end
349
 
350
 
351
local tdecode_len_base
352
local tdecode_len_nextrabits
353
local tdecode_dist_base
354
local tdecode_dist_nextrabits
355
local function parse_compressed_item(bs, os, littable, disttable)
356
  local val = littable:read(bs)
357
  --debug(val, val < 256 and string_char(val))
358
  if val < 256 then -- literal
359
    output(os, val)
360
  elseif val == 256 then -- end of block
361
    return true
362
  else
363
    if not tdecode_len_base then
364
      local t = {[257]=3}
365
      local skip = 1
366
      for i=258,285,4 do
367
        for j=i,i+3 do t[j] = t[j-1] + skip end
368
        if i ~= 258 then skip = skip * 2 end
369
      end
370
      t[285] = 258
371
      tdecode_len_base = t
372
      --for i=257,285 do debug('T1',i,t[i]) end
373
    end
374
    if not tdecode_len_nextrabits then
375
      local t = {}
376
      for i=257,285 do
377
        local j = math_max(i - 261, 0)
378
        t[i] = (j - (j % 4)) / 4
379
      end
380
      t[285] = 0
381
      tdecode_len_nextrabits = t
382
      --for i=257,285 do debug('T2',i,t[i]) end
383
    end
384
    local len_base = tdecode_len_base[val]
385
    local nextrabits = tdecode_len_nextrabits[val]
386
    local extrabits = bs:read(nextrabits)
387
    local len = len_base + extrabits
388
 
389
    if not tdecode_dist_base then
390
      local t = {[0]=1}
391
      local skip = 1
392
      for i=1,29,2 do
393
        for j=i,i+1 do t[j] = t[j-1] + skip end
394
        if i ~= 1 then skip = skip * 2 end
395
      end
396
      tdecode_dist_base = t
397
      --for i=0,29 do debug('T3',i,t[i]) end
398
    end
399
    if not tdecode_dist_nextrabits then
400
      local t = {}
401
      for i=0,29 do
402
        local j = math_max(i - 2, 0)
403
        t[i] = (j - (j % 2)) / 2
404
      end
405
      tdecode_dist_nextrabits = t
406
      --for i=0,29 do debug('T4',i,t[i]) end
407
    end
408
    local dist_val = disttable:read(bs)
409
    local dist_base = tdecode_dist_base[dist_val]
410
    local dist_nextrabits = tdecode_dist_nextrabits[dist_val]
411
    local dist_extrabits = bs:read(dist_nextrabits)
412
    local dist = dist_base + dist_extrabits
413
 
414
    --debug('BACK', len, dist)
415
    for i=1,len do
416
      local pos = (os.wnd_pos - 1 - dist) % 32768 + 1  -- 32K
417
      output(os, assert(os.wnd[pos], 'invalid distance'))
418
    end
419
  end
420
  return false
421
end
422
 
423
 
424
local function parse_block(bs, os)
425
  local bfinal = bs:read(1)
426
  local btype = bs:read(2)
427
 
428
  local BTYPE_NO_COMPRESSION = 0
429
  local BTYPE_FIXED_HUFFMAN = 1
430
  local BTYPE_DYNAMIC_HUFFMAN = 2
431
  local BTYPE_RESERVED = 3
432
 
433
  if btype == BTYPE_NO_COMPRESSION then
434
    bs:read(bs:nbits_left_in_byte())
435
    local len = bs:read(16)
436
    local nlen = noeof(bs:read(16))
437
 
438
    for i=1,len do
439
      local by = noeof(bs:read(8))
440
      output(os, by)
441
    end
442
  elseif btype == BTYPE_FIXED_HUFFMAN or btype == BTYPE_DYNAMIC_HUFFMAN then
443
    local littable, disttable
444
    if btype == BTYPE_DYNAMIC_HUFFMAN then
445
      littable, disttable = parse_huffmantables(bs)
446
    else
447
      littable  = HuffmanTable {0,8, 144,9, 256,7, 280,8, 288,nil}
448
      disttable = HuffmanTable {0,5, 32,nil}
449
    end
450
 
451
    repeat until parse_compressed_item(
452
        bs, os, littable, disttable
453
    );
454
  end
455
 
456
  return bfinal ~= 0
457
end
458
 
459
 
460
local function deflate(t)
461
  local bs, os = get_bitstream(t.input)
462
  	, make_os(get_obytestream(t.output))
463
  repeat until parse_block(bs, os)
464
end
465
 
466
return function (t)
467
  local bs = get_bitstream(t.input)
468
  local outbs = get_obytestream(t.output)
469
 
470
  parse_gzip_header(bs)
471
 
472
  deflate{input=bs, output=outbs}
473
 
474
  bs:read(bs:nbits_left_in_byte())
475
  bs:read()
476
end