[PATCH] ksmbd: add SMB Direct RDMA encryption transform

Namjae Jeon <[email protected]>
Newsgroups org.kernel.vger.linux-cifs
Message-ID <[email protected]>
Port SMB Direct RDMA payload encryption support to the current ksmbd tree.
The current tree already supports all-state lookup for encrypted expired
sessions, so the overlapping lookup hunk from the original patch is
intentionally omitted.

Signed-off-by: Namjae Jeon <[email protected]>
---
 fs/smb/client/smb2pdu.h        |  24 --
 fs/smb/common/smb2pdu.h        |  22 ++
 fs/smb/common/smb2status.h     |   1 +
 fs/smb/server/auth.c           | 158 +++++++++++
 fs/smb/server/auth.h           |   4 +
 fs/smb/server/connection.c     |   2 +
 fs/smb/server/connection.h     |   3 +
 fs/smb/server/smb2pdu.c        | 500 ++++++++++++++++++++++++++++++---
 fs/smb/server/transport_rdma.c |  10 +
 fs/smb/server/transport_rdma.h |   2 +
 10 files changed, 663 insertions(+), 63 deletions(-)

diff --git a/fs/smb/client/smb2pdu.h b/fs/smb/client/smb2pdu.h
index b9bf2fa989d5..ab6c667bebc0 100644
--- a/fs/smb/client/smb2pdu.h
+++ b/fs/smb/client/smb2pdu.h
@@ -21,30 +21,6 @@
 /* The total header size for SMB2 read and write */
 #define SMB2_READWRITE_PDU_HEADER_SIZE (48 + sizeof(struct smb2_hdr))
 
