xref: /linux/fs/smb/common/compress/compress.c (revision fab183d632628381b466a41479489541ac0e29a0)
1 // SPDX-License-Identifier: GPL-2.0-or-later
2 /*
3  * SMB2 compression transform helpers.
4  *
5  * Copyright (C) 2026 Namjae Jeon <linkinjeon@kernel.org>
6  */
7 #include <linux/module.h>
8 #include <linux/overflow.h>
9 #include <linux/string.h>
10 #include <linux/unaligned.h>
11 
12 #include "compress.h"
13 #include "lz77.h"
14 
15 #define SMB2_COMPRESSION_CHAINED_HDR_LEN \
16 	offsetof(struct smb2_compression_hdr, CompressionAlgorithm)
17 #define SMB2_COMPRESSION_PAYLOAD_BASE_LEN \
18 	(sizeof(struct smb2_compression_payload_hdr) - sizeof(__le32))
19 
20 /*
21  * A NONE payload carries bytes verbatim. Keep both cursors and remaining
22  * lengths together so every chained payload handler applies identical bounds
23  * accounting.
24  */
smb_decompress_none(const u8 ** src,u32 * slen,u8 ** dst,u32 * dlen,u32 len)25 static int smb_decompress_none(const u8 **src, u32 *slen, u8 **dst, u32 *dlen,
26 			       u32 len)
27 {
28 	if (len > *slen || len > *dlen)
29 		return -EINVAL;
30 
31 	memcpy(*dst, *src, len);
32 	*src += len;
33 	*slen -= len;
34 	*dst += len;
35 	*dlen -= len;
36 	return 0;
37 }
38 
39 /*
40  * Pattern_V1 represents a run of one byte. Its wire payload is always the
41  * fixed-size smb2_compression_pattern_v1 structure.
42  */
smb_decompress_pattern(const u8 ** src,u32 * slen,u8 ** dst,u32 * dlen,u32 len)43 static int smb_decompress_pattern(const u8 **src, u32 *slen, u8 **dst,
44 				  u32 *dlen, u32 len)
45 {
46 	const struct smb2_compression_pattern_v1 *pattern;
47 	u32 repetitions;
48 
49 	if (len != sizeof(*pattern) || len > *slen)
50 		return -EINVAL;
51 
52 	pattern = (const struct smb2_compression_pattern_v1 *)*src;
53 	repetitions = le32_to_cpu(pattern->Repetitions);
54 	if (repetitions > *dlen)
55 		return -EINVAL;
56 
57 	memset(*dst, pattern->Pattern, repetitions);
58 	*src += len;
59 	*slen -= len;
60 	*dst += repetitions;
61 	*dlen -= repetitions;
62 	return 0;
63 }
64 
65 /*
66  * LZ77 payload Length includes the four-byte OriginalPayloadSize field.
67  * Consume that field before passing the compressed stream to the raw codec.
68  */
smb_decompress_lz77_payload(const u8 ** src,u32 * slen,u8 ** dst,u32 * dlen,u32 len)69 static int smb_decompress_lz77_payload(const u8 **src, u32 *slen, u8 **dst,
70 				       u32 *dlen, u32 len)
71 {
72 	u32 orig_size;
73 	int rc;
74 
75 	if (len < sizeof(__le32) || len > *slen)
76 		return -EINVAL;
77 
78 	orig_size = get_unaligned_le32(*src);
79 	if (orig_size > *dlen)
80 		return -EINVAL;
81 
82 	*src += sizeof(__le32);
83 	*slen -= sizeof(__le32);
84 	len -= sizeof(__le32);
85 
86 	rc = smb_lz77_decompress(*src, len, *dst, orig_size);
87 	if (rc)
88 		return rc;
89 
90 	*src += len;
91 	*slen -= len;
92 	*dst += orig_size;
93 	*dlen -= orig_size;
94 	return 0;
95 }
96 
smb_decompress_chained(__le16 alg,bool allow_chained,bool allow_pattern,const struct smb2_compression_hdr * hdr,u32 slen,void * dst,u32 dlen)97 static int smb_decompress_chained(__le16 alg, bool allow_chained,
98 				  bool allow_pattern,
99 				  const struct smb2_compression_hdr *hdr,
100 				  u32 slen, void *dst, u32 dlen)
101 {
102 	const struct smb2_compression_payload_hdr *payload;
103 	const u8 *src = (const u8 *)hdr + SMB2_COMPRESSION_CHAINED_HDR_LEN;
104 	u32 orig_size = le32_to_cpu(hdr->OriginalCompressedSegmentSize);
105 	u32 remaining = slen - SMB2_COMPRESSION_CHAINED_HDR_LEN;
106 	u8 *out = dst;
107 	u32 out_remaining = dlen;
108 	bool first = true;
109 	int rc;
110 
111 	if (!allow_chained || orig_size != dlen)
112 		return -EINVAL;
113 
114 	/*
115 	 * The chained transform has an eight-byte top-level header. The next
116 	 * bytes are a sequence of payload headers whose Length fields account
117 	 * for payload data, including OriginalPayloadSize where applicable.
118 	 */
119 	while (remaining) {
120 		__le16 payload_alg;
121 		__le16 flags;
122 		u32 len;
123 
124 		if (remaining < SMB2_COMPRESSION_PAYLOAD_BASE_LEN)
125 			return -EINVAL;
126 
127 		payload = (const struct smb2_compression_payload_hdr *)src;
128 		payload_alg = payload->CompressionAlgorithm;
129 		flags = payload->Flags;
130 		len = le32_to_cpu(payload->Length);
131 
132 		/*
133 		 * CHAINED marks only the first payload. Requiring NONE on every
134 		 * later payload rejects ambiguous or independently chained data.
135 		 */
136 		if ((first && flags != cpu_to_le16(SMB2_COMPRESSION_FLAG_CHAINED)) ||
137 		    (!first && flags != cpu_to_le16(SMB2_COMPRESSION_FLAG_NONE)))
138 			return -EINVAL;
139 
140 		src += SMB2_COMPRESSION_PAYLOAD_BASE_LEN;
141 		remaining -= SMB2_COMPRESSION_PAYLOAD_BASE_LEN;
142 
143 		if (payload_alg == SMB3_COMPRESS_NONE) {
144 			rc = smb_decompress_none(&src, &remaining, &out,
145 						 &out_remaining, len);
146 		} else if (payload_alg == SMB3_COMPRESS_PATTERN) {
147 			if (!allow_pattern)
148 				return -EINVAL;
149 			rc = smb_decompress_pattern(&src, &remaining, &out,
150 						    &out_remaining, len);
151 		} else if (payload_alg == alg && alg == SMB3_COMPRESS_LZ77) {
152 			rc = smb_decompress_lz77_payload(&src, &remaining, &out,
153 							 &out_remaining, len);
154 		} else {
155 			return -EINVAL;
156 		}
157 		if (rc)
158 			return rc;
159 		first = false;
160 	}
161 
162 	return out_remaining ? -EINVAL : 0;
163 }
164 
smb_decompress_unchained(__le16 alg,const struct smb2_compression_hdr * hdr,u32 slen,void * dst,u32 dlen)165 static int smb_decompress_unchained(__le16 alg,
166 				    const struct smb2_compression_hdr *hdr,
167 				    u32 slen, void *dst, u32 dlen)
168 {
169 	u32 orig_size, offset, comp_size;
170 
171 	if (hdr->CompressionAlgorithm != alg ||
172 	    !smb_compress_alg_valid(hdr->CompressionAlgorithm, false))
173 		return -EINVAL;
174 
175 	orig_size = le32_to_cpu(hdr->OriginalCompressedSegmentSize);
176 	offset = le32_to_cpu(hdr->Offset);
177 	if (offset > slen - sizeof(*hdr) || offset > dlen ||
178 	    orig_size > dlen - offset || orig_size + offset != dlen)
179 		return -EINVAL;
180 
181 	memcpy(dst, (const u8 *)hdr + sizeof(*hdr), offset);
182 	comp_size = slen - sizeof(*hdr) - offset;
183 	return smb_lz77_decompress((const u8 *)hdr + sizeof(*hdr) + offset,
184 				   comp_size, (u8 *)dst + offset, orig_size);
185 }
186 
187 /**
188  * smb_compression_decompress() - decode an SMB2 compression transform
189  * @alg: negotiated general-purpose compression algorithm
190  * @allow_chained: whether chained transforms were negotiated
191  * @allow_pattern: whether Pattern_V1 payloads were negotiated
192  * @src: transform header followed by compressed payload data
193  * @slen: total number of bytes available at @src
194  * @dst: output buffer for the reconstructed SMB2 message
195  * @dlen: exact expected size of the reconstructed SMB2 message
196  *
197  * Validate the transform type and negotiated capabilities before dispatching
198  * to the chained or unchained decoder. The caller supplies the expected output
199  * size after applying its transport-specific message size limits.
200  *
201  * Return: 0 on success, otherwise a negative errno.
202  */
smb_compression_decompress(__le16 alg,bool allow_chained,bool allow_pattern,const void * src,u32 slen,void * dst,u32 dlen)203 int smb_compression_decompress(__le16 alg, bool allow_chained,
204 			       bool allow_pattern, const void *src, u32 slen,
205 			       void *dst, u32 dlen)
206 {
207 	const struct smb2_compression_hdr *hdr = src;
208 
209 	if (!src || !dst || slen < sizeof(*hdr) ||
210 	    hdr->ProtocolId != SMB2_COMPRESSION_TRANSFORM_ID ||
211 	    alg == SMB3_COMPRESS_NONE)
212 		return -EINVAL;
213 
214 	if (hdr->Flags == cpu_to_le16(SMB2_COMPRESSION_FLAG_CHAINED))
215 		return smb_decompress_chained(alg, allow_chained, allow_pattern,
216 					      hdr, slen, dst, dlen);
217 
218 	if (hdr->Flags != cpu_to_le16(SMB2_COMPRESSION_FLAG_NONE))
219 		return -EINVAL;
220 
221 	return smb_decompress_unchained(alg, hdr, slen, dst, dlen);
222 }
223 EXPORT_SYMBOL_GPL(smb_compression_decompress);
224 
225 struct smb_compression_builder {
226 	u8 *pos;
227 	u32 remaining;
228 	bool first;
229 };
230 
231 /*
232  * Reserve one chained payload header and initialize its common fields.
233  * OriginalPayloadSize is present only for LZNT1/LZ77/LZ77+Huffman payloads.
234  */
235 static struct smb2_compression_payload_hdr *
smb_compression_add_payload(struct smb_compression_builder * builder,__le16 alg,u32 payload_len,bool orig_size)236 smb_compression_add_payload(struct smb_compression_builder *builder,
237 			    __le16 alg, u32 payload_len, bool orig_size)
238 {
239 	struct smb2_compression_payload_hdr *payload;
240 	u32 hdr_len = SMB2_COMPRESSION_PAYLOAD_BASE_LEN;
241 	u32 total_len;
242 
243 	if (orig_size)
244 		hdr_len += sizeof(payload->OriginalPayloadSize);
245 	if (check_add_overflow(hdr_len, payload_len, &total_len) ||
246 	    total_len > builder->remaining)
247 		return NULL;
248 
249 	payload = (struct smb2_compression_payload_hdr *)builder->pos;
250 	payload->CompressionAlgorithm = alg;
251 	payload->Flags = cpu_to_le16(builder->first ?
252 		SMB2_COMPRESSION_FLAG_CHAINED : SMB2_COMPRESSION_FLAG_NONE);
253 	payload->Length = cpu_to_le32(payload_len +
254 		(orig_size ? sizeof(payload->OriginalPayloadSize) : 0));
255 
256 	builder->pos += hdr_len;
257 	builder->remaining -= hdr_len;
258 	builder->first = false;
259 	return payload;
260 }
261 
smb_compression_add_pattern(struct smb_compression_builder * builder,u8 pattern,u32 repetitions)262 static int smb_compression_add_pattern(struct smb_compression_builder *builder,
263 				       u8 pattern, u32 repetitions)
264 {
265 	struct smb2_compression_pattern_v1 *payload;
266 
267 	if (!smb_compression_add_payload(builder, SMB3_COMPRESS_PATTERN,
268 					 sizeof(*payload), false))
269 		return -ENOSPC;
270 
271 	payload = (struct smb2_compression_pattern_v1 *)builder->pos;
272 	payload->Pattern = pattern;
273 	payload->Reserved1 = 0;
274 	payload->Reserved2 = 0;
275 	payload->Repetitions = cpu_to_le32(repetitions);
276 	builder->pos += sizeof(*payload);
277 	builder->remaining -= sizeof(*payload);
278 	return 0;
279 }
280 
smb_compression_add_none(struct smb_compression_builder * builder,const u8 * src,u32 len)281 static int smb_compression_add_none(struct smb_compression_builder *builder,
282 				    const u8 *src, u32 len)
283 {
284 	if (!smb_compression_add_payload(builder, SMB3_COMPRESS_NONE, len, false))
285 		return -ENOSPC;
286 
287 	memcpy(builder->pos, src, len);
288 	builder->pos += len;
289 	builder->remaining -= len;
290 	return 0;
291 }
292 
smb_compression_add_lz77(struct smb_compression_builder * builder,const u8 * src,u32 len)293 static int smb_compression_add_lz77(struct smb_compression_builder *builder,
294 				    const u8 *src, u32 len)
295 {
296 	struct smb2_compression_payload_hdr *payload;
297 	u32 comp_len;
298 	int rc;
299 
300 	if (builder->remaining <= sizeof(*payload))
301 		return -ENOSPC;
302 
303 	comp_len = builder->remaining - sizeof(*payload);
304 	payload = smb_compression_add_payload(builder, SMB3_COMPRESS_LZ77,
305 					      comp_len, true);
306 	if (!payload)
307 		return -ENOSPC;
308 
309 	rc = smb_lz77_compress(src, len, builder->pos, &comp_len);
310 	if (rc)
311 		return rc;
312 
313 	payload->Length = cpu_to_le32(comp_len +
314 				      sizeof(payload->OriginalPayloadSize));
315 	payload->OriginalPayloadSize = cpu_to_le32(len);
316 	builder->pos += comp_len;
317 	builder->remaining -= comp_len;
318 	return 0;
319 }
320 
321 /**
322  * smb_compression_compress_chained() - build a chained SMB2 transform
323  * @alg: negotiated general-purpose compression algorithm
324  * @allow_pattern: whether Pattern_V1 was negotiated
325  * @src: complete uncompressed SMB2 message
326  * @slen: size of @src
327  * @dst: output buffer for the transform
328  * @dlen: input capacity of @dst and output transform size
329  *
330  * Following the algorithm in [MS-SMB2] 3.1.4.4, encode sufficiently long
331  * repeated runs at the front and back as Pattern_V1 payloads. Compress a
332  * middle region larger than 1 KiB with LZ77; smaller middle regions are
333  * represented by a chained NONE payload.
334  *
335  * This helper does not decide whether the final transform is smaller than the
336  * original message. The transport caller owns that policy decision.
337  *
338  * Return: 0 on success, otherwise a negative errno.
339  */
smb_compression_compress_chained(__le16 alg,bool allow_pattern,const void * src,u32 slen,void * dst,u32 * dlen)340 int smb_compression_compress_chained(__le16 alg, bool allow_pattern,
341 				     const void *src, u32 slen,
342 				     void *dst, u32 *dlen)
343 {
344 	struct smb2_compression_hdr *hdr = dst;
345 	struct smb_compression_builder builder;
346 	const u8 *input = src;
347 	u32 forward = 0, backward = 0, middle_len;
348 	int rc;
349 
350 	if (!src || !dst || !dlen || alg != SMB3_COMPRESS_LZ77 ||
351 	    *dlen <= SMB2_COMPRESSION_CHAINED_HDR_LEN || !slen)
352 		return -EINVAL;
353 
354 	hdr->ProtocolId = SMB2_COMPRESSION_TRANSFORM_ID;
355 	hdr->OriginalCompressedSegmentSize = cpu_to_le32(slen);
356 	builder.pos = (u8 *)dst + SMB2_COMPRESSION_CHAINED_HDR_LEN;
357 	builder.remaining = *dlen - SMB2_COMPRESSION_CHAINED_HDR_LEN;
358 	builder.first = true;
359 
360 	if (allow_pattern && slen > 32) {
361 		for (forward = 1; forward < slen; forward++) {
362 			if (input[forward] != input[0])
363 				break;
364 		}
365 		if (forward <= 32)
366 			forward = 0;
367 
368 		for (backward = 1; backward < slen - forward; backward++) {
369 			if (input[slen - backward - 1] != input[slen - 1])
370 				break;
371 		}
372 		if (backward <= 32)
373 			backward = 0;
374 	}
375 
376 	if (forward) {
377 		rc = smb_compression_add_pattern(&builder, input[0], forward);
378 		if (rc)
379 			return rc;
380 	}
381 
382 	middle_len = slen - forward - backward;
383 	if (middle_len > 1024)
384 		rc = smb_compression_add_lz77(&builder, input + forward,
385 					      middle_len);
386 	else if (middle_len)
387 		rc = smb_compression_add_none(&builder,
388 					      input + forward, middle_len);
389 	else
390 		rc = 0;
391 	if (rc)
392 		return rc;
393 
394 	if (backward) {
395 		rc = smb_compression_add_pattern(&builder, input[slen - 1],
396 						 backward);
397 		if (rc)
398 			return rc;
399 	}
400 
401 	*dlen = builder.pos - (u8 *)dst;
402 	return 0;
403 }
404 EXPORT_SYMBOL_GPL(smb_compression_compress_chained);
405