[PATCH v2 05/14] smb: common: compress: use a single function for chained and unchained compression

Enzo Matsumiya <[email protected]> Mon, 20 Jul 2026 16:49:16 -0300
Newsgroups org.kernel.vger.linux-cifs
Message-ID <[email protected]>
Rename smb_compression_compress_chained() to smb_compression_compress(),
and handle both unchained and chained compressions there,
similar to smb_compression_decompress().

Changes:
- add @chained arg to smb_compression_compress(), so it's possible to
  handle (chained && !allow_pattern) cases
- drop smb_compression_builder::first field -- add all payload headers
  with Flags = "NONE", set flag "CHAINED" at the end on the starting
  header
- drop @orig_size from smb_compression_add_payload(), compute it
  directly on smb_compression_add_lz77() instead
- adjust smb_compression_compress() call in server/compress.c
  - TODO: rework the code around to fully use the single function

Signed-off-by: Enzo Matsumiya <[email protected]>
---
 fs/smb/common/compress/compress.c | 119 +++++++++++++++++++-----------
 fs/smb/common/compress/compress.h |   6 +-
 fs/smb/server/compress.c          |  11 +--
 3 files changed, 83 insertions(+), 53 deletions(-)

diff --git a/fs/smb/common/compress/compress.c b/fs/smb/common/compress/compress.c
index 088e0d0d2792..75d21622c268 100644
--- a/fs/smb/common/compress/compress.c
+++ b/fs/smb/common/compress/compress.c
@@ -213,7 +213,6 @@ EXPORT_SYMBOL_GPL(smb_compression_decompress);
 struct smb_compression_builder {
 	u8 *pos;
 	u32 remaining;
-	bool first;
 };
 
 /*
@@ -222,28 +221,22 @@ struct smb_compression_builder {
  */
 static struct smb2_compression_payload_hdr *
 smb_compression_add_payload(struct smb_compression_builder *builder,
-			    __le16 alg, u32 payload_len, bool orig_size)
+			    __le16 alg, u32 payload_len)
 {
 	struct smb2_compression_payload_hdr *payload;
-	u32 hdr_len = SMB2_COMPRESSION_PAYLOAD_BASE_LEN;
 	u32 total_len;
 
-	if (orig_size)
-		hdr_len += sizeof(payload->OriginalPayloadSize);
-	if (check_add_overflow(hdr_len, payload_len, &total_len) ||
+	if (check_add_overflow(payload_len, SMB2_COMPRESSION_PAYLOAD_BASE_LEN, &total_len) ||
 	    total_len > builder->remaining)
 		return NULL;
 
 	payload = (struct smb2_compression_payload_hdr *)builder->pos;
 	payload->CompressionAlgorithm = alg;
-	payload->Flags = cpu_to_le16(builder->first ?
-		SMB2_COMPRESSION_FLAG_CHAINED : SMB2_COMPRESSION_FLAG_NONE);
-	payload->Length = cpu_to_le32(payload_len +
-		(orig_size ? sizeof(payload->OriginalPayloadSize) : 0));
-
-	builder->pos += hdr_len;
-	builder->remaining -= hdr_len;
-	builder->first = false;
+	payload->Flags = cpu_to_le16(SMB2_COMPRESSION_FLAG_NONE);
+	payload->Length = cpu_to_le32(payload_len);
+
+	builder->pos += SMB2_COMPRESSION_PAYLOAD_BASE_LEN;
+	builder->remaining -= SMB2_COMPRESSION_PAYLOAD_BASE_LEN;
 	return payload;
 }
 
@@ -252,8 +245,7 @@ static int smb_compression_add_pattern(struct smb_compression_builder *builder,
 {
 	struct smb2_compression_pattern_v1 *payload;
 
-	if (!smb_compression_add_payload(builder, SMB3_COMPRESS_PATTERN,
-					 sizeof(*payload), false))
+	if (!smb_compression_add_payload(builder, SMB3_COMPRESS_PATTERN, sizeof(*payload)))
 		return -ENOSPC;
 
 	payload = (struct smb2_compression_pattern_v1 *)builder->pos;
@@ -269,7 +261,7 @@ static int smb_compression_add_pattern(struct smb_compression_builder *builder,
 static int smb_compression_add_none(struct smb_compression_builder *builder,
 				    const u8 *src, u32 len)
 {
-	if (!smb_compression_add_payload(builder, SMB3_COMPRESS_NONE, len, false))
+	if (!smb_compression_add_payload(builder, SMB3_COMPRESS_NONE, len))
 		return -ENOSPC;
 
 	memcpy(builder->pos, src, len);
@@ -279,36 +271,47 @@ static int smb_compression_add_none(struct smb_compression_builder *builder,
 }
 
 static int smb_compression_add_lz77(struct smb_compression_builder *builder,
-				    const u8 *src, u32 len)
+				    const u8 *src, u32 len, bool chained)
 {
 	struct smb2_compression_payload_hdr *payload;
-	u32 comp_len;
+	u32 comp_len, offset = chained ? sizeof(payload->OriginalPayloadSize) : 0;
 	int rc;
 
 	if (builder->remaining <= sizeof(*payload))
 		return -ENOSPC;
 
-	comp_len = builder->remaining - sizeof(*payload);
-	payload = smb_compression_add_payload(builder, SMB3_COMPRESS_LZ77,
-					      comp_len, true);
+	/*
+	 * Always add LZ data as a payload, regardless of @chained.
+	 * The only difference is whether offset (sizeof(payload->OriginalPayloadSize)) is
+	 * accounted for.
+	 * Also, use @payload_len == 0 as we don't know compressed size yet.
+	 */
+	payload = smb_compression_add_payload(builder, SMB3_COMPRESS_LZ77, 0);
 	if (!payload)
 		return -ENOSPC;
 
+	builder->pos += offset;
+	builder->remaining -= offset;
+	comp_len = builder->remaining;
 	rc = smb_lz77_compress(src, len, builder->pos, &comp_len);
 	if (rc)
 		return rc;
 
-	payload->Length = cpu_to_le32(comp_len +
-				      sizeof(payload->OriginalPayloadSize));
-	payload->OriginalPayloadSize = cpu_to_le32(len);
+	/* These are only set when @chained. */
+	if (chained) {
+		payload->Length = cpu_to_le32(comp_len + offset);
+		payload->OriginalPayloadSize = cpu_to_le32(len);
+	}
+
 	builder->pos += comp_len;
 	builder->remaining -= comp_len;
 	return 0;
 }
 
 /**
- * smb_compression_compress_chained() - build a chained SMB2 transform
+ * smb_compression_compress() - build a SMB2 compression transform (chained optional)
  * @alg: negotiated general-purpose compression algorithm
+ * @chained: is chained compression was negotiated
  * @allow_pattern: whether Pattern_V1 was negotiated
  * @src: complete uncompressed SMB2 message
  * @slen: size of @src
@@ -325,12 +328,12 @@ static int smb_compression_add_lz77(struct smb_compression_builder *builder,
  *
  * Return: 0 on success, otherwise a negative errno.
  */
-int smb_compression_compress_chained(__le16 alg, bool allow_pattern,
-				     const void *src, u32 slen,
-				     void *dst, u32 *dlen)
+int smb_compression_compress(__le16 alg, bool chained, bool allow_pattern,
+			     const void *src, u32 slen,
+			     void *dst, u32 *dlen)
 {
-	struct smb2_compression_hdr *hdr = dst;
 	struct smb_compression_builder builder;
+	struct smb2_compression_hdr *hdr;
 	const u8 *input = src;
 	u32 forward = 0, backward = 0, middle_len;
 	int rc;
@@ -339,13 +342,21 @@ int smb_compression_compress_chained(__le16 alg, bool allow_pattern,
 	    *dlen <= SMB2_COMPRESSION_CHAINED_HDR_LEN || !slen)
 		return -EINVAL;
 
-	hdr->ProtocolId = SMB2_COMPRESSION_TRANSFORM_ID;
-	hdr->OriginalCompressedSegmentSize = cpu_to_le32(slen);
+	/* Note that the below is a bug, but (chained && !allow_pattern) is a valid combination */
+	if (WARN_ON_ONCE(!chained && allow_pattern))
+		return -EINVAL;
+
+	/*
+	 * Offset @dst for both unchained and chained cases, header is setup at the end.
+	 * If unchained, we must align @dst to the end of compression header.
+	 */
 	builder.pos = (u8 *)dst + SMB2_COMPRESSION_CHAINED_HDR_LEN;
 	builder.remaining = *dlen - SMB2_COMPRESSION_CHAINED_HDR_LEN;
-	builder.first = true;
 
-	if (allow_pattern && slen > 32) {
+	if (!chained || !allow_pattern)
+		goto do_lz;
+
+	if (slen > 32) {
 		for (forward = 1; forward < slen; forward++) {
 			if (input[forward] != input[0])
 				break;
@@ -366,18 +377,22 @@ int smb_compression_compress_chained(__le16 alg, bool allow_pattern,
 		if (rc)
 			return rc;
 	}
-
+do_lz:
+	rc = -ENODATA;
 	middle_len = slen - forward - backward;
-	if (middle_len > 1024)
-		rc = smb_compression_add_lz77(&builder, input + forward,
-					      middle_len);
-	else if (middle_len)
-		rc = smb_compression_add_none(&builder,
-					      input + forward, middle_len);
-	else
-		rc = 0;
-	if (rc)
+	if (middle_len > 1024 || !chained)
+		rc = smb_compression_add_lz77(&builder, input + forward, middle_len, chained);
+	else if (middle_len && allow_pattern)
+		rc = smb_compression_add_none(&builder, input + forward, middle_len);
+
+	if (rc) {
+		/* These will allow callers to send uncompressed request in case of small input */
+		if (rc == -ENODATA) {
+			*dlen = slen;
+			rc = 0;
+		}
 		return rc;
+	}
 
 	if (backward) {
 		rc = smb_compression_add_pattern(&builder, input[slen - 1],
@@ -387,6 +402,20 @@ int smb_compression_compress_chained(__le16 alg, bool allow_pattern,
 	}
 
 	*dlen = builder.pos - (u8 *)dst;
+
+	hdr = dst;
+	hdr->ProtocolId = SMB2_COMPRESSION_TRANSFORM_ID;
+	hdr->OriginalCompressedSegmentSize = cpu_to_le32(slen);
+
+	if (chained) {
+		/* Overwrite the first payload->Flags, the other fields must stay as is. */
+		hdr->Flags = cpu_to_le16(SMB2_COMPRESSION_FLAG_CHAINED);
+	} else {
+		hdr->CompressionAlgorithm = alg;
+		hdr->Flags = cpu_to_le16(SMB2_COMPRESSION_FLAG_NONE);
+		hdr->Offset = cpu_to_le32(0);
+	}
+
 	return 0;
 }
-EXPORT_SYMBOL_GPL(smb_compression_compress_chained);
+EXPORT_SYMBOL_GPL(smb_compression_compress);
diff --git a/fs/smb/common/compress/compress.h b/fs/smb/common/compress/compress.h
index 24b558167fe7..ade37d552067 100644
--- a/fs/smb/common/compress/compress.h
+++ b/fs/smb/common/compress/compress.h
@@ -147,8 +147,8 @@ static __always_inline bool smb_compress_alg_valid(__le16 alg, bool valid_none)
 
 int smb_compression_decompress(__le16 alg, bool allow_chained,
 			       const void *src, u32 slen, void *dst, u32 dlen);
-int smb_compression_compress_chained(__le16 alg, bool allow_pattern,
-				     const void *src, u32 slen,
-				     void *dst, u32 *dlen);
+int smb_compression_compress(__le16 alg, bool chained, bool allow_pattern,
+			     const void *src, u32 slen,
+			     void *dst, u32 *dlen);
 
 #endif /* _COMMON_SMB_COMPRESS_H */
diff --git a/fs/smb/server/compress.c b/fs/smb/server/compress.c
index 8c910f996e04..da7dea58c5a7 100644
--- a/fs/smb/server/compress.c
+++ b/fs/smb/server/compress.c
@@ -145,11 +145,12 @@ int ksmbd_compress_response(struct ksmbd_work *work)
 
 	if (work->conn->compress_chained) {
 		dst_len = max_dst_len;
-		rc = smb_compression_compress_chained(SMB3_COMPRESS_LZ77,
-						      work->conn->compress_pattern,
-						      src, src_len,
-						      out + sizeof(__be32),
-						      &dst_len);
+		rc = smb_compression_compress(SMB3_COMPRESS_LZ77,
+					      work->conn->compress_chained,
+					      work->conn->compress_pattern,
+					      src, src_len,
+					      out + sizeof(__be32),
+					      &dst_len);
 		if (rc == -EMSGSIZE || dst_len >= src_len) {
 			rc = 0;
 			goto out;
-- 
2.54.0