[PATCH 09/10] smb: client: add support for decompressing READs

Enzo Matsumiya <[email protected]>
Newsgroups org.kernel.vger.linux-cifs
Message-ID <[email protected]>
Implement decompression support for SMB2 READ messages.

Changes:
- add smb2ops.c::receive_compressed()
- add check_decompress() helper in smb2ops.c
- smb/common/compress/compress.h: add is_compress_hdr() and
  decompressed_size() helpers
- check for compress header in smb3_is_transform_hdr() as well so we
  fall into the same path handling offloaded READ decryption

Signed-off-by: Enzo Matsumiya <[email protected]>
---
 fs/smb/client/smb2ops.c           | 145 +++++++++++++++++++++++++++++-
 fs/smb/client/smb2pdu.c           |   6 ++
 fs/smb/common/compress/compress.h |  27 ++++++
 3 files changed, 177 insertions(+), 1 deletion(-)

diff --git a/fs/smb/client/smb2ops.c b/fs/smb/client/smb2ops.c
index d60bc537ecf1..e05dbf02d3b7 100644
--- a/fs/smb/client/smb2ops.c
+++ b/fs/smb/client/smb2ops.c
@@ -30,6 +30,8 @@
 #include "fs_context.h"
 #include "cached_dir.h"
 #include "reparse.h"
+#include "../common/compress/compress.h"
+#include "compress.h"
 
 /* Change credits for different ops and return the total number of credits */
 static int
@@ -4714,7 +4716,7 @@ smb3_is_transform_hdr(void *buf)
 {
 	struct smb2_transform_hdr *trhdr = buf;
 
-	return trhdr->ProtocolId == SMB2_TRANSFORM_PROTO_NUM;
+	return (trhdr->ProtocolId == SMB2_TRANSFORM_PROTO_NUM) || is_compress_hdr(buf);
 }
 
 static int
@@ -5246,6 +5248,136 @@ receive_encrypted_standard(struct TCP_Server_Info *server,
 	return ret;
 }
 
+static int receive_compressed(struct TCP_Server_Info *server)
+{
+	struct mid_q_entry *mid;
+	void *src, *dst = NULL;
+	u32 slen, dlen;
+	int ret;
+
+	slen = server->pdu_size;
+	src = kvzalloc(slen, GFP_KERNEL);
+	if (unlikely(!src))
+		return -ENOMEM;
+
+	ret = server->total_read;
+	memcpy(src, server->smallbuf, ret);
+
+	ret = cifs_read_from_socket(server, src + ret, slen - ret);
+	if (ret < 0)
+		goto err_free;
+
+	server->total_read += ret;
+
+	dlen = decompressed_size(src);
+	dst = kvzalloc(dlen, GFP_KERNEL);
+	if (unlikely(!dst)) {
+		ret = -ENOMEM;
+
+		goto err_free;
+	}
+
+	ret = smb_compression_decompress(server->compression.alg, server->compression.pattern,
+					 src, slen, dst, dlen);
+	if (likely(!ret)) {
+		const struct smb2_hdr *shdr = dst;
+
+		if (unlikely(shdr->ProtocolId != SMB2_PROTO_NUMBER)) {
+			cifs_dbg(VFS, "decompressed message is not an SMB2 message (got ProtocolId 0x%x)\n",
+				 shdr->ProtocolId);
+			ret = -ECONNRESET;
+			goto err_free;
+		}
+	} else {
+		/*
+		 * We _must_ disconnect on any failed decompression.
+		 * Actually, there's not much we can do to handle different decompression errors
+		 * here, so forcing a reconnect (which will disable compression) is our best option
+		 * as the request should be retried.
+		 */
+		ret = -ECONNRESET;
+		goto err_free;
+	}
+
+	mid = smb2_find_dequeue_mid(server, dst);
+	if (!mid) {
+		ret = -EIO;
+
+		goto err_free;
+	}
+
+	ret = handle_read_data(server, mid, dst, dlen, NULL, 0, true);
+	if (!ret) {
+		struct cifs_io_subrequest *rdata = mid->callback_data;
+#ifdef CONFIG_CIFS_STATS2
+		mid->when_received = jiffies;
+#endif
+		if (server->ops->is_network_name_deleted)
+			server->ops->is_network_name_deleted(dst, server);
+
+		rdata->iov[0].iov_base = dst;
+		rdata->iov[0].iov_len = server->vals->read_rsp_size;
+		mid_execute_callback(server, mid);
+	} else {
+		spin_lock(&server->srv_lock);
+		if (server->tcpStatus == CifsNeedReconnect) {
+			spin_lock(&server->mid_queue_lock);
+			mid->mid_state = MID_RETRY_NEEDED;
+			spin_unlock(&server->mid_queue_lock);
+			spin_unlock(&server->srv_lock);
+
+			mid_execute_callback(server, mid);
+		} else {
+			spin_lock(&server->mid_queue_lock);
+			mid->mid_state = MID_REQUEST_SUBMITTED;
+			mid->deleted_from_q = false;
+			list_add_tail(&mid->qhead, &server->pending_mid_q);
+			spin_unlock(&server->mid_queue_lock);
+			spin_unlock(&server->srv_lock);
+		}
+	}
+
+	release_mid(server, mid);
+err_free:
+	kvfree(src);
+	kvfree(dst);
+
+	if (unlikely(ret == -ECONNRESET)) {
+		spin_lock(&server->srv_lock);
+		server->tcpStatus = CifsNeedReconnect;
+		spin_unlock(&server->srv_lock);
+	}
+
+	return ret;
+}
+
+static inline int check_decompress(struct TCP_Server_Info *server, const void *buf)
+{
+	const u32 len = decompressed_size(buf);
+	u32 max_len;
+
+	if (len == 0)
+		return 0;
+
+	if (WARN_ON_ONCE(!server))
+		return -EINVAL;
+
+	if (unlikely(len < server->vals->read_rsp_size)) {
+		cifs_server_dbg(VFS, "uncompressed message too small (%u, min %zu)\n", len,
+				server->vals->read_rsp_size);
+		return -EINVAL;
+	}
+
+	max_len = 256 + SMB2_COMPRESSION_PAYLOAD_BASE_LEN +
+		  max3(server->maxBuf, server->max_read, server->max_write);
+	if (unlikely(len > max_len)) {
+		cifs_server_dbg(VFS, "uncompressed message too big (%u, max %u)\n", len, max_len);
+		return -EINVAL;
+	}
+
+	return 1;
+}
+
 static int
 smb3_receive_transform(struct TCP_Server_Info *server,
 		       struct mid_q_entry **mids, char **bufs, int *num_mids)
