xref: /linux/drivers/hv/mshv_regions.c (revision 3a2c4d55e32ad65efebdb6de44eef3bfa08bb49d)
1 // SPDX-License-Identifier: GPL-2.0-only
2 /*
3  * Copyright (c) 2025, Microsoft Corporation.
4  *
5  * Memory region management for mshv_root module.
6  *
7  * Authors: Microsoft Linux virtualization team
8  */
9 
10 #include <linux/hmm.h>
11 #include <linux/hyperv.h>
12 #include <linux/kref.h>
13 #include <linux/mm.h>
14 #include <linux/vmalloc.h>
15 
16 #include <asm/mshyperv.h>
17 
18 #include "mshv_root.h"
19 
20 #define MSHV_MAP_FAULT_IN_PAGES				PTRS_PER_PMD
21 
22 /**
23  * mshv_chunk_stride - Compute stride for mapping guest memory
24  * @page      : The page to check for huge page backing
25  * @gfn       : Guest frame number for the mapping
26  * @page_count: Total number of pages in the mapping
27  *
28  * Determines the appropriate stride (in pages) for mapping guest memory.
29  * Uses huge page stride if the backing page is huge and the guest mapping
30  * is properly aligned; otherwise falls back to single page stride.
31  *
32  * Return: Stride in pages.
33  */
34 static unsigned int mshv_chunk_stride(struct page *page, u64 gfn,
35 				      u64 page_count)
36 {
37 	unsigned int page_order = folio_order(page_folio(page));
38 
39 	/*
40 	 * Use single page stride by default. For huge page stride, the
41 	 * folio order must be at least PMD_ORDER, the page's PFN must be
42 	 * 2M-aligned (so that a 2M-aligned tail page of a larger folio is
43 	 * acceptable), and both gfn and page_count must be 2M-aligned.
44 	 */
45 	if (page_order < PMD_ORDER ||
46 	    !IS_ALIGNED(page_to_pfn(page), PTRS_PER_PMD) ||
47 	    !IS_ALIGNED(gfn, PTRS_PER_PMD) ||
48 	    !IS_ALIGNED(page_count, PTRS_PER_PMD))
49 		return 1;
50 
51 	/* Use 2M stride always i.e. process 1G folios as 2M chunks */
52 	return 1 << PMD_ORDER;
53 }
54 
55 /**
56  * mshv_region_process_chunk - Processes a contiguous chunk of memory pages
57  *                             in a region.
58  * @region     : Pointer to the memory region structure.
59  * @flags      : Flags to pass to the handler.
60  * @page_offset: Offset into the region's pages array to start processing.
61  * @page_count : Number of pages to process.
62  * @handler    : Callback function to handle the chunk.
63  *
64  * This function scans the region's pages starting from @page_offset,
65  * checking for contiguous present pages of the same size (normal or huge).
66  * It invokes @handler for the chunk of contiguous pages found. Returns the
67  * number of pages handled, or a negative error code if the first page is
68  * not present or the handler fails.
69  *
70  * Note: The @handler callback must be able to handle both normal and huge
71  * pages.
72  *
73  * Return: Number of pages handled, or negative error code.
74  */
75 static long mshv_region_process_chunk(struct mshv_mem_region *region,
76 				      u32 flags,
77 				      u64 page_offset, u64 page_count,
78 				      int (*handler)(struct mshv_mem_region *region,
79 						     u32 flags,
80 						     u64 page_offset,
81 						     u64 page_count,
82 						     bool huge_page))
83 {
84 	u64 gfn = region->start_gfn + page_offset;
85 	u64 count;
86 	struct page *page;
87 	unsigned int stride;
88 	int ret;
89 
90 	page = region->mreg_pages[page_offset];
91 	if (!page)
92 		return -EINVAL;
93 
94 	stride = mshv_chunk_stride(page, gfn, page_count);
95 
96 	/* Start at stride since the first stride is validated */
97 	for (count = stride; count < page_count; count += stride) {
98 		page = region->mreg_pages[page_offset + count];
99 
100 		/* Break if current page is not present */
101 		if (!page)
102 			break;
103 
104 		/* Break if stride size changes */
105 		if (stride != mshv_chunk_stride(page, gfn + count,
106 						page_count - count))
107 			break;
108 	}
109 
110 	ret = handler(region, flags, page_offset, count, stride > 1);
111 	if (ret)
112 		return ret;
113 
114 	return count;
115 }
116 
117 /**
118  * mshv_region_process_range - Processes a range of memory pages in a
119  *                             region.
120  * @region     : Pointer to the memory region structure.
121  * @flags      : Flags to pass to the handler.
122  * @page_offset: Offset into the region's pages array to start processing.
123  * @page_count : Number of pages to process.
124  * @handler    : Callback function to handle each chunk of contiguous
125  *               pages.
126  *
127  * Iterates over the specified range of pages in @region, skipping
128  * non-present pages. For each contiguous chunk of present pages, invokes
129  * @handler via mshv_region_process_chunk.
130  *
131  * Note: The @handler callback must be able to handle both normal and huge
132  * pages.
133  *
134  * Returns 0 on success, or a negative error code on failure.
135  */
136 static int mshv_region_process_range(struct mshv_mem_region *region,
137 				     u32 flags,
138 				     u64 page_offset, u64 page_count,
139 				     int (*handler)(struct mshv_mem_region *region,
140 						    u32 flags,
141 						    u64 page_offset,
142 						    u64 page_count,
143 						    bool huge_page))
144 {
145 	long ret;
146 
147 	if (page_offset + page_count > region->nr_pages)
148 		return -EINVAL;
149 
150 	while (page_count) {
151 		/* Skip non-present pages */
152 		if (!region->mreg_pages[page_offset]) {
153 			page_offset++;
154 			page_count--;
155 			continue;
156 		}
157 
158 		ret = mshv_region_process_chunk(region, flags,
159 						page_offset,
160 						page_count,
161 						handler);
162 		if (ret < 0)
163 			return ret;
164 
165 		page_offset += ret;
166 		page_count -= ret;
167 	}
168 
169 	return 0;
170 }
171 
172 struct mshv_mem_region *mshv_region_create(u64 guest_pfn, u64 nr_pages,
173 					   u64 uaddr, u32 flags)
174 {
175 	struct mshv_mem_region *region;
176 
177 	region = vzalloc(sizeof(*region) + sizeof(struct page *) * nr_pages);
178 	if (!region)
179 		return ERR_PTR(-ENOMEM);
180 
181 	region->nr_pages = nr_pages;
182 	region->start_gfn = guest_pfn;
183 	region->start_uaddr = uaddr;
184 	region->hv_map_flags = HV_MAP_GPA_READABLE | HV_MAP_GPA_ADJUSTABLE;
185 	if (flags & BIT(MSHV_SET_MEM_BIT_WRITABLE))
186 		region->hv_map_flags |= HV_MAP_GPA_WRITABLE;
187 	if (flags & BIT(MSHV_SET_MEM_BIT_EXECUTABLE))
188 		region->hv_map_flags |= HV_MAP_GPA_EXECUTABLE;
189 
190 	kref_init(&region->mreg_refcount);
191 
192 	return region;
193 }
194 
195 static int mshv_region_chunk_share(struct mshv_mem_region *region,
196 				   u32 flags,
197 				   u64 page_offset, u64 page_count,
198 				   bool huge_page)
199 {
200 	if (huge_page)
201 		flags |= HV_MODIFY_SPA_PAGE_HOST_ACCESS_LARGE_PAGE;
202 
203 	return hv_call_modify_spa_host_access(region->partition->pt_id,
204 					      region->mreg_pages + page_offset,
205 					      page_count,
206 					      HV_MAP_GPA_READABLE |
207 					      HV_MAP_GPA_WRITABLE,
208 					      flags, true);
209 }
210 
211 int mshv_region_share(struct mshv_mem_region *region)
212 {
213 	u32 flags = HV_MODIFY_SPA_PAGE_HOST_ACCESS_MAKE_SHARED;
214 
215 	return mshv_region_process_range(region, flags,
216 					 0, region->nr_pages,
217 					 mshv_region_chunk_share);
218 }
219 
220 static int mshv_region_chunk_unshare(struct mshv_mem_region *region,
221 				     u32 flags,
222 				     u64 page_offset, u64 page_count,
223 				     bool huge_page)
224 {
225 	if (huge_page)
226 		flags |= HV_MODIFY_SPA_PAGE_HOST_ACCESS_LARGE_PAGE;
227 
228 	return hv_call_modify_spa_host_access(region->partition->pt_id,
229 					      region->mreg_pages + page_offset,
230 					      page_count, 0,
231 					      flags, false);
232 }
233 
234 int mshv_region_unshare(struct mshv_mem_region *region)
235 {
236 	u32 flags = HV_MODIFY_SPA_PAGE_HOST_ACCESS_MAKE_EXCLUSIVE;
237 
238 	return mshv_region_process_range(region, flags,
239 					 0, region->nr_pages,
240 					 mshv_region_chunk_unshare);
241 }
242 
243 static int mshv_region_chunk_remap(struct mshv_mem_region *region,
244 				   u32 flags,
245 				   u64 page_offset, u64 page_count,
246 				   bool huge_page)
247 {
248 	if (huge_page)
249 		flags |= HV_MAP_GPA_LARGE_PAGE;
250 
251 	return hv_call_map_gpa_pages(region->partition->pt_id,
252 				     region->start_gfn + page_offset,
253 				     page_count, flags,
254 				     region->mreg_pages + page_offset);
255 }
256 
257 static int mshv_region_remap_pages(struct mshv_mem_region *region,
258 				   u32 map_flags,
259 				   u64 page_offset, u64 page_count)
260 {
261 	return mshv_region_process_range(region, map_flags,
262 					 page_offset, page_count,
263 					 mshv_region_chunk_remap);
264 }
265 
266 int mshv_region_map(struct mshv_mem_region *region)
267 {
268 	u32 map_flags = region->hv_map_flags;
269 
270 	return mshv_region_remap_pages(region, map_flags,
271 				       0, region->nr_pages);
272 }
273 
274 static void mshv_region_invalidate_pages(struct mshv_mem_region *region,
275 					 u64 page_offset, u64 page_count)
276 {
277 	if (region->mreg_type == MSHV_REGION_TYPE_MEM_PINNED)
278 		unpin_user_pages(region->mreg_pages + page_offset, page_count);
279 
280 	memset(region->mreg_pages + page_offset, 0,
281 	       page_count * sizeof(struct page *));
282 }
283 
284 void mshv_region_invalidate(struct mshv_mem_region *region)
285 {
286 	mshv_region_invalidate_pages(region, 0, region->nr_pages);
287 }
288 
289 int mshv_region_pin(struct mshv_mem_region *region)
290 {
291 	u64 done_count, nr_pages;
292 	struct page **pages;
293 	__u64 userspace_addr;
294 	int ret;
295 
296 	for (done_count = 0; done_count < region->nr_pages; done_count += ret) {
297 		pages = region->mreg_pages + done_count;
298 		userspace_addr = region->start_uaddr +
299 				 done_count * HV_HYP_PAGE_SIZE;
300 		nr_pages = min(region->nr_pages - done_count,
301 			       MSHV_PIN_PAGES_BATCH_SIZE);
302 
303 		/*
304 		 * Pinning assuming 4k pages works for large pages too.
305 		 * All page structs within the large page are returned.
306 		 *
307 		 * Pin requests are batched because pin_user_pages_fast
308 		 * with the FOLL_LONGTERM flag does a large temporary
309 		 * allocation of contiguous memory.
310 		 */
311 		ret = pin_user_pages_fast(userspace_addr, nr_pages,
312 					  FOLL_WRITE | FOLL_LONGTERM,
313 					  pages);
314 		if (ret != nr_pages)
315 			goto release_pages;
316 	}
317 
318 	return 0;
319 
320 release_pages:
321 	if (ret > 0)
322 		done_count += ret;
323 	mshv_region_invalidate_pages(region, 0, done_count);
324 	return ret < 0 ? ret : -ENOMEM;
325 }
326 
327 static int mshv_region_chunk_unmap(struct mshv_mem_region *region,
328 				   u32 flags,
329 				   u64 page_offset, u64 page_count,
330 				   bool huge_page)
331 {
332 	if (huge_page)
333 		flags |= HV_UNMAP_GPA_LARGE_PAGE;
334 
335 	return hv_call_unmap_gpa_pages(region->partition->pt_id,
336 				       region->start_gfn + page_offset,
337 				       page_count, flags);
338 }
339 
340 static int mshv_region_unmap(struct mshv_mem_region *region)
341 {
342 	return mshv_region_process_range(region, 0,
343 					 0, region->nr_pages,
344 					 mshv_region_chunk_unmap);
345 }
346 
347 static void mshv_region_destroy(struct kref *ref)
348 {
349 	struct mshv_mem_region *region =
350 		container_of(ref, struct mshv_mem_region, mreg_refcount);
351 	struct mshv_partition *partition = region->partition;
352 	int ret;
353 
354 	if (region->mreg_type == MSHV_REGION_TYPE_MEM_MOVABLE)
355 		mshv_region_movable_fini(region);
356 
357 	if (mshv_partition_encrypted(partition)) {
358 		ret = mshv_region_share(region);
359 		if (ret) {
360 			pt_err(partition,
361 			       "Failed to regain access to memory, unpinning user pages will fail and crash the host error: %d\n",
362 			       ret);
363 			return;
364 		}
365 	}
366 
367 	mshv_region_unmap(region);
368 
369 	mshv_region_invalidate(region);
370 
371 	vfree(region);
372 }
373 
374 void mshv_region_put(struct mshv_mem_region *region)
375 {
376 	kref_put(&region->mreg_refcount, mshv_region_destroy);
377 }
378 
379 int mshv_region_get(struct mshv_mem_region *region)
380 {
381 	return kref_get_unless_zero(&region->mreg_refcount);
382 }
383 
384 /**
385  * mshv_region_range_fault - Handle memory range faults for a given region.
386  * @region: Pointer to the memory region structure.
387  * @page_offset: Offset of the page within the region.
388  * @page_count: Number of pages to handle.
389  *
390  * This function resolves memory faults for a specified range of pages
391  * within a memory region. It uses HMM (Heterogeneous Memory Management)
392  * to fault in the required pages and updates the region's page array.
393  *
394  * Return: 0 on success, negative error code on failure.
395  */
396 static int mshv_region_range_fault(struct mshv_mem_region *region,
397 				   u64 page_offset, u64 page_count)
398 {
399 	struct hmm_range range = {
400 		.notifier = &region->mreg_mni,
401 		.default_flags = HMM_PFN_REQ_FAULT | HMM_PFN_REQ_WRITE,
402 	};
403 	unsigned long *pfns;
404 	int ret;
405 	u64 i;
406 
407 	pfns = kmalloc_array(page_count, sizeof(*pfns), GFP_KERNEL);
408 	if (!pfns)
409 		return -ENOMEM;
410 
411 	range.hmm_pfns = pfns;
412 	range.start = region->start_uaddr + page_offset * HV_HYP_PAGE_SIZE;
413 	range.end = range.start + page_count * HV_HYP_PAGE_SIZE;
414 
415 again:
416 	ret = hmm_range_fault_unlocked_timeout(&range, 0);
417 	if (ret)
418 		goto out;
419 
420 	mutex_lock(&region->mreg_mutex);
421 
422 	if (mmu_interval_read_retry(range.notifier, range.notifier_seq)) {
423 		mutex_unlock(&region->mreg_mutex);
424 		cond_resched();
425 		goto again;
426 	}
427 
428 	for (i = 0; i < page_count; i++)
429 		region->mreg_pages[page_offset + i] = hmm_pfn_to_page(pfns[i]);
430 
431 	ret = mshv_region_remap_pages(region, region->hv_map_flags,
432 				      page_offset, page_count);
433 
434 	mutex_unlock(&region->mreg_mutex);
435 out:
436 	kfree(pfns);
437 	return ret;
438 }
439 
440 bool mshv_region_handle_gfn_fault(struct mshv_mem_region *region, u64 gfn)
441 {
442 	u64 page_offset, page_count;
443 	int ret;
444 
445 	/* Align the page offset to the nearest MSHV_MAP_FAULT_IN_PAGES. */
446 	page_offset = ALIGN_DOWN(gfn - region->start_gfn,
447 				 MSHV_MAP_FAULT_IN_PAGES);
448 
449 	/* Map more pages than requested to reduce the number of faults. */
450 	page_count = min(region->nr_pages - page_offset,
451 			 MSHV_MAP_FAULT_IN_PAGES);
452 
453 	ret = mshv_region_range_fault(region, page_offset, page_count);
454 
455 	WARN_ONCE(ret,
456 		  "p%llu: GPA intercept failed: region %#llx-%#llx, gfn %#llx, page_offset %llu, page_count %llu\n",
457 		  region->partition->pt_id, region->start_uaddr,
458 		  region->start_uaddr + (region->nr_pages << HV_HYP_PAGE_SHIFT),
459 		  gfn, page_offset, page_count);
460 
461 	return !ret;
462 }
463 
464 /**
465  * mshv_region_interval_invalidate - Invalidate a range of memory region
466  * @mni: Pointer to the mmu_interval_notifier structure
467  * @range: Pointer to the mmu_notifier_range structure
468  * @cur_seq: Current sequence number for the interval notifier
469  *
470  * This function invalidates a memory region by remapping its pages with
471  * no access permissions. It locks the region's mutex to ensure thread safety
472  * and updates the sequence number for the interval notifier. If the range
473  * is blockable, it uses a blocking lock; otherwise, it attempts a non-blocking
474  * lock and returns false if unsuccessful.
475  *
476  * NOTE: Failure to invalidate a region is a serious error, as the pages will
477  * be considered freed while they are still mapped by the hypervisor.
478  * Any attempt to access such pages will likely crash the system.
479  *
480  * Return: true if the region was successfully invalidated, false otherwise.
481  */
482 static bool mshv_region_interval_invalidate(struct mmu_interval_notifier *mni,
483 					    const struct mmu_notifier_range *range,
484 					    unsigned long cur_seq)
485 {
486 	struct mshv_mem_region *region = container_of(mni,
487 						      struct mshv_mem_region,
488 						      mreg_mni);
489 	u64 page_offset, page_count;
490 	unsigned long mstart, mend;
491 	int ret = -EPERM;
492 
493 	mstart = max(range->start, region->start_uaddr);
494 	mend = min(range->end, region->start_uaddr +
495 		   (region->nr_pages << HV_HYP_PAGE_SHIFT));
496 
497 	page_offset = HVPFN_DOWN(mstart - region->start_uaddr);
498 	page_count = HVPFN_DOWN(mend - mstart);
499 
500 	if (mmu_notifier_range_blockable(range))
501 		mutex_lock(&region->mreg_mutex);
502 	else if (!mutex_trylock(&region->mreg_mutex))
503 		goto out_fail;
504 
505 	mmu_interval_set_seq(mni, cur_seq);
506 
507 	ret = mshv_region_remap_pages(region, HV_MAP_GPA_NO_ACCESS,
508 				      page_offset, page_count);
509 	if (ret)
510 		goto out_unlock;
511 
512 	mshv_region_invalidate_pages(region, page_offset, page_count);
513 
514 	mutex_unlock(&region->mreg_mutex);
515 
516 	return true;
517 
518 out_unlock:
519 	mutex_unlock(&region->mreg_mutex);
520 out_fail:
521 	WARN_ONCE(ret,
522 		  "Failed to invalidate region %#llx-%#llx (range %#lx-%#lx, event: %u, pages %#llx-%#llx, mm: %#llx): %d\n",
523 		  region->start_uaddr,
524 		  region->start_uaddr + (region->nr_pages << HV_HYP_PAGE_SHIFT),
525 		  range->start, range->end, range->event,
526 		  page_offset, page_offset + page_count - 1, (u64)range->mm, ret);
527 	return false;
528 }
529 
530 static const struct mmu_interval_notifier_ops mshv_region_mni_ops = {
531 	.invalidate = mshv_region_interval_invalidate,
532 };
533 
534 void mshv_region_movable_fini(struct mshv_mem_region *region)
535 {
536 	mmu_interval_notifier_remove(&region->mreg_mni);
537 }
538 
539 bool mshv_region_movable_init(struct mshv_mem_region *region)
540 {
541 	int ret;
542 
543 	ret = mmu_interval_notifier_insert(&region->mreg_mni, current->mm,
544 					   region->start_uaddr,
545 					   region->nr_pages << HV_HYP_PAGE_SHIFT,
546 					   &mshv_region_mni_ops);
547 	if (ret)
548 		return false;
549 
550 	mutex_init(&region->mreg_mutex);
551 
552 	return true;
553 }
554