xref: /linux/fs/smb/common/compress/lz77.c (revision 0eaed89c18aeedf0898baf2dbf5ff027c6795152)
1*0121b154SNamjae Jeon // SPDX-License-Identifier: GPL-2.0-only
2*0121b154SNamjae Jeon /*
3*0121b154SNamjae Jeon  * Copyright (C) 2024-2026, SUSE LLC
4*0121b154SNamjae Jeon  * Copyright (C) 2026 Namjae Jeon <linkinjeon@kernel.org>
5*0121b154SNamjae Jeon  *
6*0121b154SNamjae Jeon  * Authors: Enzo Matsumiya <ematsumiya@suse.de>
7*0121b154SNamjae Jeon  *          Namjae Jeon <linkinjeon@kernel.org>
8*0121b154SNamjae Jeon  *
9*0121b154SNamjae Jeon  * Implementation of the LZ77 "plain" compression algorithm, as per MS-XCA spec.
10*0121b154SNamjae Jeon  */
11*0121b154SNamjae Jeon #include <linux/slab.h>
12*0121b154SNamjae Jeon #include <linux/sizes.h>
13*0121b154SNamjae Jeon #include <linux/count_zeros.h>
14*0121b154SNamjae Jeon #include <linux/unaligned.h>
15*0121b154SNamjae Jeon #include <linux/module.h>
16*0121b154SNamjae Jeon #include <linux/overflow.h>
17*0121b154SNamjae Jeon 
18*0121b154SNamjae Jeon #include "lz77.h"
19*0121b154SNamjae Jeon 
20*0121b154SNamjae Jeon /*
21*0121b154SNamjae Jeon  * Compression parameters.
22*0121b154SNamjae Jeon  *
23*0121b154SNamjae Jeon  * LZ77_MATCH_MAX_DIST:		Farthest back a match can be from current position (can be 1 - 8K).
24*0121b154SNamjae Jeon  * LZ77_HASH_LOG:
25*0121b154SNamjae Jeon  * LZ77_HASH_SIZE:		ilog2 hash size (recommended to be 13 - 18, default 15 (hash size
26*0121b154SNamjae Jeon  *				32k)).
27*0121b154SNamjae Jeon  * LZ77_RSTEP_SIZE:		Number of bytes to read from input buffer for hashing and initial
28*0121b154SNamjae Jeon  *				match check (default 4 bytes, this effectivelly makes this the min
29*0121b154SNamjae Jeon  *				match len).
30*0121b154SNamjae Jeon  * LZ77_MSTEP_SIZE:		Number of bytes to extend-compare a found match (default 8 bytes).
31*0121b154SNamjae Jeon  * LZ77_SKIP_TRIGGER:		ilog2 value for adaptive skipping, i.e. to progressively skip input
32*0121b154SNamjae Jeon  *				bytes when we can't find matches.  Default is 4.
33*0121b154SNamjae Jeon  *				Higher values (>0) will decrease compression time, but will result
34*0121b154SNamjae Jeon  *				in worse compression ratio.  Lower values will give better
35*0121b154SNamjae Jeon  *				compression ratio (more matches found), but will increase time.
36*0121b154SNamjae Jeon  */
37*0121b154SNamjae Jeon #define LZ77_MATCH_MAX_DIST	SZ_8K
38*0121b154SNamjae Jeon #define LZ77_HASH_LOG		15
39*0121b154SNamjae Jeon #define LZ77_HASH_SIZE		BIT(LZ77_HASH_LOG)
40*0121b154SNamjae Jeon #define LZ77_RSTEP_SIZE		sizeof(u32)
41*0121b154SNamjae Jeon #define LZ77_MSTEP_SIZE		sizeof(u64)
42*0121b154SNamjae Jeon #define LZ77_SKIP_TRIGGER	4
43*0121b154SNamjae Jeon 
44*0121b154SNamjae Jeon #define LZ77_PREFETCH(ptr)	__builtin_prefetch((ptr), 0, 3)
45*0121b154SNamjae Jeon #define LZ77_FLAG_MAX		32
46*0121b154SNamjae Jeon 
47*0121b154SNamjae Jeon static __always_inline u8 lz77_read8(const u8 *ptr)
48*0121b154SNamjae Jeon {
49*0121b154SNamjae Jeon 	return get_unaligned(ptr);
50*0121b154SNamjae Jeon }
51*0121b154SNamjae Jeon 
52*0121b154SNamjae Jeon static __always_inline u32 lz77_read32(const u32 *ptr)
53*0121b154SNamjae Jeon {
54*0121b154SNamjae Jeon 	return get_unaligned(ptr);
55*0121b154SNamjae Jeon }
56*0121b154SNamjae Jeon 
57*0121b154SNamjae Jeon static __always_inline u64 lz77_read64(const u64 *ptr)
58*0121b154SNamjae Jeon {
59*0121b154SNamjae Jeon 	return get_unaligned(ptr);
60*0121b154SNamjae Jeon }
61*0121b154SNamjae Jeon 
62*0121b154SNamjae Jeon static __always_inline void lz77_write8(u8 *ptr, u8 v)
63*0121b154SNamjae Jeon {
64*0121b154SNamjae Jeon 	put_unaligned(v, ptr);
65*0121b154SNamjae Jeon }
66*0121b154SNamjae Jeon 
67*0121b154SNamjae Jeon static __always_inline void lz77_write16(u16 *ptr, u16 v)
68*0121b154SNamjae Jeon {
69*0121b154SNamjae Jeon 	put_unaligned_le16(v, ptr);
70*0121b154SNamjae Jeon }
71*0121b154SNamjae Jeon 
72*0121b154SNamjae Jeon static __always_inline void lz77_write32(u32 *ptr, u32 v)
73*0121b154SNamjae Jeon {
74*0121b154SNamjae Jeon 	put_unaligned_le32(v, ptr);
75*0121b154SNamjae Jeon }
76*0121b154SNamjae Jeon 
77*0121b154SNamjae Jeon static __always_inline u32 lz77_match_len(const void *match, const void *cur, const void *end)
78*0121b154SNamjae Jeon {
79*0121b154SNamjae Jeon 	const void *start = cur;
80*0121b154SNamjae Jeon 
81*0121b154SNamjae Jeon 	/* Safe for a do/while because otherwise we wouldn't reach here from the main loop. */
82*0121b154SNamjae Jeon 	do {
83*0121b154SNamjae Jeon 		const u64 diff = lz77_read64(cur) ^ lz77_read64(match);
84*0121b154SNamjae Jeon 
85*0121b154SNamjae Jeon 		if (!diff) {
86*0121b154SNamjae Jeon 			cur += LZ77_MSTEP_SIZE;
87*0121b154SNamjae Jeon 			match += LZ77_MSTEP_SIZE;
88*0121b154SNamjae Jeon 
89*0121b154SNamjae Jeon 			continue;
90*0121b154SNamjae Jeon 		}
91*0121b154SNamjae Jeon 
92*0121b154SNamjae Jeon 		/* This computes the number of common bytes in @diff. */
93*0121b154SNamjae Jeon 		cur += count_trailing_zeros(diff) >> 3;
94*0121b154SNamjae Jeon 
95*0121b154SNamjae Jeon 		return (cur - start);
96*0121b154SNamjae Jeon 	} while (likely(cur + LZ77_MSTEP_SIZE <= end));
97*0121b154SNamjae Jeon 
98*0121b154SNamjae Jeon 	/* Fallback to byte-by-byte comparison for last <8 bytes. */
99*0121b154SNamjae Jeon 	while (cur < end && lz77_read8(cur) == lz77_read8(match)) {
100*0121b154SNamjae Jeon 		cur++;
101*0121b154SNamjae Jeon 		match++;
102*0121b154SNamjae Jeon 	}
103*0121b154SNamjae Jeon 
104*0121b154SNamjae Jeon 	return (cur - start);
105*0121b154SNamjae Jeon }
106*0121b154SNamjae Jeon 
107*0121b154SNamjae Jeon /**
108*0121b154SNamjae Jeon  * lz77_encode_match() - Match encoding.
109*0121b154SNamjae Jeon  * @dst:	compressed buffer
110*0121b154SNamjae Jeon  * @nib:	pointer to an address in @dst
111*0121b154SNamjae Jeon  * @dist:	match distance
112*0121b154SNamjae Jeon  * @len:	match length
113*0121b154SNamjae Jeon  *
114*0121b154SNamjae Jeon  * Assumes all args were previously checked.
115*0121b154SNamjae Jeon  *
116*0121b154SNamjae Jeon  * Return: @dst advanced to new position
117*0121b154SNamjae Jeon  *
118*0121b154SNamjae Jeon  * Ref: MS-XCA 2.3.4 "Plain LZ77 Compression Algorithm Details" - "Processing"
119*0121b154SNamjae Jeon  */
120*0121b154SNamjae Jeon static __always_inline void *lz77_encode_match(void *dst, void **nib, u16 dist, u32 len)
121*0121b154SNamjae Jeon {
122*0121b154SNamjae Jeon 	len -= 3;
123*0121b154SNamjae Jeon 	dist--;
124*0121b154SNamjae Jeon 	dist <<= 3;
125*0121b154SNamjae Jeon 
126*0121b154SNamjae Jeon 	if (len < 7) {
127*0121b154SNamjae Jeon 		lz77_write16(dst, dist + len);
128*0121b154SNamjae Jeon 
129*0121b154SNamjae Jeon 		return dst + sizeof(u16);
130*0121b154SNamjae Jeon 	}
131*0121b154SNamjae Jeon 
132*0121b154SNamjae Jeon 	dist |= 7;
133*0121b154SNamjae Jeon 	lz77_write16(dst, dist);
134*0121b154SNamjae Jeon 	dst += sizeof(u16);
135*0121b154SNamjae Jeon 	len -= 7;
136*0121b154SNamjae Jeon 
137*0121b154SNamjae Jeon 	if (!*nib) {
138*0121b154SNamjae Jeon 		lz77_write8(dst, umin(len, 15));
139*0121b154SNamjae Jeon 		*nib = dst;
140*0121b154SNamjae Jeon 		dst++;
141*0121b154SNamjae Jeon 	} else {
142*0121b154SNamjae Jeon 		u8 *b = *nib;
143*0121b154SNamjae Jeon 
144*0121b154SNamjae Jeon 		lz77_write8(b, *b | umin(len, 15) << 4);
145*0121b154SNamjae Jeon 		*nib = NULL;
146*0121b154SNamjae Jeon 	}
147*0121b154SNamjae Jeon 
148*0121b154SNamjae Jeon 	if (len < 15)
149*0121b154SNamjae Jeon 		return dst;
150*0121b154SNamjae Jeon 
151*0121b154SNamjae Jeon 	len -= 15;
152*0121b154SNamjae Jeon 	if (len < 255) {
153*0121b154SNamjae Jeon 		lz77_write8(dst, len);
154*0121b154SNamjae Jeon 
155*0121b154SNamjae Jeon 		return dst + 1;
156*0121b154SNamjae Jeon 	}
157*0121b154SNamjae Jeon 
158*0121b154SNamjae Jeon 	lz77_write8(dst, 0xff);
159*0121b154SNamjae Jeon 	dst++;
160*0121b154SNamjae Jeon 	len += 7 + 15;
161*0121b154SNamjae Jeon 	if (len <= 0xffff) {
162*0121b154SNamjae Jeon 		lz77_write16(dst, len);
163*0121b154SNamjae Jeon 
164*0121b154SNamjae Jeon 		return dst + sizeof(u16);
165*0121b154SNamjae Jeon 	}
166*0121b154SNamjae Jeon 
167*0121b154SNamjae Jeon 	lz77_write16(dst, 0);
168*0121b154SNamjae Jeon 	dst += sizeof(u16);
169*0121b154SNamjae Jeon 	lz77_write32(dst, len);
170*0121b154SNamjae Jeon 
171*0121b154SNamjae Jeon 	return dst + sizeof(u32);
172*0121b154SNamjae Jeon }
173*0121b154SNamjae Jeon 
174*0121b154SNamjae Jeon /**
175*0121b154SNamjae Jeon  * lz77_encode_literals() - Literals encoding.
176*0121b154SNamjae Jeon  * @start:	where to start copying literals (uncompressed buffer)
177*0121b154SNamjae Jeon  * @end:	when to stop copying (uncompressed buffer)
178*0121b154SNamjae Jeon  * @dst:	compressed buffer
179*0121b154SNamjae Jeon  * @f:		pointer to current flag value
180*0121b154SNamjae Jeon  * @fc:		pointer to current flag count
181*0121b154SNamjae Jeon  * @fp:		pointer to current flag address
182*0121b154SNamjae Jeon  *
183*0121b154SNamjae Jeon  * Batch copy literals from @start to @dst, updating flag values accordingly.
184*0121b154SNamjae Jeon  * Assumes all args were previously checked.
185*0121b154SNamjae Jeon  *
186*0121b154SNamjae Jeon  * Return: @dst advanced to new position
187*0121b154SNamjae Jeon  *
188*0121b154SNamjae Jeon  * MS-XCA 2.3.4 "Plain LZ77 Compression Algorithm Details" - "Processing"
189*0121b154SNamjae Jeon  */
190*0121b154SNamjae Jeon static __always_inline void *lz77_encode_literals(const void *start, const void *end, void *dst,
191*0121b154SNamjae Jeon 						  long *f, u32 *fc, void **fp)
192*0121b154SNamjae Jeon {
193*0121b154SNamjae Jeon 	if (start >= end)
194*0121b154SNamjae Jeon 		return dst;
195*0121b154SNamjae Jeon 
196*0121b154SNamjae Jeon 	do {
197*0121b154SNamjae Jeon 		const u32 len = umin(end - start, LZ77_FLAG_MAX - *fc);
198*0121b154SNamjae Jeon 
199*0121b154SNamjae Jeon 		memcpy(dst, start, len);
200*0121b154SNamjae Jeon 
201*0121b154SNamjae Jeon 		dst += len;
202*0121b154SNamjae Jeon 		start += len;
203*0121b154SNamjae Jeon 
204*0121b154SNamjae Jeon 		*f <<= len;
205*0121b154SNamjae Jeon 		*fc += len;
206*0121b154SNamjae Jeon 		if (*fc == LZ77_FLAG_MAX) {
207*0121b154SNamjae Jeon 			lz77_write32(*fp, *f);
208*0121b154SNamjae Jeon 			*fc = 0;
209*0121b154SNamjae Jeon 			*fp = dst;
210*0121b154SNamjae Jeon 			dst += sizeof(u32);
211*0121b154SNamjae Jeon 		}
212*0121b154SNamjae Jeon 	} while (start < end);
213*0121b154SNamjae Jeon 
214*0121b154SNamjae Jeon 	return dst;
215*0121b154SNamjae Jeon }
216*0121b154SNamjae Jeon 
217*0121b154SNamjae Jeon static __always_inline u32 lz77_hash(const u32 v)
218*0121b154SNamjae Jeon {
219*0121b154SNamjae Jeon 	return ((v ^ 0x9E3779B9) * 0x85EBCA6B) >> (32 - LZ77_HASH_LOG);
220*0121b154SNamjae Jeon }
221*0121b154SNamjae Jeon 
222*0121b154SNamjae Jeon noinline int smb_lz77_compress(const void *src, const u32 slen,
223*0121b154SNamjae Jeon 			       void *dst, u32 *dlen)
224*0121b154SNamjae Jeon {
225*0121b154SNamjae Jeon 	const void *srcp, *rlim, *end, *anchor;
226*0121b154SNamjae Jeon 	u32 *htable, hash, flag_count = 0;
227*0121b154SNamjae Jeon 	void *dstp, *nib, *flag_pos;
228*0121b154SNamjae Jeon 	long flag = 0;
229*0121b154SNamjae Jeon 
230*0121b154SNamjae Jeon 	/* This is probably a bug, so throw a warning. */
231*0121b154SNamjae Jeon 	if (WARN_ON_ONCE(*dlen < smb_lz77_compressed_alloc_size(slen)))
232*0121b154SNamjae Jeon 		return -EINVAL;
233*0121b154SNamjae Jeon 
234*0121b154SNamjae Jeon 	srcp = src;
235*0121b154SNamjae Jeon 	anchor = src;
236*0121b154SNamjae Jeon 	end = srcp + slen; /* absolute end */
237*0121b154SNamjae Jeon 	rlim = end - LZ77_MSTEP_SIZE; /* read limit (for lz77_match_len()) */
238*0121b154SNamjae Jeon 	dstp = dst;
239*0121b154SNamjae Jeon 	flag_pos = dstp;
240*0121b154SNamjae Jeon 	dstp += sizeof(u32);
241*0121b154SNamjae Jeon 	nib = NULL;
242*0121b154SNamjae Jeon 
243*0121b154SNamjae Jeon 	htable = kvcalloc(LZ77_HASH_SIZE, sizeof(*htable), GFP_KERNEL);
244*0121b154SNamjae Jeon 	if (!htable)
245*0121b154SNamjae Jeon 		return -ENOMEM;
246*0121b154SNamjae Jeon 
247*0121b154SNamjae Jeon 	LZ77_PREFETCH(srcp + LZ77_RSTEP_SIZE);
248*0121b154SNamjae Jeon 
249*0121b154SNamjae Jeon 	/*
250*0121b154SNamjae Jeon 	 * Adjust @srcp so we don't get a false positive match on first iteration.
251*0121b154SNamjae Jeon 	 * Then prepare hash for first loop iteration (don't advance @srcp again).
252*0121b154SNamjae Jeon 	 */
253*0121b154SNamjae Jeon 	hash = lz77_hash(lz77_read32(srcp++));
254*0121b154SNamjae Jeon 	htable[hash] = 0;
255*0121b154SNamjae Jeon 	hash = lz77_hash(lz77_read32(srcp));
256*0121b154SNamjae Jeon 
257*0121b154SNamjae Jeon 	/*
258*0121b154SNamjae Jeon 	 * Main loop.
259*0121b154SNamjae Jeon 	 *
260*0121b154SNamjae Jeon 	 * @dlen is >= smb_lz77_compressed_alloc_size(), so run without
261*0121b154SNamjae Jeon 	 * bound-checking @dstp.
262*0121b154SNamjae Jeon 	 *
263*0121b154SNamjae Jeon 	 * This code was crafted in a way to best utilise fetch-decode-execute CPU flow.
264*0121b154SNamjae Jeon 	 * Any attempt to optimize it, or even organize it, can lead to huge performance loss.
265*0121b154SNamjae Jeon 	 */
266*0121b154SNamjae Jeon 	do {
267*0121b154SNamjae Jeon 		const void *match, *next = srcp;
268*0121b154SNamjae Jeon 		u32 len, step = 1, skip = 1U << LZ77_SKIP_TRIGGER;
269*0121b154SNamjae Jeon 
270*0121b154SNamjae Jeon 		/* Match finding (hot path -- don't change the read/check/write order). */
271*0121b154SNamjae Jeon 		do {
272*0121b154SNamjae Jeon 			const u32 cur_hash = hash;
273*0121b154SNamjae Jeon 
274*0121b154SNamjae Jeon 			srcp = next;
275*0121b154SNamjae Jeon 			next += step;
276*0121b154SNamjae Jeon 
277*0121b154SNamjae Jeon 			/*
278*0121b154SNamjae Jeon 			 * Adaptive skipping.
279*0121b154SNamjae Jeon 			 *
280*0121b154SNamjae Jeon 			 * Increment @step every (1 << LZ77_SKIP_TRIGGER, 16 in our case) bytes
281*0121b154SNamjae Jeon 			 * without a match.
282*0121b154SNamjae Jeon 			 * Reset to 1 when a match is found.
283*0121b154SNamjae Jeon 			 */
284*0121b154SNamjae Jeon 			step = (skip++ >> LZ77_SKIP_TRIGGER);
285*0121b154SNamjae Jeon 			if (unlikely(next > rlim))
286*0121b154SNamjae Jeon 				goto out;
287*0121b154SNamjae Jeon 
288*0121b154SNamjae Jeon 			hash = lz77_hash(lz77_read32(next));
289*0121b154SNamjae Jeon 			match = src + htable[cur_hash];
290*0121b154SNamjae Jeon 			htable[cur_hash] = srcp - src;
291*0121b154SNamjae Jeon 		} while (likely(match + LZ77_MATCH_MAX_DIST < srcp) ||
292*0121b154SNamjae Jeon 			 lz77_read32(match) != lz77_read32(srcp));
293*0121b154SNamjae Jeon 
294*0121b154SNamjae Jeon 		/*
295*0121b154SNamjae Jeon 		 * Match found.  Warm/cold path; begin parsing @srcp and writing to @dstp:
296*0121b154SNamjae Jeon 		 * - flush literals
297*0121b154SNamjae Jeon 		 * - compute match length (*)
298*0121b154SNamjae Jeon 		 * - encode match
299*0121b154SNamjae Jeon 		 *
300*0121b154SNamjae Jeon 		 * (*) Current minimum match length is defined by the memory read size above, so
301*0121b154SNamjae Jeon 		 * here we already know that we have 4 matching bytes, but it's just faster to
302*0121b154SNamjae Jeon 		 * redundantly compute it again in lz77_match_len() than to adjust pointers/len.
303*0121b154SNamjae Jeon 		 */
304*0121b154SNamjae Jeon 		dstp = lz77_encode_literals(anchor, srcp, dstp, &flag, &flag_count, &flag_pos);
305*0121b154SNamjae Jeon 		len = lz77_match_len(match, srcp, end);
306*0121b154SNamjae Jeon 		dstp = lz77_encode_match(dstp, &nib, srcp - match, len);
307*0121b154SNamjae Jeon 		srcp += len;
308*0121b154SNamjae Jeon 		anchor = srcp;
309*0121b154SNamjae Jeon 
310*0121b154SNamjae Jeon 		LZ77_PREFETCH(srcp);
311*0121b154SNamjae Jeon 
312*0121b154SNamjae Jeon 		flag = (flag << 1) | 1;
313*0121b154SNamjae Jeon 		flag_count++;
314*0121b154SNamjae Jeon 		if (flag_count == LZ77_FLAG_MAX) {
315*0121b154SNamjae Jeon 			lz77_write32(flag_pos, flag);
316*0121b154SNamjae Jeon 			flag_count = 0;
317*0121b154SNamjae Jeon 			flag_pos = dstp;
318*0121b154SNamjae Jeon 			dstp += sizeof(u32);
319*0121b154SNamjae Jeon 		}
320*0121b154SNamjae Jeon 
321*0121b154SNamjae Jeon 		if (unlikely(srcp > rlim))
322*0121b154SNamjae Jeon 			break;
323*0121b154SNamjae Jeon 
324*0121b154SNamjae Jeon 		/* Prepare for next loop. */
325*0121b154SNamjae Jeon 		hash = lz77_hash(lz77_read32(srcp));
326*0121b154SNamjae Jeon 	} while (srcp < end);
327*0121b154SNamjae Jeon out:
328*0121b154SNamjae Jeon 	dstp = lz77_encode_literals(anchor, end, dstp, &flag, &flag_count, &flag_pos);
329*0121b154SNamjae Jeon 
330*0121b154SNamjae Jeon 	flag_count = LZ77_FLAG_MAX - flag_count;
331*0121b154SNamjae Jeon 	flag <<= flag_count;
332*0121b154SNamjae Jeon 	flag |= (1UL << flag_count) - 1;
333*0121b154SNamjae Jeon 	lz77_write32(flag_pos, flag);
334*0121b154SNamjae Jeon 
335*0121b154SNamjae Jeon 	*dlen = dstp - dst;
336*0121b154SNamjae Jeon 	kvfree(htable);
337*0121b154SNamjae Jeon 
338*0121b154SNamjae Jeon 	if (*dlen < slen)
339*0121b154SNamjae Jeon 		return 0;
340*0121b154SNamjae Jeon 
341*0121b154SNamjae Jeon 	return -EMSGSIZE;
342*0121b154SNamjae Jeon }
343*0121b154SNamjae Jeon EXPORT_SYMBOL_GPL(smb_lz77_compress);
344*0121b154SNamjae Jeon 
345*0121b154SNamjae Jeon static int lz77_decode_match_len(const u8 **src, const u8 *end, u16 token,
346*0121b154SNamjae Jeon 				 u8 *nibble, bool *have_nibble, u32 *len)
347*0121b154SNamjae Jeon {
348*0121b154SNamjae Jeon 	u8 extra;
349*0121b154SNamjae Jeon 
350*0121b154SNamjae Jeon 	*len = (token & 0x7) + 3;
351*0121b154SNamjae Jeon 	if ((token & 0x7) != 0x7)
352*0121b154SNamjae Jeon 		return 0;
353*0121b154SNamjae Jeon 
354*0121b154SNamjae Jeon 	if (!*have_nibble) {
355*0121b154SNamjae Jeon 		if (*src >= end)
356*0121b154SNamjae Jeon 			return -EINVAL;
357*0121b154SNamjae Jeon 		*nibble = *(*src)++;
358*0121b154SNamjae Jeon 		extra = *nibble & 0xf;
359*0121b154SNamjae Jeon 		*have_nibble = true;
360*0121b154SNamjae Jeon 	} else {
361*0121b154SNamjae Jeon 		extra = *nibble >> 4;
362*0121b154SNamjae Jeon 		*have_nibble = false;
363*0121b154SNamjae Jeon 	}
364*0121b154SNamjae Jeon 
365*0121b154SNamjae Jeon 	*len += extra;
366*0121b154SNamjae Jeon 	if (extra == 0xf) {
367*0121b154SNamjae Jeon 		u8 b;
368*0121b154SNamjae Jeon 
369*0121b154SNamjae Jeon 		if (*src >= end)
370*0121b154SNamjae Jeon 			return -EINVAL;
371*0121b154SNamjae Jeon 		b = *(*src)++;
372*0121b154SNamjae Jeon 		if (b != 0xff) {
373*0121b154SNamjae Jeon 			*len += b;
374*0121b154SNamjae Jeon 		} else {
375*0121b154SNamjae Jeon 			u16 w;
376*0121b154SNamjae Jeon 
377*0121b154SNamjae Jeon 			if (end - *src < 2)
378*0121b154SNamjae Jeon 				return -EINVAL;
379*0121b154SNamjae Jeon 			w = get_unaligned_le16(*src);
380*0121b154SNamjae Jeon 			*src += 2;
381*0121b154SNamjae Jeon 			if (w) {
382*0121b154SNamjae Jeon 				*len = w + 3;
383*0121b154SNamjae Jeon 			} else {
384*0121b154SNamjae Jeon 				u32 long_len;
385*0121b154SNamjae Jeon 
386*0121b154SNamjae Jeon 				if (end - *src < 4)
387*0121b154SNamjae Jeon 					return -EINVAL;
388*0121b154SNamjae Jeon 				long_len = get_unaligned_le32(*src);
389*0121b154SNamjae Jeon 				*src += 4;
390*0121b154SNamjae Jeon 				if (check_add_overflow(long_len, 3, len))
391*0121b154SNamjae Jeon 					return -EINVAL;
392*0121b154SNamjae Jeon 			}
393*0121b154SNamjae Jeon 		}
394*0121b154SNamjae Jeon 	}
395*0121b154SNamjae Jeon 
396*0121b154SNamjae Jeon 	return 0;
397*0121b154SNamjae Jeon }
398*0121b154SNamjae Jeon 
399*0121b154SNamjae Jeon int smb_lz77_decompress(const void *src, const u32 slen, void *dst,
400*0121b154SNamjae Jeon 			const u32 dlen)
401*0121b154SNamjae Jeon {
402*0121b154SNamjae Jeon 	const u8 *sp = src, *send = sp + slen;
403*0121b154SNamjae Jeon 	u8 *dp = dst, *dend = dp + dlen;
404*0121b154SNamjae Jeon 	u32 flags = 0;
405*0121b154SNamjae Jeon 	int flag_count = 0;
406*0121b154SNamjae Jeon 	u8 nibble = 0;
407*0121b154SNamjae Jeon 	bool have_nibble = false;
408*0121b154SNamjae Jeon 
409*0121b154SNamjae Jeon 	while (dp < dend) {
410*0121b154SNamjae Jeon 		u32 len, dist;
411*0121b154SNamjae Jeon 		u16 token;
412*0121b154SNamjae Jeon 
413*0121b154SNamjae Jeon 		if (!flag_count) {
414*0121b154SNamjae Jeon 			if (send - sp < 4)
415*0121b154SNamjae Jeon 				return -EINVAL;
416*0121b154SNamjae Jeon 			flags = get_unaligned_le32(sp);
417*0121b154SNamjae Jeon 			sp += 4;
418*0121b154SNamjae Jeon 			flag_count = 32;
419*0121b154SNamjae Jeon 		}
420*0121b154SNamjae Jeon 
421*0121b154SNamjae Jeon 		if (!(flags & 0x80000000)) {
422*0121b154SNamjae Jeon 			if (sp >= send)
423*0121b154SNamjae Jeon 				return -EINVAL;
424*0121b154SNamjae Jeon 			*dp++ = *sp++;
425*0121b154SNamjae Jeon 			flags <<= 1;
426*0121b154SNamjae Jeon 			flag_count--;
427*0121b154SNamjae Jeon 			continue;
428*0121b154SNamjae Jeon 		}
429*0121b154SNamjae Jeon 
430*0121b154SNamjae Jeon 		flags <<= 1;
431*0121b154SNamjae Jeon 		flag_count--;
432*0121b154SNamjae Jeon 
433*0121b154SNamjae Jeon 		if (send - sp < 2)
434*0121b154SNamjae Jeon 			return -EINVAL;
435*0121b154SNamjae Jeon 
436*0121b154SNamjae Jeon 		token = get_unaligned_le16(sp);
437*0121b154SNamjae Jeon 		sp += 2;
438*0121b154SNamjae Jeon 
439*0121b154SNamjae Jeon 		dist = (token >> 3) + 1;
440*0121b154SNamjae Jeon 		if (dist > dp - (u8 *)dst)
441*0121b154SNamjae Jeon 			return -EINVAL;
442*0121b154SNamjae Jeon 
443*0121b154SNamjae Jeon 		if (lz77_decode_match_len(&sp, send, token, &nibble,
444*0121b154SNamjae Jeon 					  &have_nibble, &len))
445*0121b154SNamjae Jeon 			return -EINVAL;
446*0121b154SNamjae Jeon 
447*0121b154SNamjae Jeon 		if (len > dend - dp)
448*0121b154SNamjae Jeon 			return -EINVAL;
449*0121b154SNamjae Jeon 
450*0121b154SNamjae Jeon 		while (len--) {
451*0121b154SNamjae Jeon 			*dp = *(dp - dist);
452*0121b154SNamjae Jeon 			dp++;
453*0121b154SNamjae Jeon 		}
454*0121b154SNamjae Jeon 	}
455*0121b154SNamjae Jeon 
456*0121b154SNamjae Jeon 	return 0;
457*0121b154SNamjae Jeon }
458*0121b154SNamjae Jeon EXPORT_SYMBOL_GPL(smb_lz77_decompress);
459*0121b154SNamjae Jeon 
460*0121b154SNamjae Jeon MODULE_LICENSE("GPL");
461*0121b154SNamjae Jeon MODULE_DESCRIPTION("SMB plain LZ77 compression");
462