@@ -5254,6 +5386,15 @@ smb3_receive_transform(struct TCP_Server_Info *server,
 	unsigned int pdu_length = server->pdu_size;
 	struct smb2_transform_hdr *tr_hdr = (struct smb2_transform_hdr *)buf;
 	unsigned int orig_len = le32_to_cpu(tr_hdr->OriginalMessageSize);
+	int ret;
+
+	ret = check_decompress(server, buf);
+	if (ret > 0)
+		return receive_compressed(server);
+
+	if (unlikely(ret < 0))
+		/* Reset connection on any error */
+		return -ECONNABORTED;
 
 	if (pdu_length < sizeof(struct smb2_transform_hdr) +
 						sizeof(struct smb2_hdr)) {
@@ -5296,6 +5437,8 @@ static int smb2_next_header(struct TCP_Server_Info *server, char *buf,
 		*noff = le32_to_cpu(t_hdr->OriginalMessageSize);
 		if (unlikely(check_add_overflow(*noff, sizeof(*t_hdr), noff)))
 			return -EINVAL;
+	} else if (hdr->ProtocolId == SMB2_COMPRESSION_TRANSFORM_ID) {
+		*noff = 0;
 	} else {
 		*noff = le32_to_cpu(hdr->NextCommand);
 	}
diff --git a/fs/smb/client/smb2pdu.c b/fs/smb/client/smb2pdu.c
index ba2c601f290e..a94f0c0cf8de 100644
--- a/fs/smb/client/smb2pdu.c
+++ b/fs/smb/client/smb2pdu.c
@@ -4880,6 +4880,12 @@ smb2_async_readv(struct cifs_io_subrequest *rdata)
 		flags |= CIFS_HAS_CREDITS;
 	}
 
+	if (should_compress(io_parms.tcon, &rqst)) {
+		struct smb2_read_req *req = (struct smb2_read_req *)buf;
+
+		req->Flags |= SMB2_READFLAG_REQUEST_COMPRESSED;
+	}
+
 	rc = cifs_call_async(server, &rqst,
 			     cifs_readv_receive, smb2_readv_callback,
 			     smb3_handle_read_data, rdata, flags,
diff --git a/fs/smb/common/compress/compress.h b/fs/smb/common/compress/compress.h
index 93be0ad732f2..adc8b11910cf 100644
--- a/fs/smb/common/compress/compress.h
+++ b/fs/smb/common/compress/compress.h
@@ -152,6 +152,33 @@ static __always_inline bool smb_compress_alg_valid(__le16 alg, bool valid_none)
 	return false;
 }
 
+static __always_inline bool is_compress_hdr(const void *buf)
+{
+	const struct smb2_compression_hdr *hdr = buf;
+
+	if (!buf)
+		return false;
+
+	return (hdr->ProtocolId == SMB2_COMPRESSION_TRANSFORM_ID);
+}
+
+static __always_inline u32 decompressed_size(const void *buf)
+{
+	const struct smb2_compression_hdr *hdr = buf;
+	u32 len;
+
+	if (!buf || !is_compress_hdr(buf))
+		return 0;
+
+	len = le32_to_cpu(hdr->OriginalCompressedSegmentSize);
+
+	/* When unchained, we must account for the offset part (usually SMB2 header) as well. */
+	if (le16_to_cpu(hdr->Flags) == SMB2_COMPRESSION_FLAG_NONE)
+		len += le32_to_cpu(hdr->Offset);
+
+	return len;
+}
+
 /**
  * smb_compress_alloc_size() - Compute total allocation size required for compressed (dst) buffer.
  * @size:		uncompressed size
-- 
2.54.0
lmpx.com only provides a reader for public news (NNTP) servers. It is not affiliated with the servers or forums shown here and is not responsible for the content of articles, which is written by their respective authors.