-/* See MS-SMB2 2.2.43 */
-struct smb2_rdma_transform {
-	__le16 RdmaDescriptorOffset;
-	__le16 RdmaDescriptorLength;
-	__le32 Channel; /* for values see channel description in smb2 read above */
-	__le16 TransformCount;
-	__le16 Reserved1;
-	__le32 Reserved2;
-} __packed;
-
-/* TransformType */
-#define SMB2_RDMA_TRANSFORM_TYPE_ENCRYPTION	0x0001
-#define SMB2_RDMA_TRANSFORM_TYPE_SIGNING	0x0002
-
-struct smb2_rdma_crypto_transform {
-	__le16	TransformType;
-	__le16	SignatureLength;
-	__le16	NonceLength;
-	__u16	Reserved;
-	__u8	Signature[]; /* variable length */
-	/* u8 Nonce[] */
-	/* followed by padding */
-} __packed;
-
 /*
  *	Definitions for SMB2 Protocol Data Units (network frames)
  *
diff --git a/fs/smb/common/smb2pdu.h b/fs/smb/common/smb2pdu.h
index d9650aff0d3c..f9a8862cb3d4 100644
--- a/fs/smb/common/smb2pdu.h
+++ b/fs/smb/common/smb2pdu.h
@@ -743,6 +743,28 @@ struct smb2_close_rsp {
 #define SMB2_CHANNEL_RDMA_V1_INVALIDATE cpu_to_le32(0x00000002)
 #define SMB2_CHANNEL_RDMA_TRANSFORM     cpu_to_le32(0x00000003)
 
+/* See MS-SMB2 2.2.43. */
+struct smb2_rdma_transform {
+	__le16 RdmaDescriptorOffset;
+	__le16 RdmaDescriptorLength;
+	__le32 Channel;
+	__le16 TransformCount;
+	__le16 Reserved1;
+	__le32 Reserved2;
+} __packed;
+
+#define SMB2_RDMA_TRANSFORM_TYPE_ENCRYPTION	0x0001
+#define SMB2_RDMA_TRANSFORM_TYPE_SIGNING	0x0002
+
+struct smb2_rdma_crypto_transform {
+	__le16 TransformType;
+	__le16 SignatureLength;
+	__le16 NonceLength;
+	__le16 Reserved;
+	__u8 Signature[];
+	/* Followed by Nonce[] and optional alignment padding. */
+} __packed;
+
 /* SMB2 read request without RFC1001 length at the beginning */
 struct smb2_read_req {
 	struct smb2_hdr hdr;
diff --git a/fs/smb/common/smb2status.h b/fs/smb/common/smb2status.h
index b6421bc5113c..2989c3a5cb67 100644
--- a/fs/smb/common/smb2status.h
+++ b/fs/smb/common/smb2status.h
@@ -1049,6 +1049,7 @@ struct ntstatus {
 #define STATUS_WOW_ASSERTION			cpu_to_le32(0xC0009898) // -EIO
 #define STATUS_INVALID_SIGNATURE		cpu_to_le32(0xC000A000) // -EIO
 #define STATUS_HMAC_NOT_SUPPORTED		cpu_to_le32(0xC000A001) // -EIO
+#define STATUS_AUTH_TAG_MISMATCH		cpu_to_le32(0xC000A002) // -EBADMSG
 #define STATUS_IPSEC_QUEUE_OVERFLOW		cpu_to_le32(0xC000A010) // -EIO
 #define STATUS_ND_QUEUE_OVERFLOW		cpu_to_le32(0xC000A011) // -EIO
 #define STATUS_HOPLIMIT_EXCEEDED		cpu_to_le32(0xC000A012) // -EIO
diff --git a/fs/smb/server/auth.c b/fs/smb/server/auth.c
index 78491b20897e..db362c64af8d 100644
--- a/fs/smb/server/auth.c
+++ b/fs/smb/server/auth.c
@@ -833,6 +833,164 @@ static struct scatterlist *ksmbd_init_sg(struct kvec *iov, unsigned int nvec,
 	return sg;
 }
 
+/**
+ * ksmbd_init_rdma_sg() - build an AEAD scatterlist for an RDMA payload
+ * @buf: payload buffer
+ * @buflen: payload length
+ * @tag: authentication tag buffer
+ * @taglen: authentication tag length
+ *
+ * Split vmalloc-backed payloads at page boundaries and append the detached
+ * authentication tag as the final scatterlist entry.
+ *
+ * Return: allocated scatterlist, or NULL on allocation failure
+ */
+static struct scatterlist *ksmbd_init_rdma_sg(void *buf,
+					      unsigned int buflen,
+					      u8 *tag,
+					      unsigned int taglen)
+{
+	struct scatterlist *sg;
+	unsigned int nr_data = 1, nr_entries, i = 0;
+	void *data = buf;
+	int len = buflen;
+
+	if (is_vmalloc_addr(buf))
+		nr_data = DIV_ROUND_UP(offset_in_page(buf) + buflen, PAGE_SIZE);
+	nr_entries = nr_data + 1;
+
+	sg = kmalloc_objs(struct scatterlist, nr_entries, KSMBD_DEFAULT_GFP);
+	if (!sg)
+		return NULL;
+
+	sg_init_table(sg, nr_entries);
+	if (!is_vmalloc_addr(buf)) {
+		smb2_sg_set_buf(&sg[i++], buf, buflen);
+	} else {
+		while (len) {
+			unsigned int bytes = min_t(unsigned int,
+						PAGE_SIZE - offset_in_page(data), len);
+
+			sg_set_page(&sg[i++], vmalloc_to_page(data), bytes,
+				    offset_in_page(data));
+			data += bytes;
+			len -= bytes;
+		}
+	}
+	smb2_sg_set_buf(&sg[i], tag, taglen);
+	return sg;
+}
+
+/**
+ * ksmbd_crypt_rdma() - encrypt or decrypt an SMB Direct data buffer
+ * @conn: connection containing the negotiated cipher
+ * @key: session encryption or decryption key
+ * @buf: RDMA payload, transformed in place
+ * @buflen: payload length (the authentication tag is carried out of band)
+ * @nonce: transform nonce
+ * @nonce_len: nonce length
+ * @tag: authentication tag output for encryption, input for decryption
+ * @tag_len: authentication tag length
+ * @enc: true to encrypt, false to decrypt
+ *
+ * SMB2_RDMA_CRYPTO_TRANSFORM carries the nonce and authentication tag in the
+ * SMB2 message while only the payload is transferred through RDMA.  Therefore
+ * this uses AEAD without the normal SMB3 transform header as associated data.
+ *
+ * Return: 0 on success, otherwise a negative errno
+ */
+int ksmbd_crypt_rdma(struct ksmbd_conn *conn, const u8 *key,
+		     void *buf, unsigned int buflen, const u8 *nonce,
+		     unsigned int nonce_len, u8 *tag, unsigned int tag_len,
+		     bool enc)
+{
+	struct ksmbd_crypto_ctx *ctx;
+	struct crypto_aead *tfm;
+	struct aead_request *req = NULL;
+	struct scatterlist *sg = NULL;
+	unsigned int iv_len, crypt_len;
+	u8 auth_tag[SMB2_SIGNATURE_SIZE] = {};
+	u8 *iv = NULL;
+	int rc;
+	DECLARE_CRYPTO_WAIT(wait);
+
+	if (!buflen || !tag_len || tag_len > SMB2_SIGNATURE_SIZE)
+		return -EINVAL;
+	if (!enc)
+		memcpy(auth_tag, tag, tag_len);
+
+	if (conn->cipher_type == SMB2_ENCRYPTION_AES128_GCM ||
+	    conn->cipher_type == SMB2_ENCRYPTION_AES256_GCM) {
+		if (nonce_len != SMB3_AES_GCM_NONCE)
+			return -EINVAL;
+		ctx = ksmbd_crypto_ctx_find_gcm();
+	} else {
+		if (nonce_len != SMB3_AES_CCM_NONCE)
+			return -EINVAL;
+		ctx = ksmbd_crypto_ctx_find_ccm();
+	}
+	if (!ctx)
+		return -ENOMEM;
+
+	tfm = (conn->cipher_type == SMB2_ENCRYPTION_AES128_GCM ||
+	       conn->cipher_type == SMB2_ENCRYPTION_AES256_GCM) ?
+		CRYPTO_GCM(ctx) : CRYPTO_CCM(ctx);
+	if (conn->cipher_type == SMB2_ENCRYPTION_AES256_CCM ||
+	    conn->cipher_type == SMB2_ENCRYPTION_AES256_GCM)
+		rc = crypto_aead_setkey(tfm, key, SMB3_GCM256_CRYPTKEY_SIZE);
+	else
+		rc = crypto_aead_setkey(tfm, key, SMB3_GCM128_CRYPTKEY_SIZE);
+	if (rc)
+		goto out;
+
+	rc = crypto_aead_setauthsize(tfm, tag_len);
+	if (rc)
+		goto out;
+
+	req = aead_request_alloc(tfm, KSMBD_DEFAULT_GFP);
+	if (!req) {
+		rc = -ENOMEM;
+		goto out;
+	}
+
+	sg = ksmbd_init_rdma_sg(buf, buflen, auth_tag, tag_len);
+	if (!sg) {
+		rc = -ENOMEM;
+		goto out;
+	}
+
+	iv_len = crypto_aead_ivsize(tfm);
+	iv = kzalloc(iv_len, KSMBD_DEFAULT_GFP);
+	if (!iv) {
+		rc = -ENOMEM;
+		goto out;
+	}
+	if (conn->cipher_type == SMB2_ENCRYPTION_AES128_GCM ||
+	    conn->cipher_type == SMB2_ENCRYPTION_AES256_GCM) {
+		memcpy(iv, nonce, nonce_len);
+	} else {
+		iv[0] = 3;
+		memcpy(iv + 1, nonce, nonce_len);
+	}
+
+	crypt_len = buflen + (enc ? 0 : tag_len);
+	aead_request_set_crypt(req, sg, sg, crypt_len, iv);
+	aead_request_set_ad(req, 0);
+	aead_request_set_callback(req, CRYPTO_TFM_REQ_MAY_BACKLOG |
+				  CRYPTO_TFM_REQ_MAY_SLEEP,
+				  crypto_req_done, &wait);
+	rc = crypto_wait_req(enc ? crypto_aead_encrypt(req) :
+			     crypto_aead_decrypt(req), &wait);
+	if (!rc && enc)
+		memcpy(tag, auth_tag, tag_len);
+out:
+	kfree(iv);
+	kfree(sg);
+	aead_request_free(req);
+	ksmbd_release_crypto_ctx(ctx);
+	return rc;
+}
+
 int ksmbd_crypt_message(struct ksmbd_work *work, struct kvec *iov,
 			unsigned int nvec, int enc)
 {
diff --git a/fs/smb/server/auth.h b/fs/smb/server/auth.h
index f14b7c033264..7ce9c42d58f1 100644
--- a/fs/smb/server/auth.h
+++ b/fs/smb/server/auth.h
@@ -38,6 +38,10 @@ struct kvec;
 
 int ksmbd_crypt_message(struct ksmbd_work *work, struct kvec *iov,
 			unsigned int nvec, int enc);
+int ksmbd_crypt_rdma(struct ksmbd_conn *conn, const u8 *key,
+		     void *buf, unsigned int buflen, const u8 *nonce,
+		     unsigned int nonce_len, u8 *tag, unsigned int tag_len,
+		     bool enc);
 void ksmbd_copy_gss_neg_header(void *buf);
 int ksmbd_auth_ntlmv2(struct ksmbd_conn *conn, struct ksmbd_session *sess,
 		      struct ntlmv2_resp *ntlmv2, int blen, char *domain_name,
diff --git a/fs/smb/server/connection.c b/fs/smb/server/connection.c
index 5d729473dd18..d32f4f3cef93 100644
--- a/fs/smb/server/connection.c
+++ b/fs/smb/server/connection.c
@@ -77,6 +77,8 @@ static void proc_show_conn_features(struct seq_file *m,
 	proc_show_conn_feature(m, &separator,
 			       conn->compress_algorithm != SMB3_COMPRESS_NONE,
 			       "compress");
+	proc_show_conn_feature(m, &separator, conn->rdma_transform_ids,
+			       "rdma-transform");
 	proc_show_conn_feature(m, &separator, conn->posix_ext_supported, "posix");
 	if (!separator)
 		seq_puts(m, "none");
diff --git a/fs/smb/server/connection.h b/fs/smb/server/connection.h
index 421907aed473..63484c8efbbd 100644
--- a/fs/smb/server/connection.h
+++ b/fs/smb/server/connection.h
@@ -139,6 +139,9 @@ struct ksmbd_conn {
 	/* Negotiated SMB 3.1.1 compression capabilities. */
 	bool				compress_chained;
 	bool				compress_pattern;
+	/* Bitmap indexed by SMB2_RDMA_TRANSFORM_* IDs. */
+	unsigned long			rdma_transform_ids;
+	bool				rdma_transform_negotiated;
 	bool				posix_ext_supported;
 	bool				signing_negotiated;
 	__le16				signing_algorithm;
diff --git a/fs/smb/server/smb2pdu.c b/fs/smb/server/smb2pdu.c
index a564535132e5..b48eff02dbf8 100644
--- a/fs/smb/server/smb2pdu.c
+++ b/fs/smb/server/smb2pdu.c
@@ -1356,6 +1356,37 @@ static void build_compress_ctxt(struct smb2_compression_capabilities_context *pn
 	pneg_ctxt->CompressionAlgorithms[3] = 0;
 }
 
+/**
+ * build_rdma_ctx() - build an RDMA transform negotiate response context
+ * @ctxt: response context header to populate
+ * @transform_ids: bitmap of transforms common to the client and server
+ *
+ * Return: encoded negotiate context length
+ */
+static int build_rdma_ctx(struct smb2_neg_context *ctxt,
+			  unsigned long transform_ids)
+{
+	struct smb2_rdma_transform_capabilities_context *pneg_ctxt;
+	int count = 0;
+
+	pneg_ctxt = (void *)ctxt;
+	pneg_ctxt->ContextType = SMB2_RDMA_TRANSFORM_CAPABILITIES;
+	pneg_ctxt->Reserved = 0;
+	pneg_ctxt->Reserved1 = 0;
+	pneg_ctxt->Reserved2 = 0;
+	if (transform_ids & BIT(SMB2_RDMA_TRANSFORM_ENCRYPTION))
+		pneg_ctxt->RDMATransformIds[count++] =
+			cpu_to_le16(SMB2_RDMA_TRANSFORM_ENCRYPTION);
+	if (!count)
+		pneg_ctxt->RDMATransformIds[count++] =
+			cpu_to_le16(SMB2_RDMA_TRANSFORM_NONE);
+
+	pneg_ctxt->TransformCount = cpu_to_le16(count);
+	pneg_ctxt->DataLength = cpu_to_le16(8 + count * sizeof(__le16));
+	return sizeof(struct smb2_neg_context) +
+		le16_to_cpu(pneg_ctxt->DataLength);
+}
+
 static void build_sign_cap_ctxt(struct smb2_signing_capabilities *pneg_ctxt,
 				__le16 sign_algo)
 {
@@ -1431,6 +1462,18 @@ static unsigned int assemble_neg_contexts(struct ksmbd_conn *conn,
 			(conn->compress_pattern ? 12 : 10);
 	}
 
+	if (conn->rdma_transform_negotiated) {
+		struct smb2_neg_context *rdma_ctxt;
+
+		ctxt_size = round_up(ctxt_size, 8);
+		ksmbd_debug(SMB,
+			    "assemble SMB2_RDMA_TRANSFORM_CAPABILITIES context\n");
+		rdma_ctxt = (void *)(pneg_ctxt + ctxt_size);
+		ctxt_size += build_rdma_ctx(rdma_ctxt,
+					     conn->rdma_transform_ids);
+		neg_ctxt_cnt++;
+	}
+
 	if (conn->posix_ext_supported) {
 		ctxt_size = round_up(ctxt_size, 8);
 		ksmbd_debug(SMB,
@@ -1631,6 +1674,46 @@ static void decode_sign_cap_ctxt(struct ksmbd_conn *conn,
 	}
 }
 
+/**
+ * decode_rdma_ctx() - decode an RDMA transform negotiate request context
+ * @conn: connection being negotiated
+ * @ctxt: request context header to decode
+ * @ctxt_len: total context length, including the negotiate context header
+ *
+ * Record transforms supported by both peers only for SMB Direct connections.
+ *
+ * Return: NT status describing the decode result
+ */
+static __le32 decode_rdma_ctx(struct ksmbd_conn *conn,
+			      struct smb2_neg_context *ctxt, int ctxt_len)
+{
+	struct smb2_rdma_transform_capabilities_context *pneg_ctxt;
+	unsigned int count, i;
+
+	pneg_ctxt = (void *)ctxt;
+	/* RDMA transforms are a node capability, not just a transport capability. */
+	if (!ksmbd_rdma_enabled())
+		return STATUS_SUCCESS;
+
+	if (ctxt_len < sizeof(*pneg_ctxt))
+		return STATUS_INVALID_PARAMETER;
+
+	count = le16_to_cpu(pneg_ctxt->TransformCount);
+	if (!count || count >
+	    (ctxt_len - sizeof(*pneg_ctxt)) / sizeof(__le16))
+		return STATUS_INVALID_PARAMETER;
+
+	conn->rdma_transform_negotiated = true;
+	conn->rdma_transform_ids = 0;
+	for (i = 0; i < count; i++) {
+		u16 id = le16_to_cpu(pneg_ctxt->RDMATransformIds[i]);
+
+		if (id == SMB2_RDMA_TRANSFORM_ENCRYPTION)
+			conn->rdma_transform_ids |= BIT(id);
+	}
+	return STATUS_SUCCESS;
+}
+
 static __le32 deassemble_neg_contexts(struct ksmbd_conn *conn,
 				      struct smb2_negotiate_req *req,
 				      unsigned int len_of_smb)
@@ -1641,7 +1724,7 @@ static __le32 deassemble_neg_contexts(struct ksmbd_conn *conn,
 	unsigned int offset = le32_to_cpu(req->NegotiateContextOffset);
 	unsigned int neg_ctxt_cnt = le16_to_cpu(req->NegotiateContextCount);
 	__le32 status = STATUS_INVALID_PARAMETER;
-	int compress_ctxt_cnt = 0;
+	int compress_ctxt_cnt = 0, rdma_transform_ctxt_cnt = 0;
 
 	ksmbd_debug(SMB, "decoding %d negotiate contexts\n", neg_ctxt_cnt);
 	if (len_of_smb <= offset) {
@@ -1700,6 +1783,17 @@ static __le32 deassemble_neg_contexts(struct ksmbd_conn *conn,
 		} else if (pctx->ContextType == SMB2_NETNAME_NEGOTIATE_CONTEXT_ID) {
 			ksmbd_debug(SMB,
 				    "deassemble SMB2_NETNAME_NEGOTIATE_CONTEXT_ID context\n");
+		} else if (pctx->ContextType == SMB2_RDMA_TRANSFORM_CAPABILITIES) {
+			ksmbd_debug(SMB,
+				    "deassemble SMB2_RDMA_TRANSFORM_CAPABILITIES context\n");
+			if (ksmbd_rdma_enabled() &&
+			    rdma_transform_ctxt_cnt++) {
+				status = STATUS_INVALID_PARAMETER;
+				break;
+			}
+			status = decode_rdma_ctx(conn, pctx, ctxt_len);
+			if (status != STATUS_SUCCESS)
+				break;
 		} else if (pctx->ContextType == SMB2_POSIX_EXTENSIONS_AVAILABLE) {
 			ksmbd_debug(SMB,
 				    "deassemble SMB2_POSIX_EXTENSIONS_AVAILABLE context\n");
@@ -1807,6 +1901,9 @@ int smb2_handle_negotiate(struct ksmbd_work *work)
 			conn->preauth_info = NULL;
 			goto err_out;
 		}
+		if (!conn->cipher_type)
+			conn->rdma_transform_ids &=
+				~BIT(SMB2_RDMA_TRANSFORM_ENCRYPTION);
 
 		rc = init_smb3_11_server(conn);
 		if (rc < 0) {
@@ -8553,18 +8650,31 @@ static noinline int smb2_read_pipe(struct ksmbd_work *work)
 	return err;
 }
 
-static int smb2_set_remote_key_for_rdma(struct ksmbd_work *work,
-					struct smbdirect_buffer_descriptor_v1 *desc,
-					__le32 Channel,
-					__le16 ChannelInfoLength)
+/**
+ * smb2_set_rdma_key() - validate descriptors and save invalidation state
+ * @work: request work item
+ * @desc: first RDMA buffer descriptor
+ * @Channel: nested RDMA channel type
+ * @channel_info_len: descriptor array length
+ *
+ * Return: 0 on success, otherwise -EINVAL
+ */
+static int smb2_set_rdma_key(struct ksmbd_work *work,
+			     struct smbdirect_buffer_descriptor_v1 *desc,
+			     __le32 Channel, __le16 channel_info_len)
 {
 	unsigned int i, ch_count;
 
+	if (Channel != SMB2_CHANNEL_RDMA_V1 &&
+	    Channel != SMB2_CHANNEL_RDMA_V1_INVALIDATE)
+		return -EINVAL;
 	if (work->conn->dialect == SMB30_PROT_ID &&
 	    Channel != SMB2_CHANNEL_RDMA_V1)
 		return -EINVAL;
+	if (le16_to_cpu(channel_info_len) % sizeof(*desc))
+		return -EINVAL;
 
-	ch_count = le16_to_cpu(ChannelInfoLength) / sizeof(*desc);
+	ch_count = le16_to_cpu(channel_info_len) / sizeof(*desc);
 	if (ksmbd_debug_types & KSMBD_DEBUG_RDMA) {
 		for (i = 0; i < ch_count; i++) {
 			pr_info("RDMA r/w request %#x: token %#x, length %#x\n",
@@ -8583,9 +8693,223 @@ static int smb2_set_remote_key_for_rdma(struct ksmbd_work *work,
 	return 0;
 }
 
-static ssize_t smb2_read_rdma_channel(struct ksmbd_work *work,
-				      struct smb2_read_req *req, void *data_buf,
-				      size_t length)
+/**
+ * smb2_prep_rdma_read() - transform an RDMA READ payload
+ * @work: request work item
+ * @req: READ request controlling encryption or signing
+ * @rsp: READ response receiving transform metadata
+ * @data: data that will be transferred through RDMA
+ * @datalen: data length
+ *
+ * Encrypt the payload in place and encode the detached crypto metadata in
+ * the response buffer.
+ *
+ * Return: metadata length, zero when no transform applies, or negative errno
+ */
+static int smb2_prep_rdma_read(struct ksmbd_work *work,
+			       struct smb2_read_req *req,
+			       struct smb2_read_rsp *rsp,
+			       void *data, unsigned int datalen)
+{
+	struct ksmbd_conn *conn = work->conn;
+	struct smb2_rdma_transform *transform;
+	struct smb2_rdma_crypto_transform *crypto;
+	u8 *nonce;
+	unsigned int nonce_len = 0, transform_len;
+	u16 transform_type;
+	int err;
+
+	if (work->encrypted &&
+	    (conn->rdma_transform_ids & BIT(SMB2_RDMA_TRANSFORM_ENCRYPTION))) {
+		transform_type = SMB2_RDMA_TRANSFORM_TYPE_ENCRYPTION;
+		nonce_len = (conn->cipher_type == SMB2_ENCRYPTION_AES128_GCM ||
+			     conn->cipher_type == SMB2_ENCRYPTION_AES256_GCM) ?
+			SMB3_AES_GCM_NONCE : SMB3_AES_CCM_NONCE;
+	} else {
+		return 0;
+	}
+
+	transform = (struct smb2_rdma_transform *)rsp->Buffer;
+	crypto = (struct smb2_rdma_crypto_transform *)(transform + 1);
+	memset(transform, 0, sizeof(*transform) + sizeof(*crypto) +
+	       SMB2_SIGNATURE_SIZE + nonce_len);
+	transform->Channel = SMB2_CHANNEL_NONE;
+	transform->TransformCount = cpu_to_le16(1);
+
+	crypto->TransformType = cpu_to_le16(transform_type);
+	crypto->SignatureLength = cpu_to_le16(SMB2_SIGNATURE_SIZE);
+	crypto->NonceLength = cpu_to_le16(nonce_len);
+	nonce = crypto->Signature + SMB2_SIGNATURE_SIZE;
+
+	get_random_bytes(nonce, nonce_len);
+	err = ksmbd_crypt_rdma(conn,
+			       work->sess->smb3encryptionkey,
+			       data, datalen, nonce, nonce_len,
+			       crypto->Signature,
+			       SMB2_SIGNATURE_SIZE, true);
+	if (err)
+		return err;
+
+	transform_len = sizeof(*transform) + sizeof(*crypto) +
+		SMB2_SIGNATURE_SIZE + nonce_len;
+	rsp->Flags = SMB2_READFLAG_RESPONSE_RDMA_TRANSFORM;
+	rsp->DataLength = cpu_to_le32(transform_len);
+	return transform_len;
+}
+
+struct smb2_rdma_write_transform {
+	struct smbdirect_buffer_descriptor_v1 *desc;
+	struct smb2_rdma_crypto_transform *crypto;
+	u8 *nonce;
+	unsigned int desc_len;
+	unsigned int nonce_len;
+	unsigned int signature_len;
+	u16 type;
+	__le32 channel;
+};
+
+/**
+ * smb2_current_req_len() - return the current compound request element size
+ * @work: request work item
+ * @hdr: current SMB2 header
+ *
+ * Return: current request element length measured from the SMB2 header
+ */
+static unsigned int smb2_current_req_len(struct ksmbd_work *work,
+					 struct smb2_hdr *hdr)
+{
+	if (hdr->NextCommand)
+		return le32_to_cpu(hdr->NextCommand);
+	return get_rfc1002_len(work->request_buf) -
+		work->next_smb2_rcv_hdr_off;
+}
+
+/**
+ * check_rdma_desc() - validate an RDMA descriptor array
+ * @desc: descriptor array
+ * @desc_len: descriptor array length
+ * @required_len: minimum aggregate buffer length
+ *
+ * Return: 0 when the descriptors cover the transfer, otherwise -EINVAL
+ */
+static int check_rdma_desc(struct smbdirect_buffer_descriptor_v1 *desc,
+			   unsigned int desc_len,
+			   unsigned int required_len)
+{
+	unsigned int i, count;
+	u64 described_len = 0;
+
+	if (!desc_len || desc_len % sizeof(*desc))
+		return -EINVAL;
+	count = desc_len / sizeof(*desc);
+	if (!le32_to_cpu(desc[0].length))
+		return -EINVAL;
+	for (i = 0; i < count; i++)
+		described_len += le32_to_cpu(desc[i].length);
+	return described_len < required_len ? -EINVAL : 0;
+}
+
+/**
+ * smb2_parse_rdma_write_transform() - validate RDMA WRITE transform metadata
+ * @work: request work item
+ * @req: WRITE request containing the transform
+ * @info: parsed transform information
+ *
+ * Validate transform counts, crypto fields, descriptor alignment and bounds,
+ * negotiated algorithms, and the nested RDMA channel.
+ *
+ * Return: 0 on success, otherwise a negative errno
+ */
+static int smb2_parse_rdma_write_transform(struct ksmbd_work *work,
+					   struct smb2_write_req *req,
+					   struct smb2_rdma_write_transform *info)
+{
+	struct smb2_rdma_transform *transform;
+	struct smb2_rdma_crypto_transform *crypto;
+	unsigned int req_len = smb2_current_req_len(work, &req->hdr);
+	unsigned int offset = le16_to_cpu(req->WriteChannelInfoOffset);
+	unsigned int length = le16_to_cpu(req->WriteChannelInfoLength);
+	unsigned int desc_offset, desc_len, crypto_len, expected_desc_offset;
+
+	if (!work->conn->rdma_transform_ids ||
+	    offset < offsetof(struct smb2_write_req, Buffer) ||
+	    length < sizeof(*transform) || offset > req_len ||
+	    length > req_len - offset)
+		return -EINVAL;
+
+	transform = (struct smb2_rdma_transform *)((char *)req + offset);
+	if (le16_to_cpu(transform->TransformCount) != 1 ||
+	    (transform->Channel != SMB2_CHANNEL_RDMA_V1 &&
+	     transform->Channel != SMB2_CHANNEL_RDMA_V1_INVALIDATE))
+		return -EINVAL;
+
+	desc_offset = le16_to_cpu(transform->RdmaDescriptorOffset);
+	desc_len = le16_to_cpu(transform->RdmaDescriptorLength);
+	if (!desc_len || desc_len % sizeof(*info->desc) ||
+	    desc_offset < sizeof(*transform) || desc_offset > length ||
+	    desc_len > length - desc_offset)
+		return -EINVAL;
+
+	crypto = (struct smb2_rdma_crypto_transform *)(transform + 1);
+	if (length - sizeof(*transform) < sizeof(*crypto))
+		return -EINVAL;
+	info->type = le16_to_cpu(crypto->TransformType);
+	info->signature_len = le16_to_cpu(crypto->SignatureLength);
+	info->nonce_len = le16_to_cpu(crypto->NonceLength);
+	if (!info->signature_len)
+		return info->type == SMB2_RDMA_TRANSFORM_TYPE_ENCRYPTION ?
+			-EBADMSG : -EINVAL;
+	if (info->signature_len > SMB2_SIGNATURE_SIZE)
+		return info->type == SMB2_RDMA_TRANSFORM_TYPE_ENCRYPTION ?
+			-EBADMSG : -EINVAL;
+	if (info->signature_len > length - sizeof(*transform) - sizeof(*crypto) ||
+	    info->nonce_len > length - sizeof(*transform) - sizeof(*crypto) -
+				 info->signature_len)
+		return -EINVAL;
+
+	crypto_len = sizeof(*crypto) + info->signature_len + info->nonce_len;
+	expected_desc_offset = ALIGN(sizeof(*transform) + crypto_len, 8);
+	if (desc_offset != expected_desc_offset)
+		return -EINVAL;
+
+	if (info->type == SMB2_RDMA_TRANSFORM_TYPE_ENCRYPTION) {
+		unsigned int expected_nonce_len;
+
+		if (!(work->conn->rdma_transform_ids &
+		      BIT(SMB2_RDMA_TRANSFORM_ENCRYPTION)) || !work->encrypted)
+			return -EINVAL;
+		expected_nonce_len =
+			(work->conn->cipher_type == SMB2_ENCRYPTION_AES128_GCM ||
+			 work->conn->cipher_type == SMB2_ENCRYPTION_AES256_GCM) ?
+			SMB3_AES_GCM_NONCE : SMB3_AES_CCM_NONCE;
+		if (info->nonce_len != expected_nonce_len)
+			return -EBADMSG;
+	} else {
+		return -EINVAL;
+	}
+
+	info->desc = (struct smbdirect_buffer_descriptor_v1 *)
+		((char *)transform + desc_offset);
+	info->desc_len = desc_len;
+	info->crypto = crypto;
+	info->nonce = crypto->Signature + info->signature_len;
+	info->channel = transform->Channel;
+	return check_rdma_desc(info->desc, info->desc_len,
+			       le32_to_cpu(req->RemainingBytes));
+}
+
+/**
+ * smb2_read_rdma() - transfer READ data to client RDMA buffers
+ * @work: request work item
+ * @req: READ request containing client descriptors
+ * @data_buf: data to transfer
+ * @length: data length
+ *
+ * Return: transferred length on success, otherwise a negative errno
+ */
+static ssize_t smb2_read_rdma(struct ksmbd_work *work,
+			      struct smb2_read_req *req, void *data_buf,
+			      size_t length)
 {
 	int err;
 
@@ -8615,6 +8939,7 @@ int smb2_read(struct ksmbd_work *work)
 	size_t length, mincount;
 	ssize_t nbytes = 0, remain_bytes = 0;
 	int err = 0;
+	int rdma_transform_len = 0;
 	bool is_rdma_channel = false, async_interim = false;
 	unsigned int max_read_size = conn->vals->max_read_size;
 	unsigned int id = KSMBD_NO_FID, pid = KSMBD_NO_FID;
@@ -8649,6 +8974,12 @@ int smb2_read(struct ksmbd_work *work)
 		pid = req->PersistentFileId;
 	}
 
+	if (req->Channel != SMB2_CHANNEL_NONE &&
+	    req->Channel != SMB2_CHANNEL_RDMA_V1 &&
+	    req->Channel != SMB2_CHANNEL_RDMA_V1_INVALIDATE) {
+		err = -EINVAL;
+		goto out;
+	}
 	if (req->Channel == SMB2_CHANNEL_RDMA_V1_INVALIDATE ||
 	    req->Channel == SMB2_CHANNEL_RDMA_V1) {
 		is_rdma_channel = true;
@@ -8661,16 +8992,24 @@ int smb2_read(struct ksmbd_work *work)
 
 	if (is_rdma_channel == true) {
 		unsigned int ch_offset = le16_to_cpu(req->ReadChannelInfoOffset);
+		unsigned int ch_len = le16_to_cpu(req->ReadChannelInfoLength);
+		unsigned int req_len = smb2_current_req_len(work, &req->hdr);
+		struct smbdirect_buffer_descriptor_v1 *desc;
 
-		if (ch_offset < offsetof(struct smb2_read_req, Buffer)) {
+		if (!le32_to_cpu(req->Length) ||
+		    ch_offset < offsetof(struct smb2_read_req, Buffer) ||
+		    ch_offset > req_len || ch_len > req_len - ch_offset) {
 			err = -EINVAL;
 			goto out;
 		}
-		err = smb2_set_remote_key_for_rdma(work,
-						   (struct smbdirect_buffer_descriptor_v1 *)
-						   ((char *)req + ch_offset),
-						   req->Channel,
-						   req->ReadChannelInfoLength);
+		desc = (struct smbdirect_buffer_descriptor_v1 *)
+			((char *)req + ch_offset);
+		err = check_rdma_desc(desc, ch_len, le32_to_cpu(req->Length));
+		if (err)
+			goto out;
+		err = smb2_set_rdma_key(work, desc,
+					req->Channel,
+					req->ReadChannelInfoLength);
 		if (err)
 			goto out;
 	}
@@ -8753,10 +9092,19 @@ int smb2_read(struct ksmbd_work *work)
 		    nbytes, offset, mincount);
 
 	if (is_rdma_channel == true) {
+		rdma_transform_len = smb2_prep_rdma_read(work, req,
+							 rsp,
+							 aux_payload_buf,
+							 nbytes);
+		if (rdma_transform_len < 0) {
+			kvfree(aux_payload_buf);
+			err = rdma_transform_len;
+			goto out;
+		}
 		/* write data to the client using rdma channel */
-		remain_bytes = smb2_read_rdma_channel(work, req,
-						      aux_payload_buf,
-						      nbytes);
+		remain_bytes = smb2_read_rdma(work, req,
+					      aux_payload_buf,
+					      nbytes);
 		kvfree(aux_payload_buf);
 		aux_payload_buf = NULL;
 		nbytes = 0;
@@ -8769,11 +9117,13 @@ int smb2_read(struct ksmbd_work *work)
 	rsp->StructureSize = cpu_to_le16(17);
 	rsp->DataOffset = 80;
 	rsp->Reserved = 0;
-	rsp->DataLength = cpu_to_le32(nbytes);
+	rsp->DataLength = cpu_to_le32(rdma_transform_len ?: nbytes);
 	rsp->DataRemaining = cpu_to_le32(remain_bytes);
-	rsp->Flags = 0;
+	rsp->Flags = rdma_transform_len ?
+		SMB2_READFLAG_RESPONSE_RDMA_TRANSFORM : 0;
 	err = ksmbd_iov_pin_rsp_read(work, (void *)rsp,
-				     offsetof(struct smb2_read_rsp, Buffer),
+				     offsetof(struct smb2_read_rsp, Buffer) +
+				     rdma_transform_len,
 				     aux_payload_buf, nbytes);
 	if (err) {
 		kvfree(aux_payload_buf);
@@ -8885,10 +9235,28 @@ static noinline int smb2_write_pipe(struct ksmbd_work *work)
 	return err;
 }
 
-static ssize_t smb2_write_rdma_channel(struct ksmbd_work *work,
-				       struct smb2_write_req *req,
-				       struct ksmbd_file *fp,
-				       loff_t offset, size_t length, bool sync)
+/**
+ * smb2_write_rdma() - receive and store an RDMA WRITE payload
+ * @work: request work item
+ * @desc: client RDMA buffer descriptors
+ * @desc_len: descriptor array length
+ * @transform: parsed transform, or NULL for an untransformed transfer
+ * @fp: target open file
+ * @offset: target file offset
+ * @length: transfer length
+ * @sync: request synchronous storage completion
+ *
+ * Receive the payload, authenticate or decrypt it when required, and write it
+ * to the target file.
+ *
+ * Return: written byte count on success, otherwise a negative errno
+ */
+static ssize_t smb2_write_rdma(struct ksmbd_work *work,
+			       struct smbdirect_buffer_descriptor_v1 *desc,
+			       unsigned int desc_len,
+			       struct smb2_rdma_write_transform *transform,
+			       struct ksmbd_file *fp, loff_t offset,
+			       size_t length, bool sync)
 {
 	char *data_buf;
 	int ret;
@@ -8898,15 +9266,27 @@ static ssize_t smb2_write_rdma_channel(struct ksmbd_work *work,
 	if (!data_buf)
 		return -ENOMEM;
 
-	ret = ksmbd_conn_rdma_read(work->conn, data_buf, length,
-				   (struct smbdirect_buffer_descriptor_v1 *)
-				   ((char *)req + le16_to_cpu(req->WriteChannelInfoOffset)),
-				   le16_to_cpu(req->WriteChannelInfoLength));
+	ret = ksmbd_conn_rdma_read(work->conn, data_buf, length, desc,
+				   desc_len);
 	if (ret < 0) {
 		kvfree(data_buf);
 		return ret;
 	}
 
+	if (transform &&
+	    transform->type == SMB2_RDMA_TRANSFORM_TYPE_ENCRYPTION) {
+		ret = ksmbd_crypt_rdma(work->conn,
+				       work->sess->smb3decryptionkey,
+				       data_buf, length, transform->nonce,
+				       transform->nonce_len,
+				       transform->crypto->Signature,
+				       transform->signature_len, false);
+		if (ret) {
+			kvfree(data_buf);
+			return ret == -ENOMEM ? ret : -EBADMSG;
+		}
+	}
+
 	ret = ksmbd_vfs_write(work, fp, data_buf, length, &offset, sync, &nbytes);
 	kvfree(data_buf);
 	if (ret < 0)
@@ -8925,6 +9305,10 @@ int smb2_write(struct ksmbd_work *work)
 {
 	struct smb2_write_req *req;
 	struct smb2_write_rsp *rsp;
+	struct smb2_rdma_write_transform rdma_transform = {};
+	struct smb2_rdma_write_transform *rdma_info = NULL;
+	struct smbdirect_buffer_descriptor_v1 *rdma_desc = NULL;
+	unsigned int rdma_desc_len = 0;
 	struct ksmbd_file *fp = NULL;
 	loff_t offset;
 	size_t length;
@@ -8969,8 +9353,21 @@ int smb2_write(struct ksmbd_work *work)
 	}
 	length = le32_to_cpu(req->Length);
 
+	if (req->Channel != SMB2_CHANNEL_NONE &&
+	    req->Channel != SMB2_CHANNEL_RDMA_V1 &&
+	    req->Channel != SMB2_CHANNEL_RDMA_V1_INVALIDATE &&
+	    req->Channel != SMB2_CHANNEL_RDMA_TRANSFORM) {
+		err = -EINVAL;
+		goto out;
+	}
+	if (req->Channel == SMB2_CHANNEL_RDMA_TRANSFORM &&
+	    work->conn->dialect != SMB311_PROT_ID) {
+		err = -EINVAL;
+		goto out;
+	}
 	if (req->Channel == SMB2_CHANNEL_RDMA_V1 ||
-	    req->Channel == SMB2_CHANNEL_RDMA_V1_INVALIDATE) {
+	    req->Channel == SMB2_CHANNEL_RDMA_V1_INVALIDATE ||
+	    req->Channel == SMB2_CHANNEL_RDMA_TRANSFORM) {
 		is_rdma_channel = true;
 		max_write_size = get_smbd_max_read_write_size(work->conn->transport);
 		if (max_write_size == 0) {
@@ -8995,17 +9392,37 @@ int smb2_write(struct ksmbd_work *work)
 
 	if (is_rdma_channel == true) {
 		unsigned int ch_offset = le16_to_cpu(req->WriteChannelInfoOffset);
+		unsigned int ch_len = le16_to_cpu(req->WriteChannelInfoLength);
+		unsigned int req_len = smb2_current_req_len(work, &req->hdr);
 
-		if (req->Length != 0 || req->DataOffset != 0 ||
-		    ch_offset < offsetof(struct smb2_write_req, Buffer)) {
+		if (!length || req->Length != 0 || req->DataOffset != 0 ||
+		    ch_offset < offsetof(struct smb2_write_req, Buffer) ||
+		    ch_offset > req_len || ch_len > req_len - ch_offset) {
 			err = -EINVAL;
 			goto out;
 		}
-		err = smb2_set_remote_key_for_rdma(work,
-						   (struct smbdirect_buffer_descriptor_v1 *)
-						   ((char *)req + ch_offset),
-						   req->Channel,
-						   req->WriteChannelInfoLength);
+		if (req->Channel == SMB2_CHANNEL_RDMA_TRANSFORM) {
+			err = smb2_parse_rdma_write_transform(work, req,
+							      &rdma_transform);
+			if (err)
+				goto out;
+			rdma_desc = rdma_transform.desc;
+			rdma_desc_len = rdma_transform.desc_len;
+			rdma_info = &rdma_transform;
+			err = smb2_set_rdma_key(work, rdma_desc,
+						rdma_transform.channel,
+						cpu_to_le16(rdma_desc_len));
+		} else {
+			rdma_desc = (struct smbdirect_buffer_descriptor_v1 *)
+				((char *)req + ch_offset);
+			rdma_desc_len = ch_len;
+			err = check_rdma_desc(rdma_desc, rdma_desc_len, length);
+			if (err)
+				goto out;
+			err = smb2_set_rdma_key(work, rdma_desc,
+						req->Channel,
+						req->WriteChannelInfoLength);
+		}
 		if (err)
 			goto out;
 	}
@@ -9074,8 +9491,9 @@ int smb2_write(struct ksmbd_work *work)
 		/* read data from the client using rdma channel, and
 		 * write the data.
 		 */
-		nbytes = smb2_write_rdma_channel(work, req, fp, offset, length,
-						 writethrough);
+		nbytes = smb2_write_rdma(work, rdma_desc, rdma_desc_len,
+					 rdma_info, fp, offset, length,
+					 writethrough);
 		if (nbytes < 0) {
 			err = (int)nbytes;
 			goto out;
@@ -9112,6 +9530,10 @@ int smb2_write(struct ksmbd_work *work)
 		rsp->hdr.Status = STATUS_SHARING_VIOLATION;
 	else if (err == -EINVAL)
 		rsp->hdr.Status = STATUS_INVALID_PARAMETER;
+	else if (err == -EBADMSG)
+		rsp->hdr.Status = STATUS_AUTH_TAG_MISMATCH;
+	else if (err == -EKEYREJECTED)
+		rsp->hdr.Status = STATUS_INVALID_SIGNATURE;
 	else if (rsp->hdr.Status == 0)
 		rsp->hdr.Status = STATUS_INVALID_HANDLE;
 
diff --git a/fs/smb/server/transport_rdma.c b/fs/smb/server/transport_rdma.c
index 85d12c4c354c..ee28a4d1cc86 100644
--- a/fs/smb/server/transport_rdma.c
+++ b/fs/smb/server/transport_rdma.c
@@ -76,6 +76,8 @@ static int smb_direct_max_receive_size = 1364;
 
 static int smb_direct_max_read_write_size = SMBD_DEFAULT_IOSIZE;
 
+static bool smb_direct_enabled;
+
 static struct smb_direct_listener {
 	int			port;
 
@@ -512,18 +514,26 @@ int ksmbd_rdma_init(void)
 	ksmbd_debug(RDMA, "iWarp RDMA listener. socket=%p\n",
 		    smb_direct_iw_listener.socket);
 
+	WRITE_ONCE(smb_direct_enabled, true);
 	return 0;
 err:
+	WRITE_ONCE(smb_direct_enabled, false);
 	ksmbd_rdma_stop_listening();
 	return ret;
 }
 
 void ksmbd_rdma_stop_listening(void)
 {
+	WRITE_ONCE(smb_direct_enabled, false);
 	smb_direct_listener_destroy(&smb_direct_ib_listener);
 	smb_direct_listener_destroy(&smb_direct_iw_listener);
 }
 
+bool ksmbd_rdma_enabled(void)
+{
+	return READ_ONCE(smb_direct_enabled);
+}
+
 bool ksmbd_rdma_capable_netdev(struct net_device *netdev)
 {
 	u8 node_type = smbdirect_netdev_rdma_capable_node_type(netdev);
diff --git a/fs/smb/server/transport_rdma.h b/fs/smb/server/transport_rdma.h
index 8b78917a1795..23247713b5c3 100644
--- a/fs/smb/server/transport_rdma.h
+++ b/fs/smb/server/transport_rdma.h
@@ -14,12 +14,14 @@
 #ifdef CONFIG_SMB_SERVER_SMBDIRECT
 int ksmbd_rdma_init(void);
 void ksmbd_rdma_stop_listening(void);
+bool ksmbd_rdma_enabled(void);
 bool ksmbd_rdma_capable_netdev(struct net_device *netdev);
 void init_smbd_max_io_size(unsigned int sz);
 unsigned int get_smbd_max_read_write_size(struct ksmbd_transport *kt);
 #else
 static inline int ksmbd_rdma_init(void) { return 0; }
 static inline void ksmbd_rdma_stop_listening(void) { }
+static inline bool ksmbd_rdma_enabled(void) { return false; }
 static inline bool ksmbd_rdma_capable_netdev(struct net_device *netdev) { return false; }
 static inline void init_smbd_max_io_size(unsigned int sz) { }
 static inline unsigned int get_smbd_max_read_write_size(struct ksmbd_transport *kt) { return 0; }
-- 
2.25.1
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.