xref: /linux/tools/testing/selftests/liveupdate/lib/lu_utils.c (revision bcf2a0b55e322a663c9602c9e773770b87d3e629)
1 // SPDX-License-Identifier: GPL-2.0-only
2 
3 /*
4  * Copyright (c) 2025, Google LLC.
5  * Pasha Tatashin <pasha.tatashin@soleen.com>
6  */
7 
8 #define _GNU_SOURCE
9 
10 #include <stdio.h>
11 #include <stdlib.h>
12 #include <string.h>
13 #include <getopt.h>
14 #include <fcntl.h>
15 #include <unistd.h>
16 #include <sys/ioctl.h>
17 #include <sys/syscall.h>
18 #include <sys/mman.h>
19 #include <sys/types.h>
20 #include <sys/resource.h>
21 #include <sys/stat.h>
22 #include <errno.h>
23 #include <stdarg.h>
24 #include <linux/unistd.h>
25 
26 #include <libliveupdate.h>
27 
28 int luo_open_device(void)
29 {
30 	return open(LUO_DEVICE, O_RDWR);
31 }
32 
33 int luo_ensure_nofile_limit(long min_limit)
34 {
35 	struct rlimit hl;
36 
37 	/* Allow to extra files to be used by test itself */
38 	min_limit += 32;
39 
40 	if (getrlimit(RLIMIT_NOFILE, &hl) < 0)
41 		return -errno;
42 
43 	if (hl.rlim_cur >= min_limit)
44 		return 0;
45 
46 	hl.rlim_cur = min_limit;
47 	if (hl.rlim_cur > hl.rlim_max)
48 		hl.rlim_max = hl.rlim_cur;
49 
50 	if (setrlimit(RLIMIT_NOFILE, &hl) < 0)
51 		return -errno;
52 
53 	return 0;
54 }
55 
56 int luo_create_session(int luo_fd, const char *name)
57 {
58 	struct liveupdate_ioctl_create_session arg = { .size = sizeof(arg) };
59 
60 	snprintf((char *)arg.name, LIVEUPDATE_SESSION_NAME_LENGTH, "%.*s",
61 		 LIVEUPDATE_SESSION_NAME_LENGTH - 1, name);
62 
63 	if (ioctl(luo_fd, LIVEUPDATE_IOCTL_CREATE_SESSION, &arg))
64 		return -errno;
65 
66 	return arg.fd;
67 }
68 
69 int luo_retrieve_session(int luo_fd, const char *name)
70 {
71 	struct liveupdate_ioctl_retrieve_session arg = { .size = sizeof(arg) };
72 
73 	snprintf((char *)arg.name, LIVEUPDATE_SESSION_NAME_LENGTH, "%.*s",
74 		 LIVEUPDATE_SESSION_NAME_LENGTH - 1, name);
75 
76 	if (ioctl(luo_fd, LIVEUPDATE_IOCTL_RETRIEVE_SESSION, &arg))
77 		return -errno;
78 
79 	return arg.fd;
80 }
81 
82 int luo_session_preserve_fd(int session_fd, int fd, __u64 token)
83 {
84 	struct liveupdate_session_preserve_fd arg = {
85 		.size = sizeof(arg),
86 		.fd = fd,
87 		.token = token,
88 	};
89 
90 	if (ioctl(session_fd, LIVEUPDATE_SESSION_PRESERVE_FD, &arg))
91 		return -errno;
92 
93 	return 0;
94 }
95 
96 int luo_session_retrieve_fd(int session_fd, __u64 token)
97 {
98 	struct liveupdate_session_retrieve_fd arg = {
99 		.size = sizeof(arg),
100 		.token = token,
101 	};
102 
103 	if (ioctl(session_fd, LIVEUPDATE_SESSION_RETRIEVE_FD, &arg))
104 		return -errno;
105 
106 	return arg.fd;
107 }
108 
109 /* Helper function to get a session name via ioctl. */
110 int luo_get_session_name(int session_fd, char *name, size_t name_len)
111 {
112 	struct liveupdate_session_get_name args = {};
113 
114 	args.size = sizeof(args);
115 
116 	if (ioctl(session_fd, LIVEUPDATE_SESSION_GET_NAME, &args))
117 		return -errno;
118 
119 	strncpy(name, (char *)args.name, name_len - 1);
120 	name[name_len - 1] = '\0';
121 
122 	return 0;
123 }
124 
125 int create_and_preserve_memfd(int session_fd, int token, const char *data)
126 {
127 	long page_size = getpagesize();
128 	void *map = MAP_FAILED;
129 	int mfd = -1, ret = -1;
130 
131 	mfd = memfd_create("test_mfd", 0);
132 	if (mfd < 0)
133 		return -errno;
134 
135 	if (ftruncate(mfd, page_size) != 0)
136 		goto out;
137 
138 	map = mmap(NULL, page_size, PROT_WRITE, MAP_SHARED, mfd, 0);
139 	if (map == MAP_FAILED)
140 		goto out;
141 
142 	snprintf(map, page_size, "%s", data);
143 	munmap(map, page_size);
144 
145 	ret = luo_session_preserve_fd(session_fd, mfd, token);
146 	if (ret)
147 		goto out;
148 
149 	ret = 0;
150 out:
151 	if (ret != 0 && errno != 0)
152 		ret = -errno;
153 	if (mfd >= 0)
154 		close(mfd);
155 	return ret;
156 }
157 
158 int restore_and_verify_memfd(int session_fd, int token,
159 			     const char *expected_data)
160 {
161 	long page_size = getpagesize();
162 	void *map = MAP_FAILED;
163 	int mfd = -1, ret = -1;
164 
165 	mfd = luo_session_retrieve_fd(session_fd, token);
166 	if (mfd < 0)
167 		return mfd;
168 
169 	map = mmap(NULL, page_size, PROT_READ, MAP_SHARED, mfd, 0);
170 	if (map == MAP_FAILED)
171 		goto out;
172 
173 	if (expected_data && strcmp(expected_data, map) != 0) {
174 		ksft_print_msg("Data mismatch! Expected '%s', Got '%s'\n",
175 			       expected_data, (char *)map);
176 		ret = -EINVAL;
177 		goto out_munmap;
178 	}
179 
180 	ret = mfd;
181 out_munmap:
182 	munmap(map, page_size);
183 out:
184 	if (ret < 0 && errno != 0)
185 		ret = -errno;
186 	if (ret < 0 && mfd >= 0)
187 		close(mfd);
188 	return ret;
189 }
190 
191 int luo_session_finish(int session_fd)
192 {
193 	struct liveupdate_session_finish arg = { .size = sizeof(arg) };
194 
195 	if (ioctl(session_fd, LIVEUPDATE_SESSION_FINISH, &arg) < 0)
196 		return -errno;
197 
198 	return 0;
199 }
200 
201 void create_state_file(int luo_fd, const char *session_name, int token,
202 		       int next_stage)
203 {
204 	char buf[32];
205 	int state_session_fd;
206 
207 	state_session_fd = luo_create_session(luo_fd, session_name);
208 	if (state_session_fd < 0)
209 		fail_exit("luo_create_session for state tracking");
210 
211 	snprintf(buf, sizeof(buf), "%d", next_stage);
212 	if (create_and_preserve_memfd(state_session_fd, token, buf) < 0)
213 		fail_exit("create_and_preserve_memfd for state tracking");
214 
215 	/*
216 	 * DO NOT close session FD, otherwise it is going to be unpreserved
217 	 */
218 }
219 
220 void restore_and_read_stage(int state_session_fd, int token, int *stage)
221 {
222 	char buf[32] = {0};
223 	int mfd;
224 
225 	mfd = restore_and_verify_memfd(state_session_fd, token, NULL);
226 	if (mfd < 0)
227 		fail_exit("failed to restore state memfd");
228 
229 	if (read(mfd, buf, sizeof(buf) - 1) < 0)
230 		fail_exit("failed to read state mfd");
231 
232 	*stage = atoi(buf);
233 
234 	close(mfd);
235 }
236 
237 void daemonize_and_wait(void)
238 {
239 	pid_t pid;
240 
241 	ksft_print_msg("[STAGE 1] Forking persistent child to hold sessions...\n");
242 
243 	pid = fork();
244 	if (pid < 0)
245 		fail_exit("fork failed");
246 
247 	if (pid > 0) {
248 		ksft_print_msg("[STAGE 1] Child PID: %d. Resources are pinned.\n", pid);
249 		ksft_print_msg("[STAGE 1] You may now perform kexec reboot.\n");
250 		exit(EXIT_SUCCESS);
251 	}
252 
253 	/* Detach from terminal so closing the window doesn't kill us */
254 	if (setsid() < 0)
255 		fail_exit("setsid failed");
256 
257 	close(STDIN_FILENO);
258 	close(STDOUT_FILENO);
259 	close(STDERR_FILENO);
260 
261 	/* Change dir to root to avoid locking filesystems */
262 	if (chdir("/") < 0)
263 		exit(EXIT_FAILURE);
264 
265 	while (1)
266 		sleep(60);
267 }
268 
269 static int parse_stage_args(int argc, char *argv[])
270 {
271 	int stage = 1;
272 	int opt;
273 
274 	optind = 1;
275 	while ((opt = getopt(argc, argv, "s:")) != -1) {
276 		switch (opt) {
277 		case 's':
278 			stage = atoi(optarg);
279 			if (stage != 1 && stage != 2)
280 				fail_exit("Invalid stage argument");
281 			break;
282 		default:
283 			fail_exit("Unknown argument");
284 		}
285 	}
286 
287 	return stage;
288 }
289 
290 int luo_test(int argc, char *argv[],
291 	     const char *state_session_name,
292 	     luo_test_stage1_fn stage1,
293 	     luo_test_stage2_fn stage2)
294 {
295 	int target_stage = parse_stage_args(argc, argv);
296 	int luo_fd = luo_open_device();
297 	int state_session_fd;
298 	int detected_stage;
299 
300 	if (luo_fd < 0) {
301 		ksft_exit_skip("Failed to open %s. Is the luo module loaded?\n",
302 			       LUO_DEVICE);
303 	}
304 
305 	state_session_fd = luo_retrieve_session(luo_fd, state_session_name);
306 	if (state_session_fd == -ENOENT)
307 		detected_stage = 1;
308 	else if (state_session_fd >= 0)
309 		detected_stage = 2;
310 	else
311 		fail_exit("Failed to check for state session");
312 
313 	if (target_stage != detected_stage) {
314 		ksft_exit_fail_msg("Stage mismatch Requested stage %d, but system is in stage %d.\n"
315 				   "(State session %s: %s)\n",
316 				   target_stage, detected_stage, state_session_name,
317 				   (detected_stage == 2) ? "EXISTS" : "MISSING");
318 	}
319 
320 	if (target_stage == 1)
321 		stage1(luo_fd);
322 	else
323 		stage2(luo_fd, state_session_fd);
324 
325 	return 0;
326 }
327