diff --git a/fs/smb/common/compress/compress.c b/fs/smb/common/compress/compress.c index b07a317597a4..a4123c8f1c0a 100644 --- a/fs/smb/common/compress/compress.c +++ b/fs/smb/common/compress/compress.c @@ -95,6 +95,7 @@ static int smb_decompress_lz77_payload(const u8 **src, u32 *slen, u8 **dst, } static int smb_decompress_chained(__le16 alg, bool allow_chained, + bool allow_pattern, const struct smb2_compression_hdr *hdr, u32 slen, void *dst, u32 dlen) { @@ -143,6 +144,8 @@ static int smb_decompress_chained(__le16 alg, bool allow_chained, rc = smb_decompress_none(&src, &remaining, &out, &out_remaining, len); } else if (payload_alg == SMB3_COMPRESS_PATTERN) { + if (!allow_pattern) + return -EINVAL; rc = smb_decompress_pattern(&src, &remaining, &out, &out_remaining, len); } else if (payload_alg == alg && alg == SMB3_COMPRESS_LZ77) { @@ -185,6 +188,7 @@ static int smb_decompress_unchained(__le16 alg, * smb_compression_decompress() - decode an SMB2 compression transform * @alg: negotiated general-purpose compression algorithm * @allow_chained: whether chained transforms were negotiated + * @allow_pattern: whether Pattern_V1 payloads were negotiated * @src: transform header followed by compressed payload data * @slen: total number of bytes available at @src * @dst: output buffer for the reconstructed SMB2 message @@ -197,7 +201,8 @@ static int smb_decompress_unchained(__le16 alg, * Return: 0 on success, otherwise a negative errno. */ int smb_compression_decompress(__le16 alg, bool allow_chained, - const void *src, u32 slen, void *dst, u32 dlen) + bool allow_pattern, const void *src, u32 slen, + void *dst, u32 dlen) { const struct smb2_compression_hdr *hdr = src; @@ -207,8 +212,8 @@ int smb_compression_decompress(__le16 alg, bool allow_chained, return -EINVAL; if (hdr->Flags == cpu_to_le16(SMB2_COMPRESSION_FLAG_CHAINED)) - return smb_decompress_chained(alg, allow_chained, hdr, slen, - dst, dlen); + return smb_decompress_chained(alg, allow_chained, allow_pattern, + hdr, slen, dst, dlen); if (hdr->Flags != cpu_to_le16(SMB2_COMPRESSION_FLAG_NONE)) return -EINVAL; diff --git a/fs/smb/common/compress/compress.h b/fs/smb/common/compress/compress.h index 7ace3bf4b664..d6916669f887 100644 --- a/fs/smb/common/compress/compress.h +++ b/fs/smb/common/compress/compress.h @@ -20,7 +20,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); + bool allow_pattern, 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); diff --git a/fs/smb/server/compress.c b/fs/smb/server/compress.c index 95e48fa6b448..01d1771ff663 100644 --- a/fs/smb/server/compress.c +++ b/fs/smb/server/compress.c @@ -46,16 +46,25 @@ int ksmbd_decompress_request(struct ksmbd_conn *conn) return -EINVAL; orig_size = le32_to_cpu(hdr->OriginalCompressedSegmentSize); + /* + * For chained transforms the top-level header is only eight bytes; the + * Flags field overlays the first payload header. Reject unknown Flags + * and unnegotiated chained mode before allocating the output buffer. + */ if (hdr->Flags == cpu_to_le16(SMB2_COMPRESSION_FLAG_CHAINED)) { + if (!conn->compress_chained) + return -EINVAL; out_size = orig_size; - } else { + } else if (hdr->Flags == cpu_to_le16(SMB2_COMPRESSION_FLAG_NONE)) { offset = le32_to_cpu(hdr->Offset); if (offset > pdu_size - sizeof(*hdr) || check_add_overflow(orig_size, offset, &out_size)) return -EINVAL; + } else { + return -EINVAL; } - max_allowed_pdu_size = SMB3_MAX_MSGSIZE + conn->vals->max_write_size; + max_allowed_pdu_size = ksmbd_max_allowed_pdu_size(conn); if (out_size < sizeof(struct smb2_pdu) || out_size > max_allowed_pdu_size || out_size > MAX_STREAM_PROT_LEN) @@ -69,6 +78,7 @@ int ksmbd_decompress_request(struct ksmbd_conn *conn) *(__be32 *)out = cpu_to_be32(out_size); rc = smb_compression_decompress(conn->compress_algorithm, conn->compress_chained, + conn->compress_pattern, buf, pdu_size, out + 4, out_size); if (rc) { kvfree(out); diff --git a/fs/smb/server/connection.c b/fs/smb/server/connection.c index dee8e4aced99..ef6f202f4024 100644 --- a/fs/smb/server/connection.c +++ b/fs/smb/server/connection.c @@ -488,11 +488,7 @@ int ksmbd_conn_handler_loop(void *p) pdu_size = get_rfc1002_len(hdr_buf); ksmbd_debug(CONN, "RFC1002 header %u bytes\n", pdu_size); - if (ksmbd_conn_good(conn)) - max_allowed_pdu_size = - SMB3_MAX_MSGSIZE + conn->vals->max_write_size; - else - max_allowed_pdu_size = SMB3_MAX_MSGSIZE; + max_allowed_pdu_size = ksmbd_max_allowed_pdu_size(conn); if (pdu_size > max_allowed_pdu_size) { pr_err_ratelimited("PDU length(%u) exceeded maximum allowed pdu size(%u) on connection(%d)\n", diff --git a/fs/smb/server/connection.h b/fs/smb/server/connection.h index 2a194ee36fb4..0e4ebfac5558 100644 --- a/fs/smb/server/connection.h +++ b/fs/smb/server/connection.h @@ -210,6 +210,15 @@ static inline bool ksmbd_conn_good(struct ksmbd_conn *conn) return READ_ONCE(conn->status) == KSMBD_SESS_GOOD; } +static inline unsigned int +ksmbd_max_allowed_pdu_size(struct ksmbd_conn *conn) +{ + if (ksmbd_conn_good(conn)) + return SMB3_MAX_MSGSIZE + conn->vals->max_write_size; + + return SMB3_MAX_MSGSIZE; +} + static inline bool ksmbd_conn_need_negotiate(struct ksmbd_conn *conn) { return READ_ONCE(conn->status) == KSMBD_SESS_NEED_NEGOTIATE;