xref: /linux/fs/smb/server/auth.c (revision 26ba30221c03364d6ed9910be8da4c1fd871b07b)
1 // SPDX-License-Identifier: GPL-2.0-or-later
2 /*
3  *   Copyright (C) 2016 Namjae Jeon <linkinjeon@kernel.org>
4  *   Copyright (C) 2018 Samsung Electronics Co., Ltd.
5  */
6 
7 #include <linux/kernel.h>
8 #include <linux/fs.h>
9 #include <linux/uaccess.h>
10 #include <linux/backing-dev.h>
11 #include <linux/writeback.h>
12 #include <linux/uio.h>
13 #include <linux/xattr.h>
14 #include <crypto/aead.h>
15 #include <crypto/aes-cbc-macs.h>
16 #include <crypto/md5.h>
17 #include <crypto/sha2.h>
18 #include <crypto/utils.h>
19 #include <linux/random.h>
20 #include <linux/scatterlist.h>
21 
22 #include "auth.h"
23 #include "glob.h"
24 
25 #include <linux/fips.h>
26 #include <crypto/arc4.h>
27 #include <crypto/des.h>
28 
29 #include "server.h"
30 #include "smb_common.h"
31 #include "connection.h"
32 #include "mgmt/user_session.h"
33 #include "mgmt/user_config.h"
34 #include "crypto_ctx.h"
35 #include "transport_ipc.h"
36 
37 /*
38  * Fixed format data defining GSS header and fixed string
39  * "not_defined_in_RFC4178@please_ignore".
40  * So sec blob data in neg phase could be generated statically.
41  */
42 static char NEGOTIATE_GSS_HEADER[AUTH_GSS_LENGTH] = {
43 #ifdef CONFIG_SMB_SERVER_KERBEROS5
44 	0x60, 0x5e, 0x06, 0x06, 0x2b, 0x06, 0x01, 0x05,
45 	0x05, 0x02, 0xa0, 0x54, 0x30, 0x52, 0xa0, 0x24,
46 	0x30, 0x22, 0x06, 0x09, 0x2a, 0x86, 0x48, 0x86,
47 	0xf7, 0x12, 0x01, 0x02, 0x02, 0x06, 0x09, 0x2a,
48 	0x86, 0x48, 0x82, 0xf7, 0x12, 0x01, 0x02, 0x02,
49 	0x06, 0x0a, 0x2b, 0x06, 0x01, 0x04, 0x01, 0x82,
50 	0x37, 0x02, 0x02, 0x0a, 0xa3, 0x2a, 0x30, 0x28,
51 	0xa0, 0x26, 0x1b, 0x24, 0x6e, 0x6f, 0x74, 0x5f,
52 	0x64, 0x65, 0x66, 0x69, 0x6e, 0x65, 0x64, 0x5f,
53 	0x69, 0x6e, 0x5f, 0x52, 0x46, 0x43, 0x34, 0x31,
54 	0x37, 0x38, 0x40, 0x70, 0x6c, 0x65, 0x61, 0x73,
55 	0x65, 0x5f, 0x69, 0x67, 0x6e, 0x6f, 0x72, 0x65
56 #else
57 	0x60, 0x48, 0x06, 0x06, 0x2b, 0x06, 0x01, 0x05,
58 	0x05, 0x02, 0xa0, 0x3e, 0x30, 0x3c, 0xa0, 0x0e,
59 	0x30, 0x0c, 0x06, 0x0a, 0x2b, 0x06, 0x01, 0x04,
60 	0x01, 0x82, 0x37, 0x02, 0x02, 0x0a, 0xa3, 0x2a,
61 	0x30, 0x28, 0xa0, 0x26, 0x1b, 0x24, 0x6e, 0x6f,
62 	0x74, 0x5f, 0x64, 0x65, 0x66, 0x69, 0x6e, 0x65,
63 	0x64, 0x5f, 0x69, 0x6e, 0x5f, 0x52, 0x46, 0x43,
64 	0x34, 0x31, 0x37, 0x38, 0x40, 0x70, 0x6c, 0x65,
65 	0x61, 0x73, 0x65, 0x5f, 0x69, 0x67, 0x6e, 0x6f,
66 	0x72, 0x65
67 #endif
68 };
69 
70 void ksmbd_copy_gss_neg_header(void *buf)
71 {
72 	memcpy(buf, NEGOTIATE_GSS_HEADER, AUTH_GSS_LENGTH);
73 }
74 
75 static int calc_ntlmv2_hash(struct ksmbd_conn *conn, struct ksmbd_session *sess,
76 			    char *ntlmv2_hash, char *dname)
77 {
78 	int ret, len, conv_len;
79 	wchar_t *domain = NULL;
80 	__le16 *uniname = NULL;
81 	struct hmac_md5_ctx ctx;
82 
83 	hmac_md5_init_usingrawkey(&ctx, user_passkey(sess->user),
84 				  CIFS_ENCPWD_SIZE);
85 
86 	/* convert user_name to unicode */
87 	len = strlen(user_name(sess->user));
88 	uniname = kzalloc(2 + UNICODE_LEN(len), KSMBD_DEFAULT_GFP);
89 	if (!uniname) {
90 		ret = -ENOMEM;
91 		goto out;
92 	}
93 
94 	conv_len = smb_strtoUTF16(uniname, user_name(sess->user), len,
95 				  conn->local_nls);
96 	if (conv_len < 0 || conv_len > len) {
97 		ret = -EINVAL;
98 		goto out;
99 	}
100 	UniStrupr(uniname);
101 
102 	hmac_md5_update(&ctx, (const u8 *)uniname, UNICODE_LEN(conv_len));
103 
104 	/* Convert domain name or conn name to unicode and uppercase */
105 	len = strlen(dname);
106 	domain = kzalloc(2 + UNICODE_LEN(len), KSMBD_DEFAULT_GFP);
107 	if (!domain) {
108 		ret = -ENOMEM;
109 		goto out;
110 	}
111 
112 	conv_len = smb_strtoUTF16((__le16 *)domain, dname, len,
113 				  conn->local_nls);
114 	if (conv_len < 0 || conv_len > len) {
115 		ret = -EINVAL;
116 		goto out;
117 	}
118 
119 	hmac_md5_update(&ctx, (const u8 *)domain, UNICODE_LEN(conv_len));
120 	hmac_md5_final(&ctx, ntlmv2_hash);
121 	ret = 0;
122 out:
123 	kfree(uniname);
124 	kfree(domain);
125 	return ret;
126 }
127 
128 /**
129  * ksmbd_auth_ntlmv2() - NTLMv2 authentication handler
130  * @conn:		connection
131  * @sess:		session of connection
132  * @ntlmv2:		NTLMv2 challenge response
133  * @blen:		NTLMv2 blob length
134  * @domain_name:	domain name
135  * @cryptkey:		session crypto key
136  * @sess_key:		derived session key output buffer
137  *
138  * Return:	0 on success, error number on error
139  */
140 int ksmbd_auth_ntlmv2(struct ksmbd_conn *conn, struct ksmbd_session *sess,
141 		      struct ntlmv2_resp *ntlmv2, int blen, char *domain_name,
142 		      char *cryptkey, char *sess_key)
143 {
144 	char ntlmv2_hash[CIFS_ENCPWD_SIZE];
145 	char ntlmv2_rsp[CIFS_HMAC_MD5_HASH_SIZE];
146 	char base_key[SMB2_NTLMV2_SESSKEY_SIZE];
147 	struct hmac_md5_ctx ctx;
148 	int rc;
149 
150 	if (fips_enabled) {
151 		ksmbd_debug(AUTH, "NTLMv2 support is disabled due to FIPS\n");
152 		return -EOPNOTSUPP;
153 	}
154 
155 	rc = calc_ntlmv2_hash(conn, sess, ntlmv2_hash, domain_name);
156 	if (rc) {
157 		ksmbd_debug(AUTH, "could not get v2 hash rc %d\n", rc);
158 		return rc;
159 	}
160 
161 	hmac_md5_init_usingrawkey(&ctx, ntlmv2_hash, CIFS_HMAC_MD5_HASH_SIZE);
162 	hmac_md5_update(&ctx, cryptkey, CIFS_CRYPTO_KEY_SIZE);
163 	hmac_md5_update(&ctx, (const u8 *)&ntlmv2->blob_signature, blen);
164 	hmac_md5_final(&ctx, ntlmv2_rsp);
165 
166 	/* Generate the session key */
167 	hmac_md5_usingrawkey(ntlmv2_hash, CIFS_HMAC_MD5_HASH_SIZE,
168 			     ntlmv2_rsp, CIFS_HMAC_MD5_HASH_SIZE,
169 			     base_key);
170 
171 	if (crypto_memneq(ntlmv2->ntlmv2_hash, ntlmv2_rsp,
172 			  CIFS_HMAC_MD5_HASH_SIZE)) {
173 		rc = -EINVAL;
174 		goto out;
175 	}
176 
177 	memcpy(sess_key, base_key, sizeof(base_key));
178 	rc = 0;
179 out:
180 	memzero_explicit(ntlmv2_hash, sizeof(ntlmv2_hash));
181 	memzero_explicit(ntlmv2_rsp, sizeof(ntlmv2_rsp));
182 	memzero_explicit(base_key, sizeof(base_key));
183 	return rc;
184 }
185 
186 /**
187  * ksmbd_decode_ntlmssp_auth_blob() - helper function to construct
188  * authenticate blob
189  * @authblob:	authenticate blob source pointer
190  * @blob_len:	length of the @authblob message
191  * @conn:	connection
192  * @sess:	session of connection
193  * @sess_key:	derived session key output buffer
194  *
195  * Return:	0 on success, error number on error
196  */
197 int ksmbd_decode_ntlmssp_auth_blob(struct authenticate_message *authblob,
198 				   int blob_len, struct ksmbd_conn *conn,
199 				   struct ksmbd_session *sess, char *sess_key)
200 {
201 	char *domain_name;
202 	unsigned int nt_off, dn_off;
203 	unsigned short nt_len, dn_len;
204 	int ret;
205 
206 	if (blob_len < sizeof(struct authenticate_message)) {
207 		ksmbd_debug(AUTH, "negotiate blob len %d too small\n",
208 			    blob_len);
209 		return -EINVAL;
210 	}
211 
212 	if (memcmp(authblob->Signature, "NTLMSSP", 8)) {
213 		ksmbd_debug(AUTH, "blob signature incorrect %s\n",
214 			    authblob->Signature);
215 		return -EINVAL;
216 	}
217 
218 	nt_off = le32_to_cpu(authblob->NtChallengeResponse.BufferOffset);
219 	nt_len = le16_to_cpu(authblob->NtChallengeResponse.Length);
220 	dn_off = le32_to_cpu(authblob->DomainName.BufferOffset);
221 	dn_len = le16_to_cpu(authblob->DomainName.Length);
222 
223 	if (blob_len < (u64)dn_off + dn_len || blob_len < (u64)nt_off + nt_len ||
224 	    nt_len < CIFS_ENCPWD_SIZE)
225 		return -EINVAL;
226 
227 	/* TODO : use domain name that imported from configuration file */
228 	domain_name = smb_strndup_from_utf16((const char *)authblob + dn_off,
229 					     dn_len, true, conn->local_nls);
230 	if (IS_ERR(domain_name))
231 		return PTR_ERR(domain_name);
232 
233 	/* process NTLMv2 authentication */
234 	ksmbd_debug(AUTH, "decode_ntlmssp_authenticate_blob dname%s\n",
235 		    domain_name);
236 	ret = ksmbd_auth_ntlmv2(conn, sess,
237 				(struct ntlmv2_resp *)((char *)authblob + nt_off),
238 				nt_len - CIFS_ENCPWD_SIZE,
239 				domain_name, conn->ntlmssp.cryptkey, sess_key);
240 	kfree(domain_name);
241 	if (ret)
242 		return ret;
243 
244 	/* The recovered secondary session key */
245 	if (conn->ntlmssp.client_flags & NTLMSSP_NEGOTIATE_KEY_XCH) {
246 		struct arc4_ctx *ctx_arc4;
247 		unsigned int sess_key_off, sess_key_len;
248 
249 		sess_key_off = le32_to_cpu(authblob->SessionKey.BufferOffset);
250 		sess_key_len = le16_to_cpu(authblob->SessionKey.Length);
251 
252 		if (blob_len < (u64)sess_key_off + sess_key_len)
253 			return -EINVAL;
254 
255 		if (sess_key_len > CIFS_KEY_SIZE)
256 			return -EINVAL;
257 
258 		ctx_arc4 = kmalloc_obj(*ctx_arc4, KSMBD_DEFAULT_GFP);
259 		if (!ctx_arc4)
260 			return -ENOMEM;
261 
262 		arc4_setkey(ctx_arc4, sess_key, SMB2_NTLMV2_SESSKEY_SIZE);
263 		arc4_crypt(ctx_arc4, sess_key,
264 			   (char *)authblob + sess_key_off, sess_key_len);
265 		kfree_sensitive(ctx_arc4);
266 	}
267 
268 	return ret;
269 }
270 
271 /**
272  * ksmbd_decode_ntlmssp_neg_blob() - helper function to construct
273  * negotiate blob
274  * @negblob: negotiate blob source pointer
275  * @blob_len:	length of the @authblob message
276  * @conn:	connection
277  *
278  */
279 int ksmbd_decode_ntlmssp_neg_blob(struct negotiate_message *negblob,
280 				  int blob_len, struct ksmbd_conn *conn)
281 {
282 	if (blob_len < sizeof(struct negotiate_message)) {
283 		ksmbd_debug(AUTH, "negotiate blob len %d too small\n",
284 			    blob_len);
285 		return -EINVAL;
286 	}
287 
288 	if (memcmp(negblob->Signature, "NTLMSSP", 8)) {
289 		ksmbd_debug(AUTH, "blob signature incorrect %s\n",
290 			    negblob->Signature);
291 		return -EINVAL;
292 	}
293 
294 	conn->ntlmssp.client_flags = le32_to_cpu(negblob->NegotiateFlags);
295 	return 0;
296 }
297 
298 /**
299  * ksmbd_build_ntlmssp_challenge_blob() - helper function to construct
300  * challenge blob
301  * @chgblob: challenge blob source pointer to initialize
302  * @conn:	connection
303  *
304  */
305 unsigned int
306 ksmbd_build_ntlmssp_challenge_blob(struct challenge_message *chgblob,
307 				   struct ksmbd_conn *conn)
308 {
309 	struct target_info *tinfo;
310 	wchar_t *name;
311 	__u8 *target_name;
312 	unsigned int flags, blob_off, blob_len, type, target_info_len = 0;
313 	int len, uni_len, conv_len;
314 	int cflags = conn->ntlmssp.client_flags;
315 
316 	memcpy(chgblob->Signature, NTLMSSP_SIGNATURE, 8);
317 	chgblob->MessageType = NtLmChallenge;
318 
319 	flags = NTLMSSP_NEGOTIATE_UNICODE |
320 		NTLMSSP_NEGOTIATE_NTLM | NTLMSSP_TARGET_TYPE_SERVER |
321 		NTLMSSP_NEGOTIATE_TARGET_INFO;
322 
323 	if (cflags & NTLMSSP_NEGOTIATE_SIGN) {
324 		flags |= NTLMSSP_NEGOTIATE_SIGN;
325 		flags |= cflags & (NTLMSSP_NEGOTIATE_128 |
326 				   NTLMSSP_NEGOTIATE_56);
327 	}
328 
329 	if (cflags & NTLMSSP_NEGOTIATE_SEAL && smb3_encryption_negotiated(conn))
330 		flags |= NTLMSSP_NEGOTIATE_SEAL;
331 
332 	if (cflags & NTLMSSP_NEGOTIATE_ALWAYS_SIGN)
333 		flags |= NTLMSSP_NEGOTIATE_ALWAYS_SIGN;
334 
335 	if (cflags & NTLMSSP_REQUEST_TARGET)
336 		flags |= NTLMSSP_REQUEST_TARGET;
337 
338 	if (conn->use_spnego &&
339 	    (cflags & NTLMSSP_NEGOTIATE_EXTENDED_SEC))
340 		flags |= NTLMSSP_NEGOTIATE_EXTENDED_SEC;
341 
342 	if (cflags & NTLMSSP_NEGOTIATE_KEY_XCH)
343 		flags |= NTLMSSP_NEGOTIATE_KEY_XCH;
344 
345 	chgblob->NegotiateFlags = cpu_to_le32(flags);
346 	len = strlen(ksmbd_netbios_name());
347 	name = kmalloc(2 + UNICODE_LEN(len), KSMBD_DEFAULT_GFP);
348 	if (!name)
349 		return -ENOMEM;
350 
351 	conv_len = smb_strtoUTF16((__le16 *)name, ksmbd_netbios_name(), len,
352 				  conn->local_nls);
353 	if (conv_len < 0 || conv_len > len) {
354 		kfree(name);
355 		return -EINVAL;
356 	}
357 
358 	uni_len = UNICODE_LEN(conv_len);
359 
360 	blob_off = sizeof(struct challenge_message);
361 	blob_len = blob_off + uni_len;
362 
363 	chgblob->TargetName.Length = cpu_to_le16(uni_len);
364 	chgblob->TargetName.MaximumLength = cpu_to_le16(uni_len);
365 	chgblob->TargetName.BufferOffset = cpu_to_le32(blob_off);
366 
367 	/* Initialize random conn challenge */
368 	get_random_bytes(conn->ntlmssp.cryptkey, sizeof(__u64));
369 	memcpy(chgblob->Challenge, conn->ntlmssp.cryptkey,
370 	       CIFS_CRYPTO_KEY_SIZE);
371 
372 	/* Add Target Information to security buffer */
373 	chgblob->TargetInfoArray.BufferOffset = cpu_to_le32(blob_len);
374 
375 	target_name = (__u8 *)chgblob + blob_off;
376 	memcpy(target_name, name, uni_len);
377 	tinfo = (struct target_info *)(target_name + uni_len);
378 
379 	chgblob->TargetInfoArray.Length = 0;
380 	/* Add target info list for NetBIOS/DNS settings */
381 	for (type = NTLMSSP_AV_NB_COMPUTER_NAME;
382 	     type <= NTLMSSP_AV_DNS_DOMAIN_NAME; type++) {
383 		tinfo->Type = cpu_to_le16(type);
384 		tinfo->Length = cpu_to_le16(uni_len);
385 		memcpy(tinfo->Content, name, uni_len);
386 		tinfo = (struct target_info *)((char *)tinfo + 4 + uni_len);
387 		target_info_len += 4 + uni_len;
388 	}
389 
390 	/* Add terminator subblock */
391 	tinfo->Type = 0;
392 	tinfo->Length = 0;
393 	target_info_len += 4;
394 
395 	chgblob->TargetInfoArray.Length = cpu_to_le16(target_info_len);
396 	chgblob->TargetInfoArray.MaximumLength = cpu_to_le16(target_info_len);
397 	blob_len += target_info_len;
398 	kfree(name);
399 	ksmbd_debug(AUTH, "NTLMSSP SecurityBufferLength %d\n", blob_len);
400 	return blob_len;
401 }
402 
403 #ifdef CONFIG_SMB_SERVER_KERBEROS5
404 int ksmbd_krb5_authenticate(struct ksmbd_session *sess, char *in_blob,
405 			    int in_len, char *out_blob, int *out_len,
406 			    char *sess_key)
407 {
408 	struct ksmbd_spnego_authen_response *resp;
409 	struct ksmbd_login_response_ext *resp_ext = NULL;
410 	struct ksmbd_user *user = NULL;
411 	int retval;
412 
413 	resp = ksmbd_ipc_spnego_authen_request(in_blob, in_len);
414 	if (!resp) {
415 		ksmbd_debug(AUTH, "SPNEGO_AUTHEN_REQUEST failure\n");
416 		return -EINVAL;
417 	}
418 
419 	if (!(resp->login_response.status & KSMBD_USER_FLAG_OK)) {
420 		ksmbd_debug(AUTH, "krb5 authentication failure\n");
421 		retval = -EPERM;
422 		goto out;
423 	}
424 
425 	if (*out_len <= resp->spnego_blob_len) {
426 		ksmbd_debug(AUTH, "buf len %d, but blob len %d\n",
427 			    *out_len, resp->spnego_blob_len);
428 		retval = -EINVAL;
429 		goto out;
430 	}
431 
432 	if (resp->session_key_len > sizeof(sess->sess_key)) {
433 		ksmbd_debug(AUTH, "session key is too long\n");
434 		retval = -EINVAL;
435 		goto out;
436 	}
437 
438 	if (resp->login_response.status & KSMBD_USER_FLAG_EXTENSION)
439 		resp_ext = ksmbd_ipc_login_request_ext(resp->login_response.account);
440 
441 	user = ksmbd_alloc_user(&resp->login_response, resp_ext);
442 	if (!user) {
443 		ksmbd_debug(AUTH, "login failure\n");
444 		retval = -ENOMEM;
445 		goto out;
446 	}
447 
448 	if (!sess->user) {
449 		/* First successful authentication */
450 		sess->user = user;
451 	} else {
452 		if (!ksmbd_compare_user(sess->user, user)) {
453 			ksmbd_debug(AUTH, "different user tried to reuse session\n");
454 			retval = -EKEYREJECTED;
455 			ksmbd_free_user(user);
456 			goto out;
457 		}
458 		ksmbd_free_user(user);
459 	}
460 
461 	memcpy(sess_key, resp->payload, resp->session_key_len);
462 	memcpy(out_blob, resp->payload + resp->session_key_len,
463 	       resp->spnego_blob_len);
464 	*out_len = resp->spnego_blob_len;
465 	retval = 0;
466 out:
467 	kvfree(resp);
468 	return retval;
469 }
470 #else
471 int ksmbd_krb5_authenticate(struct ksmbd_session *sess, char *in_blob,
472 			    int in_len, char *out_blob, int *out_len,
473 			    char *sess_key)
474 {
475 	return -EOPNOTSUPP;
476 }
477 #endif
478 
479 /**
480  * ksmbd_sign_smb2_pdu() - function to generate packet signing
481  * @conn:	connection
482  * @key:	signing key
483  * @iov:        buffer iov array
484  * @n_vec:	number of iovecs
485  * @sig:	signature value generated for client request packet
486  *
487  */
488 void ksmbd_sign_smb2_pdu(struct ksmbd_conn *conn, char *key, struct kvec *iov,
489 			 int n_vec, char *sig)
490 {
491 	struct hmac_sha256_ctx ctx;
492 	int i;
493 
494 	hmac_sha256_init_usingrawkey(&ctx, key, SMB2_NTLMV2_SESSKEY_SIZE);
495 	for (i = 0; i < n_vec; i++)
496 		hmac_sha256_update(&ctx, iov[i].iov_base, iov[i].iov_len);
497 	hmac_sha256_final(&ctx, sig);
498 }
499 
500 /**
501  * ksmbd_sign_smb3_pdu() - function to generate packet signing
502  * @conn:	connection
503  * @key:	signing key
504  * @iov:        buffer iov array
505  * @n_vec:	number of iovecs
506  * @sig:	signature value generated for client request packet
507  *
508  */
509 void ksmbd_sign_smb3_pdu(struct ksmbd_conn *conn, char *key, struct kvec *iov,
510 			 int n_vec, char *sig)
511 {
512 	struct aes_cmac_key cmac_key;
513 	struct aes_cmac_ctx cmac_ctx;
514 	int i;
515 
516 	/* This cannot fail, since we always pass a valid key length. */
517 	static_assert(SMB2_CMACAES_SIZE == AES_KEYSIZE_128);
518 	aes_cmac_preparekey(&cmac_key, key, SMB2_CMACAES_SIZE);
519 
520 	aes_cmac_init(&cmac_ctx, &cmac_key);
521 	for (i = 0; i < n_vec; i++)
522 		aes_cmac_update(&cmac_ctx, iov[i].iov_base, iov[i].iov_len);
523 	aes_cmac_final(&cmac_ctx, sig);
524 }
525 
526 struct derivation {
527 	struct kvec label;
528 	struct kvec context;
529 	bool binding;
530 };
531 
532 static void generate_key(struct ksmbd_conn *conn, const char *sess_key,
533 			 struct kvec label, struct kvec context, __u8 *key,
534 			 unsigned int key_size)
535 {
536 	unsigned char zero = 0x0;
537 	__u8 i[4] = {0, 0, 0, 1};
538 	__u8 L128[4] = {0, 0, 0, 128};
539 	__u8 L256[4] = {0, 0, 1, 0};
540 	unsigned char prfhash[SMB2_HMACSHA256_SIZE];
541 	struct hmac_sha256_ctx ctx;
542 
543 	hmac_sha256_init_usingrawkey(&ctx, sess_key,
544 				     SMB2_NTLMV2_SESSKEY_SIZE);
545 	hmac_sha256_update(&ctx, i, 4);
546 	hmac_sha256_update(&ctx, label.iov_base, label.iov_len);
547 	hmac_sha256_update(&ctx, &zero, 1);
548 	hmac_sha256_update(&ctx, context.iov_base, context.iov_len);
549 
550 	if (key_size == SMB3_ENC_DEC_KEY_SIZE &&
551 	    (conn->cipher_type == SMB2_ENCRYPTION_AES256_CCM ||
552 	     conn->cipher_type == SMB2_ENCRYPTION_AES256_GCM))
553 		hmac_sha256_update(&ctx, L256, 4);
554 	else
555 		hmac_sha256_update(&ctx, L128, 4);
556 
557 	hmac_sha256_final(&ctx, prfhash);
558 	memcpy(key, prfhash, key_size);
559 }
560 
561 static int generate_smb3signingkey(struct ksmbd_session *sess,
562 				   struct ksmbd_conn *conn,
563 				   const struct derivation *signing)
564 {
565 	struct channel *chann;
566 	char *key, *sess_key;
567 
568 	chann = lookup_chann_list(sess, conn);
569 	if (!chann)
570 		return 0;
571 
572 	if (conn->dialect >= SMB30_PROT_ID && signing->binding) {
573 		key = chann->smb3signingkey;
574 		sess_key = chann->sess_key;
575 	} else {
576 		key = sess->smb3signingkey;
577 		sess_key = sess->sess_key;
578 	}
579 
580 	generate_key(conn, sess_key, signing->label, signing->context, key,
581 		     SMB3_SIGN_KEY_SIZE);
582 
583 	if (!(conn->dialect >= SMB30_PROT_ID && signing->binding))
584 		memcpy(chann->smb3signingkey, key, SMB3_SIGN_KEY_SIZE);
585 
586 	ksmbd_debug(AUTH, "generated SMB3 signing key\n");
587 	ksmbd_debug(AUTH, "Session Id    %llu\n", sess->id);
588 	return 0;
589 }
590 
591 int ksmbd_gen_smb30_signingkey(struct ksmbd_session *sess,
592 			       struct ksmbd_conn *conn)
593 {
594 	struct derivation d;
595 
596 	d.label.iov_base = "SMB2AESCMAC";
597 	d.label.iov_len = 12;
598 	d.context.iov_base = "SmbSign";
599 	d.context.iov_len = 8;
600 	d.binding = conn->binding;
601 
602 	return generate_smb3signingkey(sess, conn, &d);
603 }
604 
605 int ksmbd_gen_smb311_signingkey(struct ksmbd_session *sess,
606 				struct ksmbd_conn *conn)
607 {
608 	struct derivation d;
609 
610 	d.label.iov_base = "SMBSigningKey";
611 	d.label.iov_len = 14;
612 	if (conn->binding) {
613 		struct preauth_session *preauth_sess;
614 
615 		preauth_sess = ksmbd_preauth_session_lookup(conn, sess->id);
616 		if (!preauth_sess)
617 			return -ENOENT;
618 		d.context.iov_base = preauth_sess->Preauth_HashValue;
619 	} else {
620 		d.context.iov_base = sess->Preauth_HashValue;
621 	}
622 	d.context.iov_len = 64;
623 	d.binding = conn->binding;
624 
625 	return generate_smb3signingkey(sess, conn, &d);
626 }
627 
628 struct derivation_twin {
629 	struct derivation encryption;
630 	struct derivation decryption;
631 };
632 
633 static void generate_smb3encryptionkey(struct ksmbd_conn *conn,
634 				       struct ksmbd_session *sess,
635 				       const struct derivation_twin *ptwin)
636 {
637 	generate_key(conn, sess->sess_key, ptwin->encryption.label,
638 		     ptwin->encryption.context, sess->smb3encryptionkey,
639 		     SMB3_ENC_DEC_KEY_SIZE);
640 
641 	generate_key(conn, sess->sess_key, ptwin->decryption.label,
642 		     ptwin->decryption.context,
643 		     sess->smb3decryptionkey, SMB3_ENC_DEC_KEY_SIZE);
644 
645 	ksmbd_debug(AUTH, "generated SMB3 encryption/decryption keys\n");
646 	ksmbd_debug(AUTH, "Cipher type   %d\n", conn->cipher_type);
647 	ksmbd_debug(AUTH, "Session Id    %llu\n", sess->id);
648 }
649 
650 void ksmbd_gen_smb30_encryptionkey(struct ksmbd_conn *conn,
651 				   struct ksmbd_session *sess)
652 {
653 	struct derivation_twin twin;
654 	struct derivation *d;
655 
656 	d = &twin.encryption;
657 	d->label.iov_base = "SMB2AESCCM";
658 	d->label.iov_len = 11;
659 	d->context.iov_base = "ServerOut";
660 	d->context.iov_len = 10;
661 
662 	d = &twin.decryption;
663 	d->label.iov_base = "SMB2AESCCM";
664 	d->label.iov_len = 11;
665 	d->context.iov_base = "ServerIn ";
666 	d->context.iov_len = 10;
667 
668 	generate_smb3encryptionkey(conn, sess, &twin);
669 }
670 
671 void ksmbd_gen_smb311_encryptionkey(struct ksmbd_conn *conn,
672 				    struct ksmbd_session *sess)
673 {
674 	struct derivation_twin twin;
675 	struct derivation *d;
676 
677 	d = &twin.encryption;
678 	d->label.iov_base = "SMBS2CCipherKey";
679 	d->label.iov_len = 16;
680 	d->context.iov_base = sess->Preauth_HashValue;
681 	d->context.iov_len = 64;
682 
683 	d = &twin.decryption;
684 	d->label.iov_base = "SMBC2SCipherKey";
685 	d->label.iov_len = 16;
686 	d->context.iov_base = sess->Preauth_HashValue;
687 	d->context.iov_len = 64;
688 
689 	generate_smb3encryptionkey(conn, sess, &twin);
690 }
691 
692 int ksmbd_gen_preauth_integrity_hash(struct ksmbd_conn *conn, char *buf,
693 				     __u8 *pi_hash)
694 {
695 	struct smb2_hdr *rcv_hdr = smb_get_msg(buf);
696 	char *all_bytes_msg = (char *)&rcv_hdr->ProtocolId;
697 	int msg_size = get_rfc1002_len(buf);
698 	struct sha512_ctx sha_ctx;
699 
700 	if (conn->preauth_info->Preauth_HashId !=
701 	    SMB2_PREAUTH_INTEGRITY_SHA512)
702 		return -EINVAL;
703 
704 	sha512_init(&sha_ctx);
705 	sha512_update(&sha_ctx, pi_hash, 64);
706 	sha512_update(&sha_ctx, all_bytes_msg, msg_size);
707 	sha512_final(&sha_ctx, pi_hash);
708 	return 0;
709 }
710 
711 static int ksmbd_get_encryption_key(struct ksmbd_work *work, __u64 ses_id,
712 				    int enc, u8 *key)
713 {
714 	struct ksmbd_session *sess;
715 	u8 *ses_enc_key;
716 
717 	if (enc)
718 		sess = work->sess;
719 	else
720 		sess = ksmbd_session_lookup_all(work->conn, ses_id);
721 	if (!sess)
722 		return -EINVAL;
723 
724 	ses_enc_key = enc ? sess->smb3encryptionkey :
725 		sess->smb3decryptionkey;
726 	memcpy(key, ses_enc_key, SMB3_ENC_DEC_KEY_SIZE);
727 	if (!enc)
728 		ksmbd_user_session_put(sess);
729 
730 	return 0;
731 }
732 
733 static inline void smb2_sg_set_buf(struct scatterlist *sg, const void *buf,
734 				   unsigned int buflen)
735 {
736 	void *addr;
737 
738 	if (is_vmalloc_addr(buf))
739 		addr = vmalloc_to_page(buf);
740 	else
741 		addr = virt_to_page(buf);
742 	sg_set_page(sg, addr, buflen, offset_in_page(buf));
743 }
744 
745 static struct scatterlist *ksmbd_init_sg(struct kvec *iov, unsigned int nvec,
746 					 u8 *sign)
747 {
748 	struct scatterlist *sg;
749 	unsigned int assoc_data_len = sizeof(struct smb2_transform_hdr) - 20;
750 	int i, *nr_entries, total_entries = 0, sg_idx = 0;
751 
752 	if (!nvec)
753 		return NULL;
754 
755 	nr_entries = kzalloc_objs(int, nvec, KSMBD_DEFAULT_GFP);
756 	if (!nr_entries)
757 		return NULL;
758 
759 	for (i = 0; i < nvec - 1; i++) {
760 		unsigned long kaddr = (unsigned long)iov[i + 1].iov_base;
761 
762 		if (is_vmalloc_addr(iov[i + 1].iov_base)) {
763 			nr_entries[i] = ((kaddr + iov[i + 1].iov_len +
764 					PAGE_SIZE - 1) >> PAGE_SHIFT) -
765 				(kaddr >> PAGE_SHIFT);
766 		} else {
767 			nr_entries[i]++;
768 		}
769 		total_entries += nr_entries[i];
770 	}
771 
772 	/* Add two entries for transform header and signature */
773 	total_entries += 2;
774 
775 	sg = kmalloc_objs(struct scatterlist, total_entries, KSMBD_DEFAULT_GFP);
776 	if (!sg) {
777 		kfree(nr_entries);
778 		return NULL;
779 	}
780 
781 	sg_init_table(sg, total_entries);
782 	smb2_sg_set_buf(&sg[sg_idx++], iov[0].iov_base + 24, assoc_data_len);
783 	for (i = 0; i < nvec - 1; i++) {
784 		void *data = iov[i + 1].iov_base;
785 		int len = iov[i + 1].iov_len;
786 
787 		if (is_vmalloc_addr(data)) {
788 			int j, offset = offset_in_page(data);
789 
790 			for (j = 0; j < nr_entries[i]; j++) {
791 				unsigned int bytes = PAGE_SIZE - offset;
792 
793 				if (!len)
794 					break;
795 
796 				if (bytes > len)
797 					bytes = len;
798 
799 				sg_set_page(&sg[sg_idx++],
800 					    vmalloc_to_page(data), bytes,
801 					    offset_in_page(data));
802 
803 				data += bytes;
804 				len -= bytes;
805 				offset = 0;
806 			}
807 		} else {
808 			sg_set_page(&sg[sg_idx++], virt_to_page(data), len,
809 				    offset_in_page(data));
810 		}
811 	}
812 	smb2_sg_set_buf(&sg[sg_idx], sign, SMB2_SIGNATURE_SIZE);
813 	kfree(nr_entries);
814 	return sg;
815 }
816 
817 int ksmbd_crypt_message(struct ksmbd_work *work, struct kvec *iov,
818 			unsigned int nvec, int enc)
819 {
820 	struct ksmbd_conn *conn = work->conn;
821 	struct smb2_transform_hdr *tr_hdr = smb_get_msg(iov[0].iov_base);
822 	unsigned int assoc_data_len = sizeof(struct smb2_transform_hdr) - 20;
823 	int rc;
824 	DECLARE_CRYPTO_WAIT(wait);
825 	struct scatterlist *sg;
826 	u8 sign[SMB2_SIGNATURE_SIZE] = {};
827 	u8 key[SMB3_ENC_DEC_KEY_SIZE];
828 	struct aead_request *req;
829 	char *iv;
830 	unsigned int iv_len;
831 	struct crypto_aead *tfm;
832 	unsigned int crypt_len = le32_to_cpu(tr_hdr->OriginalMessageSize);
833 	struct ksmbd_crypto_ctx *ctx;
834 
835 	rc = ksmbd_get_encryption_key(work,
836 				      le64_to_cpu(tr_hdr->SessionId),
837 				      enc,
838 				      key);
839 	if (rc) {
840 		pr_err("Could not get %scryption key\n", enc ? "en" : "de");
841 		return rc;
842 	}
843 
844 	if (conn->cipher_type == SMB2_ENCRYPTION_AES128_GCM ||
845 	    conn->cipher_type == SMB2_ENCRYPTION_AES256_GCM)
846 		ctx = ksmbd_crypto_ctx_find_gcm();
847 	else
848 		ctx = ksmbd_crypto_ctx_find_ccm();
849 	if (!ctx) {
850 		pr_err("crypto alloc failed\n");
851 		return -ENOMEM;
852 	}
853 
854 	if (conn->cipher_type == SMB2_ENCRYPTION_AES128_GCM ||
855 	    conn->cipher_type == SMB2_ENCRYPTION_AES256_GCM)
856 		tfm = CRYPTO_GCM(ctx);
857 	else
858 		tfm = CRYPTO_CCM(ctx);
859 
860 	if (conn->cipher_type == SMB2_ENCRYPTION_AES256_CCM ||
861 	    conn->cipher_type == SMB2_ENCRYPTION_AES256_GCM)
862 		rc = crypto_aead_setkey(tfm, key, SMB3_GCM256_CRYPTKEY_SIZE);
863 	else
864 		rc = crypto_aead_setkey(tfm, key, SMB3_GCM128_CRYPTKEY_SIZE);
865 	if (rc) {
866 		pr_err("Failed to set aead key %d\n", rc);
867 		goto free_ctx;
868 	}
869 
870 	rc = crypto_aead_setauthsize(tfm, SMB2_SIGNATURE_SIZE);
871 	if (rc) {
872 		pr_err("Failed to set authsize %d\n", rc);
873 		goto free_ctx;
874 	}
875 
876 	req = aead_request_alloc(tfm, KSMBD_DEFAULT_GFP);
877 	if (!req) {
878 		rc = -ENOMEM;
879 		goto free_ctx;
880 	}
881 
882 	if (!enc) {
883 		memcpy(sign, &tr_hdr->Signature, SMB2_SIGNATURE_SIZE);
884 		crypt_len += SMB2_SIGNATURE_SIZE;
885 	}
886 
887 	sg = ksmbd_init_sg(iov, nvec, sign);
888 	if (!sg) {
889 		pr_err("Failed to init sg\n");
890 		rc = -ENOMEM;
891 		goto free_req;
892 	}
893 
894 	iv_len = crypto_aead_ivsize(tfm);
895 	iv = kzalloc(iv_len, KSMBD_DEFAULT_GFP);
896 	if (!iv) {
897 		rc = -ENOMEM;
898 		goto free_sg;
899 	}
900 
901 	if (conn->cipher_type == SMB2_ENCRYPTION_AES128_GCM ||
902 	    conn->cipher_type == SMB2_ENCRYPTION_AES256_GCM) {
903 		memcpy(iv, (char *)tr_hdr->Nonce, SMB3_AES_GCM_NONCE);
904 	} else {
905 		iv[0] = 3;
906 		memcpy(iv + 1, (char *)tr_hdr->Nonce, SMB3_AES_CCM_NONCE);
907 	}
908 
909 	aead_request_set_crypt(req, sg, sg, crypt_len, iv);
910 	aead_request_set_ad(req, assoc_data_len);
911 	aead_request_set_callback(req, CRYPTO_TFM_REQ_MAY_BACKLOG |
912 				  CRYPTO_TFM_REQ_MAY_SLEEP,
913 				  crypto_req_done, &wait);
914 
915 	rc = crypto_wait_req(enc ? crypto_aead_encrypt(req) :
916 			     crypto_aead_decrypt(req), &wait);
917 	if (rc)
918 		goto free_iv;
919 
920 	if (enc)
921 		memcpy(&tr_hdr->Signature, sign, SMB2_SIGNATURE_SIZE);
922 
923 free_iv:
924 	kfree(iv);
925 free_sg:
926 	kfree(sg);
927 free_req:
928 	aead_request_free(req);
929 free_ctx:
930 	ksmbd_release_crypto_ctx(ctx);
931 	return rc;
932 }
933