xref: /linux/fs/smb/server/mgmt/user_session.c (revision d7fd1f98607f2cd358e583f9548bdf0173090a86)
1 // SPDX-License-Identifier: GPL-2.0-or-later
2 /*
3  *   Copyright (C) 2018 Samsung Electronics Co., Ltd.
4  */
5 
6 #include <linux/list.h>
7 #include <linux/slab.h>
8 #include <linux/rwsem.h>
9 #include <linux/xarray.h>
10 
11 #include "ksmbd_ida.h"
12 #include "user_session.h"
13 #include "user_config.h"
14 #include "tree_connect.h"
15 #include "share_config.h"
16 #include "../transport_ipc.h"
17 #include "../connection.h"
18 #include "../vfs_cache.h"
19 #include "../misc.h"
20 #include "../stats.h"
21 
22 static DEFINE_IDA(session_ida);
23 
24 #define SESSION_HASH_BITS		12
25 #define KSMBD_MAX_PENDING_SESSIONS	1
26 static DEFINE_HASHTABLE(sessions_table, SESSION_HASH_BITS);
27 static DECLARE_RWSEM(sessions_table_lock);
28 
29 struct ksmbd_session_rpc {
30 	int			id;
31 	unsigned int		method;
32 };
33 
34 #ifdef CONFIG_PROC_FS
35 
36 static const struct ksmbd_const_name ksmbd_sess_cap_const_names[] = {
37 	{SMB2_GLOBAL_CAP_DFS, "dfs"},
38 	{SMB2_GLOBAL_CAP_LEASING, "lease"},
39 	{SMB2_GLOBAL_CAP_LARGE_MTU, "large-mtu"},
40 	{SMB2_GLOBAL_CAP_MULTI_CHANNEL, "multi-channel"},
41 	{SMB2_GLOBAL_CAP_PERSISTENT_HANDLES, "persistent-handles"},
42 	{SMB2_GLOBAL_CAP_DIRECTORY_LEASING, "dir-lease"},
43 	{SMB2_GLOBAL_CAP_ENCRYPTION, "encryption"}
44 };
45 
46 static const struct ksmbd_const_name ksmbd_cipher_const_names[] = {
47 	{le16_to_cpu(SMB2_ENCRYPTION_AES128_CCM), "aes128-ccm"},
48 	{le16_to_cpu(SMB2_ENCRYPTION_AES128_GCM), "aes128-gcm"},
49 	{le16_to_cpu(SMB2_ENCRYPTION_AES256_CCM), "aes256-ccm"},
50 	{le16_to_cpu(SMB2_ENCRYPTION_AES256_GCM), "aes256-gcm"},
51 };
52 
53 static const struct ksmbd_const_name ksmbd_signing_const_names[] = {
54 	{SIGNING_ALG_HMAC_SHA256, "hmac-sha256"},
55 	{SIGNING_ALG_AES_CMAC, "aes-cmac"},
56 	{SIGNING_ALG_AES_GMAC, "aes-gmac"},
57 };
58 
59 static const char *session_state_string(struct ksmbd_session *session)
60 {
61 	switch (session->state) {
62 	case SMB2_SESSION_VALID:
63 		return "valid";
64 	case SMB2_SESSION_IN_PROGRESS:
65 		return "progress";
66 	case SMB2_SESSION_EXPIRED:
67 		return "expired";
68 	default:
69 		return "";
70 	}
71 }
72 
73 static const char *session_user_name(struct ksmbd_session *session)
74 {
75 	if (user_guest(session->user))
76 		return "(Guest)";
77 	else if (ksmbd_anonymous_user(session->user))
78 		return "(Anonymous)";
79 	return session->user->name;
80 }
81 
82 static const char *session_account_type(struct ksmbd_session *session)
83 {
84 	if (user_guest(session->user))
85 		return "guest";
86 	if (ksmbd_anonymous_user(session->user))
87 		return "anonymous";
88 	return "user";
89 }
90 
91 static unsigned int session_open_file_count(struct ksmbd_session *session)
92 {
93 	struct ksmbd_file *fp;
94 	unsigned int count = 0;
95 	unsigned int id;
96 
97 	read_lock(&session->file_table.lock);
98 	idr_for_each_entry(session->file_table.idr, fp, id)
99 		count++;
100 	read_unlock(&session->file_table.lock);
101 	return count;
102 }
103 
104 static int show_proc_session(struct seq_file *m, void *v)
105 {
106 	struct ksmbd_session *sess;
107 	struct ksmbd_tree_connect *tree_conn;
108 	struct ksmbd_share_config *share_conf;
109 	struct channel *chan;
110 	unsigned long id;
111 	int i = 0;
112 
113 	sess = (struct ksmbd_session *)m->private;
114 	ksmbd_user_session_get(sess);
115 
116 	seq_printf(m, "user:\t%s\n", session_user_name(sess));
117 	seq_printf(m, "account_type:\t%s\n",
118 		   session_account_type(sess));
119 	seq_printf(m, "id:\t%llu\n", sess->id);
120 	seq_printf(m, "state:\t%s\n", session_state_string(sess));
121 	seq_printf(m, "dialect:\t0x%04x\n", sess->dialect);
122 	seq_printf(m, "last_active_seconds:\t%lu\n",
123 		   jiffies_to_msecs(jiffies - sess->last_active) / MSEC_PER_SEC);
124 	seq_printf(m, "open_files:\t%u\n",
125 		   session_open_file_count(sess));
126 
127 	i = 0;
128 	down_read(&sess->chann_lock);
129 	xa_for_each(&sess->ksmbd_chann_list, id, chan) {
130 		const char *name;
131 
132 #if IS_ENABLED(CONFIG_IPV6)
133 		if (chan->conn->inet_addr)
134 			seq_printf(m, "client:\t%pI4\n",
135 					&chan->conn->inet_addr);
136 		else
137 			seq_printf(m, "client:\t%pI6c\n",
138 					&chan->conn->inet6_addr);
139 #else
140 		seq_printf(m, "client:\t%pI4\n",
141 				&chan->conn->inet_addr);
142 #endif
143 		seq_puts(m, "capabilities:\t");
144 		ksmbd_proc_show_flag_names(m,
145 				ksmbd_sess_cap_const_names,
146 				ARRAY_SIZE(ksmbd_sess_cap_const_names),
147 				chan->conn->vals->req_capabilities);
148 		seq_putc(m, '\n');
149 		seq_printf(m, "posix_extensions:\t%s\n",
150 			   chan->conn->posix_ext_supported ? "yes" : "no");
151 
152 		if (sess->sign) {
153 			unsigned int algorithm =
154 				le16_to_cpu(chan->conn->signing_algorithm);
155 
156 			name = ksmbd_proc_const_name(ksmbd_signing_const_names,
157 						     ARRAY_SIZE(ksmbd_signing_const_names),
158 						     algorithm);
159 			if (name)
160 				seq_printf(m, "signing:\t%s\n", name);
161 			else
162 				seq_printf(m, "signing:\t0x%04x\n",
163 					   algorithm);
164 		}
165 		if (sess->enc) {
166 			unsigned int cipher = le16_to_cpu(chan->conn->cipher_type);
167 
168 			name = ksmbd_proc_const_name(ksmbd_cipher_const_names,
169 						     ARRAY_SIZE(ksmbd_cipher_const_names),
170 						     cipher);
171 			if (name)
172 				seq_printf(m, "encryption:\t%s\n", name);
173 			else
174 				seq_printf(m, "encryption:\t0x%04x\n",
175 					   cipher);
176 		}
177 		i++;
178 	}
179 	up_read(&sess->chann_lock);
180 
181 	seq_printf(m, "channels:\t%d\n", i);
182 
183 	i = 0;
184 	down_read(&sess->tree_conns_lock);
185 	xa_for_each(&sess->tree_conns, id, tree_conn) {
186 		share_conf = tree_conn->share_conf;
187 		seq_printf(m, "share:\t%s\n", share_conf->name);
188 		seq_printf(m, "tree_id:\t%d\n", tree_conn->id);
189 		seq_printf(m, "share_type:\t%s\n",
190 			   test_share_config_flag(share_conf, KSMBD_SHARE_FLAG_PIPE) ?
191 			   "pipe" : "disk");
192 		i++;
193 	}
194 	up_read(&sess->tree_conns_lock);
195 	seq_printf(m, "tree_connects:\t%d\n", i);
196 
197 	ksmbd_user_session_put(sess);
198 	return 0;
199 }
200 
201 static int create_proc_session(struct ksmbd_session *sess)
202 {
203 	char name[30];
204 
205 	snprintf(name, sizeof(name), "sessions/%llu", sess->id);
206 	sess->proc_entry = ksmbd_proc_create(name,
207 					     show_proc_session, sess);
208 	if (!sess->proc_entry)
209 		return -ENOMEM;
210 	return 0;
211 }
212 
213 static void delete_proc_session(struct ksmbd_session *sess)
214 {
215 	if (sess->proc_entry)
216 		proc_remove(sess->proc_entry);
217 }
218 
219 static int show_proc_sessions(struct seq_file *m, void *v)
220 {
221 	struct ksmbd_session *session;
222 	struct channel *chan;
223 	int i;
224 	unsigned long id;
225 
226 	down_read(&sessions_table_lock);
227 	hash_for_each(sessions_table, i, session, hlist) {
228 		down_read(&session->chann_lock);
229 		xa_for_each(&session->ksmbd_chann_list, id, chan) {
230 			down_read(&chan->conn->session_lock);
231 			ksmbd_user_session_get(session);
232 
233 #if IS_ENABLED(CONFIG_IPV6)
234 			if (!chan->conn->inet_addr)
235 				seq_printf(m, "client:\t%pI6c\n", &chan->conn->inet6_addr);
236 			else
237 #endif
238 				seq_printf(m, "client:\t%pI4\n", &chan->conn->inet_addr);
239 			seq_printf(m, "user:\t%s\n", session_user_name(session));
240 			seq_printf(m, "id:\t%llu\n", session->id);
241 			seq_printf(m, "state:\t%s\n\n",
242 				   session_state_string(session));
243 
244 			ksmbd_user_session_put(session);
245 			up_read(&chan->conn->session_lock);
246 		}
247 		up_read(&session->chann_lock);
248 	}
249 	up_read(&sessions_table_lock);
250 	return 0;
251 }
252 
253 int create_proc_sessions(void)
254 {
255 	if (!ksmbd_proc_create("sessions/sessions",
256 			       show_proc_sessions, NULL))
257 		return -ENOMEM;
258 	return 0;
259 }
260 #else
261 int create_proc_sessions(void) { return 0; }
262 static int create_proc_session(struct ksmbd_session *sess) { return 0; }
263 static void delete_proc_session(struct ksmbd_session *sess) {}
264 #endif
265 
266 static void free_channel_list(struct ksmbd_session *sess)
267 {
268 	struct channel *chann;
269 	unsigned long index;
270 
271 	down_write(&sess->chann_lock);
272 	xa_for_each(&sess->ksmbd_chann_list, index, chann) {
273 		xa_erase(&sess->ksmbd_chann_list, index);
274 		kfree_sensitive(chann);
275 	}
276 
277 	xa_destroy(&sess->ksmbd_chann_list);
278 	up_write(&sess->chann_lock);
279 }
280 
281 static void __session_rpc_close(struct ksmbd_session *sess,
282 				struct ksmbd_session_rpc *entry)
283 {
284 	struct ksmbd_rpc_command *resp;
285 
286 	resp = ksmbd_rpc_close(sess, entry->id);
287 	if (!resp)
288 		pr_err("Unable to close RPC pipe %d\n", entry->id);
289 
290 	kvfree(resp);
291 	ksmbd_rpc_id_free(entry->id);
292 	kfree(entry);
293 }
294 
295 static void ksmbd_session_rpc_clear_list(struct ksmbd_session *sess)
296 {
297 	struct ksmbd_session_rpc *entry;
298 	long index;
299 
300 	down_write(&sess->rpc_lock);
301 	xa_for_each(&sess->rpc_handle_list, index, entry) {
302 		xa_erase(&sess->rpc_handle_list, index);
303 		__session_rpc_close(sess, entry);
304 	}
305 	up_write(&sess->rpc_lock);
306 
307 	xa_destroy(&sess->rpc_handle_list);
308 }
309 
310 static int __rpc_method(char *rpc_name)
311 {
312 	if (!strcmp(rpc_name, "\\srvsvc") || !strcmp(rpc_name, "srvsvc"))
313 		return KSMBD_RPC_SRVSVC_METHOD_INVOKE;
314 
315 	if (!strcmp(rpc_name, "\\wkssvc") || !strcmp(rpc_name, "wkssvc"))
316 		return KSMBD_RPC_WKSSVC_METHOD_INVOKE;
317 
318 	if (!strcmp(rpc_name, "LANMAN") || !strcmp(rpc_name, "lanman"))
319 		return KSMBD_RPC_RAP_METHOD;
320 
321 	if (!strcmp(rpc_name, "\\samr") || !strcmp(rpc_name, "samr"))
322 		return KSMBD_RPC_SAMR_METHOD_INVOKE;
323 
324 	if (!strcmp(rpc_name, "\\lsarpc") || !strcmp(rpc_name, "lsarpc"))
325 		return KSMBD_RPC_LSARPC_METHOD_INVOKE;
326 
327 	if (!strcmp(rpc_name, "\\mdssvc") || !strcmp(rpc_name, "mdssvc"))
328 		return -ENOENT;
329 
330 	pr_err("Unsupported RPC: %s\n", rpc_name);
331 	return -ENOENT;
332 }
333 
334 int ksmbd_session_rpc_open(struct ksmbd_session *sess, char *rpc_name)
335 {
336 	struct ksmbd_session_rpc *entry, *old;
337 	struct ksmbd_rpc_command *resp;
338 	int method, id;
339 
340 	method = __rpc_method(rpc_name);
341 	if (method < 0)
342 		return method;
343 
344 	entry = kzalloc_obj(struct ksmbd_session_rpc, KSMBD_DEFAULT_GFP);
345 	if (!entry)
346 		return -ENOMEM;
347 
348 	entry->method = method;
349 	entry->id = id = ksmbd_ipc_id_alloc();
350 	if (id < 0)
351 		goto free_entry;
352 
353 	down_write(&sess->rpc_lock);
354 	old = xa_store(&sess->rpc_handle_list, id, entry, KSMBD_DEFAULT_GFP);
355 	if (xa_is_err(old)) {
356 		up_write(&sess->rpc_lock);
357 		goto free_id;
358 	}
359 
360 	resp = ksmbd_rpc_open(sess, id);
361 	if (!resp) {
362 		xa_erase(&sess->rpc_handle_list, entry->id);
363 		up_write(&sess->rpc_lock);
364 		goto free_id;
365 	}
366 
367 	up_write(&sess->rpc_lock);
368 	kvfree(resp);
369 	return id;
370 free_id:
371 	ksmbd_rpc_id_free(entry->id);
372 free_entry:
373 	kfree(entry);
374 	return -EINVAL;
375 }
376 
377 void ksmbd_session_rpc_close(struct ksmbd_session *sess, int id)
378 {
379 	struct ksmbd_session_rpc *entry;
380 
381 	down_write(&sess->rpc_lock);
382 	entry = xa_erase(&sess->rpc_handle_list, id);
383 	if (entry)
384 		__session_rpc_close(sess, entry);
385 	up_write(&sess->rpc_lock);
386 }
387 
388 int ksmbd_session_rpc_method(struct ksmbd_session *sess, int id)
389 {
390 	struct ksmbd_session_rpc *entry;
391 
392 	lockdep_assert_held(&sess->rpc_lock);
393 	entry = xa_load(&sess->rpc_handle_list, id);
394 
395 	return entry ? entry->method : 0;
396 }
397 
398 void ksmbd_session_destroy(struct ksmbd_session *sess)
399 {
400 	if (!sess)
401 		return;
402 
403 	delete_proc_session(sess);
404 	ksmbd_tree_conn_session_logoff(sess);
405 	ksmbd_destroy_file_table(sess);
406 	if (sess->user)
407 		ksmbd_free_user(sess->user);
408 	ksmbd_launch_ksmbd_durable_scavenger();
409 	ksmbd_session_rpc_clear_list(sess);
410 	free_channel_list(sess);
411 	kfree_sensitive(sess->Preauth_HashValue);
412 	ksmbd_release_id(&session_ida, sess->id);
413 	ida_destroy(&sess->tree_conn_ida);
414 	kfree_sensitive(sess);
415 }
416 
417 static void ksmbd_session_remove_from_table(struct ksmbd_session *sess)
418 {
419 	hash_del(&sess->hlist);
420 	ksmbd_counter_dec(KSMBD_COUNTER_SESSIONS);
421 }
422 
423 struct ksmbd_session *__session_lookup(unsigned long long id)
424 {
425 	struct ksmbd_session *sess;
426 
427 	hash_for_each_possible(sessions_table, sess, hlist, id) {
428 		if (id == sess->id) {
429 			sess->last_active = jiffies;
430 			return sess;
431 		}
432 	}
433 	return NULL;
434 }
435 
436 static bool ksmbd_too_many_session_setups(struct ksmbd_conn *conn)
437 {
438 	unsigned long id;
439 	struct ksmbd_session *sess;
440 	unsigned int pending = 0;
441 
442 	down_write(&sessions_table_lock);
443 	down_write(&conn->session_lock);
444 	xa_for_each(&conn->sessions, id, sess) {
445 		if (READ_ONCE(sess->state) != SMB2_SESSION_IN_PROGRESS)
446 			continue;
447 
448 		if (atomic_read(&sess->refcnt) <= 1 &&
449 		    time_after(jiffies, sess->last_active +
450 			       KSMBD_UNAUTHENTICATED_CONN_TIMEOUT)) {
451 			xa_erase(&conn->sessions, sess->id);
452 			ksmbd_session_remove_from_table(sess);
453 			ksmbd_session_destroy(sess);
454 			continue;
455 		}
456 		pending++;
457 	}
458 	up_write(&conn->session_lock);
459 	up_write(&sessions_table_lock);
460 	return pending >= KSMBD_MAX_PENDING_SESSIONS;
461 }
462 
463 int ksmbd_session_register(struct ksmbd_conn *conn,
464 			   struct ksmbd_session *sess)
465 {
466 	int ret;
467 
468 	sess->dialect = conn->dialect;
469 	memcpy(sess->ClientGUID, conn->ClientGUID, SMB2_CLIENT_GUID_SIZE);
470 	/* Bound abandoned SessionId-zero authentication exchanges. */
471 	if (ksmbd_too_many_session_setups(conn))
472 		ret = -ENOSPC;
473 	else
474 		ret = xa_err(xa_store(&conn->sessions, sess->id, sess,
475 				      KSMBD_DEFAULT_GFP));
476 	if (ret) {
477 		down_write(&sessions_table_lock);
478 		ksmbd_session_remove_from_table(sess);
479 		up_write(&sessions_table_lock);
480 		ksmbd_user_session_put(sess);
481 	}
482 
483 	return ret;
484 }
485 
486 void ksmbd_session_unregister(struct ksmbd_conn *conn,
487 			      struct ksmbd_session *sess)
488 {
489 	struct ksmbd_conn *session_conns[KSMBD_MAX_CHANNELS];
490 	struct channel *chann;
491 	unsigned long index;
492 	unsigned int nr_conns = 0, i;
493 	bool removed = false;
494 
495 	down_write(&sessions_table_lock);
496 	if (!hlist_unhashed(&sess->hlist)) {
497 		/* Keep each channel connection stable under sessions_table_lock. */
498 		down_read(&sess->chann_lock);
499 		xa_for_each(&sess->ksmbd_chann_list, index, chann) {
500 			if (nr_conns == ARRAY_SIZE(session_conns))
501 				break;
502 			session_conns[nr_conns++] = chann->conn;
503 		}
504 		up_read(&sess->chann_lock);
505 
506 		ksmbd_session_remove_from_table(sess);
507 		removed = true;
508 	}
509 
510 	down_write(&conn->session_lock);
511 	if (xa_load(&conn->sessions, sess->id) == sess)
512 		xa_erase(&conn->sessions, sess->id);
513 	up_write(&conn->session_lock);
514 	for (i = 0; i < nr_conns; i++) {
515 		if (session_conns[i] == conn)
516 			continue;
517 		down_write(&session_conns[i]->session_lock);
518 		if (xa_load(&session_conns[i]->sessions, sess->id) == sess)
519 			xa_erase(&session_conns[i]->sessions, sess->id);
520 		up_write(&session_conns[i]->session_lock);
521 	}
522 	up_write(&sessions_table_lock);
523 
524 	if (removed)
525 		ksmbd_user_session_put(sess);
526 }
527 
528 bool ksmbd_conn_has_valid_or_expired_session(struct ksmbd_conn *conn)
529 {
530 	struct ksmbd_session *sess;
531 	unsigned long id;
532 	int state, bkt;
533 	bool found = false;
534 
535 	down_read(&conn->session_lock);
536 	xa_for_each(&conn->sessions, id, sess) {
537 		state = READ_ONCE(sess->state);
538 		if (state == SMB2_SESSION_VALID ||
539 		    state == SMB2_SESSION_EXPIRED) {
540 			found = true;
541 			break;
542 		}
543 	}
544 	up_read(&conn->session_lock);
545 	if (found)
546 		return true;
547 
548 	/* A session bound through SMB3 multichannel is not in conn->sessions. */
549 	down_read(&sessions_table_lock);
550 	hash_for_each(sessions_table, bkt, sess, hlist) {
551 		state = READ_ONCE(sess->state);
552 		if (state != SMB2_SESSION_VALID &&
553 		    state != SMB2_SESSION_EXPIRED)
554 			continue;
555 
556 		down_read(&sess->chann_lock);
557 		found = xa_load(&sess->ksmbd_chann_list, (long)conn);
558 		up_read(&sess->chann_lock);
559 		if (found)
560 			break;
561 	}
562 	up_read(&sessions_table_lock);
563 	return found;
564 }
565 
566 void ksmbd_expire_sessions(void)
567 {
568 	struct ksmbd_session *sess;
569 	u64 now = ktime_get_real_seconds();
570 	int bkt;
571 
572 	down_read(&sessions_table_lock);
573 	hash_for_each(sessions_table, bkt, sess, hlist) {
574 		if (READ_ONCE(sess->state) != SMB2_SESSION_VALID ||
575 		    !sess->kerberos_expiry || now < sess->kerberos_expiry)
576 			continue;
577 
578 		if (cmpxchg(&sess->state, SMB2_SESSION_VALID,
579 			    SMB2_SESSION_EXPIRED) == SMB2_SESSION_VALID)
580 			ksmbd_counter_inc(KSMBD_COUNTER_SESSION_TIMEOUTS);
581 	}
582 	up_read(&sessions_table_lock);
583 }
584 
585 static int ksmbd_chann_del(struct ksmbd_conn *conn, struct ksmbd_session *sess)
586 {
587 	struct channel *chann;
588 
589 	down_write(&sess->chann_lock);
590 	chann = xa_erase(&sess->ksmbd_chann_list, (long)conn);
591 	up_write(&sess->chann_lock);
592 	if (!chann)
593 		return -ENOENT;
594 
595 	kfree_sensitive(chann);
596 	return 0;
597 }
598 
599 void ksmbd_conn_sessions_cleanup(struct ksmbd_conn *conn)
600 {
601 	struct ksmbd_session *sess;
602 	unsigned long id;
603 	struct hlist_node *tmp;
604 	int bkt;
605 
606 	down_write(&sessions_table_lock);
607 	hash_for_each_safe(sessions_table, bkt, tmp, sess, hlist) {
608 		if (!ksmbd_chann_del(conn, sess) &&
609 		    xa_empty(&sess->ksmbd_chann_list)) {
610 			ksmbd_session_remove_from_table(sess);
611 			down_write(&conn->session_lock);
612 			xa_erase(&conn->sessions, sess->id);
613 			up_write(&conn->session_lock);
614 			if (atomic_dec_and_test(&sess->refcnt))
615 				ksmbd_session_destroy(sess);
616 		}
617 	}
618 
619 	down_write(&conn->session_lock);
620 	xa_for_each(&conn->sessions, id, sess) {
621 		ksmbd_chann_del(conn, sess);
622 		if (xa_empty(&sess->ksmbd_chann_list)) {
623 			xa_erase(&conn->sessions, sess->id);
624 			ksmbd_session_remove_from_table(sess);
625 			if (atomic_dec_and_test(&sess->refcnt))
626 				ksmbd_session_destroy(sess);
627 		}
628 	}
629 	up_write(&conn->session_lock);
630 	up_write(&sessions_table_lock);
631 }
632 
633 bool is_ksmbd_session_in_connection(struct ksmbd_conn *conn,
634 				   unsigned long long id)
635 {
636 	struct ksmbd_session *sess;
637 
638 	down_read(&conn->session_lock);
639 	sess = xa_load(&conn->sessions, id);
640 	if (sess) {
641 		up_read(&conn->session_lock);
642 		return true;
643 	}
644 	up_read(&conn->session_lock);
645 
646 	return false;
647 }
648 
649 struct ksmbd_session *ksmbd_session_lookup(struct ksmbd_conn *conn,
650 					   unsigned long long id)
651 {
652 	struct ksmbd_session *sess;
653 
654 	down_read(&conn->session_lock);
655 	sess = xa_load(&conn->sessions, id);
656 	if (sess) {
657 		sess->last_active = jiffies;
658 		ksmbd_user_session_get(sess);
659 	}
660 	up_read(&conn->session_lock);
661 	return sess;
662 }
663 
664 struct ksmbd_session *ksmbd_session_lookup_slowpath(unsigned long long id)
665 {
666 	struct ksmbd_session *sess;
667 
668 	down_read(&sessions_table_lock);
669 	sess = __session_lookup(id);
670 	if (sess)
671 		ksmbd_user_session_get(sess);
672 	up_read(&sessions_table_lock);
673 
674 	return sess;
675 }
676 
677 struct ksmbd_session *ksmbd_session_lookup_all_states(struct ksmbd_conn *conn,
678 						      unsigned long long id)
679 {
680 	struct ksmbd_session *sess;
681 	bool channel_found;
682 
683 	sess = ksmbd_session_lookup(conn, id);
684 	if (!sess) {
685 		sess = ksmbd_session_lookup_slowpath(id);
686 		if (!sess)
687 			return NULL;
688 
689 		down_read(&sess->chann_lock);
690 		channel_found = xa_load(&sess->ksmbd_chann_list, (long)conn);
691 		up_read(&sess->chann_lock);
692 		if (!channel_found) {
693 			ksmbd_user_session_put(sess);
694 			sess = NULL;
695 		}
696 	}
697 	return sess;
698 }
699 
700 struct ksmbd_session *ksmbd_session_lookup_all(struct ksmbd_conn *conn,
701 					       unsigned long long id)
702 {
703 	struct ksmbd_session *sess;
704 
705 	sess = ksmbd_session_lookup_all_states(conn, id);
706 	if (sess && sess->state != SMB2_SESSION_VALID) {
707 		ksmbd_user_session_put(sess);
708 		sess = NULL;
709 	}
710 	return sess;
711 }
712 
713 void ksmbd_user_session_get(struct ksmbd_session *sess)
714 {
715 	atomic_inc(&sess->refcnt);
716 }
717 
718 void ksmbd_user_session_put(struct ksmbd_session *sess)
719 {
720 	if (!sess)
721 		return;
722 
723 	if (atomic_read(&sess->refcnt) <= 0)
724 		WARN_ON(1);
725 	else if (atomic_dec_and_test(&sess->refcnt))
726 		ksmbd_session_destroy(sess);
727 }
728 
729 struct preauth_session *ksmbd_preauth_session_alloc(struct ksmbd_conn *conn,
730 						    u64 sess_id)
731 {
732 	struct preauth_session *sess;
733 
734 	sess = kmalloc_obj(struct preauth_session, KSMBD_DEFAULT_GFP);
735 	if (!sess)
736 		return NULL;
737 
738 	sess->id = sess_id;
739 	memcpy(sess->Preauth_HashValue, conn->preauth_info->Preauth_HashValue,
740 	       PREAUTH_HASHVALUE_SIZE);
741 	list_add(&sess->preauth_entry, &conn->preauth_sess_table);
742 
743 	return sess;
744 }
745 
746 void ksmbd_preauth_session_destroy(struct ksmbd_conn *conn)
747 {
748 	struct preauth_session *sess, *tmp;
749 
750 	list_for_each_entry_safe(sess, tmp, &conn->preauth_sess_table,
751 				 preauth_entry) {
752 		list_del(&sess->preauth_entry);
753 		kfree(sess);
754 	}
755 }
756 
757 void destroy_previous_session(struct ksmbd_conn *conn,
758 			      struct ksmbd_user *user, u64 id)
759 {
760 	struct ksmbd_session *prev_sess;
761 	struct ksmbd_user *prev_user;
762 	int err;
763 
764 	down_write(&sessions_table_lock);
765 	down_write(&conn->session_lock);
766 	prev_sess = __session_lookup(id);
767 	if (!prev_sess || prev_sess->state == SMB2_SESSION_EXPIRED)
768 		goto out;
769 
770 	prev_user = prev_sess->user;
771 	if (!prev_user ||
772 	    strcmp(user->name, prev_user->name) ||
773 	    user->passkey_sz != prev_user->passkey_sz ||
774 	    memcmp(user->passkey, prev_user->passkey, user->passkey_sz))
775 		goto out;
776 
777 	down_write(&prev_sess->chann_lock);
778 	if (prev_sess->tearing_down) {
779 		up_write(&prev_sess->chann_lock);
780 		goto out;
781 	}
782 	prev_sess->tearing_down = true;
783 	up_write(&prev_sess->chann_lock);
784 
785 	ksmbd_all_conn_set_status(prev_sess, KSMBD_SESS_NEED_RECONNECT);
786 	err = ksmbd_conn_wait_idle_sess(conn, prev_sess);
787 	if (err) {
788 		down_write(&prev_sess->chann_lock);
789 		prev_sess->tearing_down = false;
790 		up_write(&prev_sess->chann_lock);
791 		ksmbd_all_conn_set_status(prev_sess, KSMBD_SESS_GOOD);
792 		goto out;
793 	}
794 
795 	ksmbd_destroy_file_table(prev_sess);
796 	prev_sess->kerberos_expiry = 0;
797 	prev_sess->state = SMB2_SESSION_EXPIRED;
798 	ksmbd_all_conn_set_status(prev_sess, KSMBD_SESS_NEED_SETUP);
799 	ksmbd_launch_ksmbd_durable_scavenger();
800 out:
801 	up_write(&conn->session_lock);
802 	up_write(&sessions_table_lock);
803 }
804 
805 static bool ksmbd_preauth_session_id_match(struct preauth_session *sess,
806 					   unsigned long long id)
807 {
808 	return sess->id == id;
809 }
810 
811 struct preauth_session *ksmbd_preauth_session_lookup(struct ksmbd_conn *conn,
812 						     unsigned long long id)
813 {
814 	struct preauth_session *sess = NULL;
815 
816 	list_for_each_entry(sess, &conn->preauth_sess_table, preauth_entry) {
817 		if (ksmbd_preauth_session_id_match(sess, id))
818 			return sess;
819 	}
820 	return NULL;
821 }
822 
823 static int __init_smb2_session(struct ksmbd_session *sess)
824 {
825 	int id = ksmbd_acquire_smb2_uid(&session_ida);
826 
827 	if (id < 0)
828 		return -EINVAL;
829 	sess->id = id;
830 	return 0;
831 }
832 
833 static struct ksmbd_session *__session_create(int protocol)
834 {
835 	struct ksmbd_session *sess;
836 	int ret;
837 
838 	if (protocol != CIFDS_SESSION_FLAG_SMB2)
839 		return NULL;
840 
841 	sess = kzalloc_obj(struct ksmbd_session, KSMBD_DEFAULT_GFP);
842 	if (!sess)
843 		return NULL;
844 
845 	ida_init(&sess->tree_conn_ida);
846 
847 	if (ksmbd_init_file_table(&sess->file_table))
848 		goto error;
849 
850 	sess->last_active = jiffies;
851 	sess->state = SMB2_SESSION_IN_PROGRESS;
852 	set_session_flag(sess, protocol);
853 	xa_init(&sess->tree_conns);
854 	xa_init(&sess->ksmbd_chann_list);
855 	xa_init(&sess->rpc_handle_list);
856 	sess->sequence_number = 1;
857 	atomic_set(&sess->refcnt, 2);
858 	init_rwsem(&sess->tree_conns_lock);
859 	init_rwsem(&sess->rpc_lock);
860 	init_rwsem(&sess->chann_lock);
861 
862 	ret = __init_smb2_session(sess);
863 	if (ret)
864 		goto error;
865 
866 	down_write(&sessions_table_lock);
867 	hash_add(sessions_table, &sess->hlist, sess->id);
868 	ksmbd_counter_inc(KSMBD_COUNTER_SESSIONS);
869 	up_write(&sessions_table_lock);
870 
871 	if (create_proc_session(sess))
872 		pr_warn_ratelimited("Unable to create session %llu procfs entry\n", sess->id);
873 	return sess;
874 
875 error:
876 	ksmbd_session_destroy(sess);
877 	return NULL;
878 }
879 
880 struct ksmbd_session *ksmbd_smb2_session_create(void)
881 {
882 	return __session_create(CIFDS_SESSION_FLAG_SMB2);
883 }
884 
885 int ksmbd_acquire_tree_conn_id(struct ksmbd_session *sess)
886 {
887 	int id = -EINVAL;
888 
889 	if (test_session_flag(sess, CIFDS_SESSION_FLAG_SMB2))
890 		id = ksmbd_acquire_smb2_tid(&sess->tree_conn_ida);
891 
892 	return id;
893 }
894 
895 void ksmbd_release_tree_conn_id(struct ksmbd_session *sess, int id)
896 {
897 	if (id >= 0)
898 		ksmbd_release_id(&sess->tree_conn_ida, id);
899 }
900