1#include <asm.h>
2#include <errno.h>
3#include <kassert.h>
4#include <math/align.h>
5#include <mem/anon_vma.h>
6#include <mem/fixed_size_alloc.h>
7#include <mem/mm.h>
8#include <mem/page.h>
9#include <mem/vmm.h>
10#include <smp/perdomain.h>
11#include <structures/rbt.h>
12#include <types/refcount.h>
13
14#define MM_USER_MIN 0x0000000000010000UL
15#define MM_USER_MAX 0x0000800000000000UL
16
17FIXED_SIZE_RANGE_PERDOMAIN_DECLARE(mm, .obj_size = sizeof(struct mm),
18 .obj_align = _Alignof(struct mm));
19
20/* VMA tree is rbit keyed [vma_range_start, vma_range_end - 1], carrying custom
21 * agumentation to track max_gap to make gap search O(log n)
22 *
23 * Non overlapping VMAs make 'left subtree max end' exactly the
24 * predecessor's end, so no neighbor ptr needed here */
25static size_t vma_range_node_min_low(struct rbit_node *n) {
26 return n ? rbit_entry(n, struct vma_range, mm_node)->min_low : 0;
27}
28
29static size_t vma_range_node_gap(struct rbit_node *n) {
30 return n ? rbit_entry(n, struct vma_range, mm_node)->max_gap : 0;
31}
32
33static bool vma_range_tree_augment(struct rbit_node *n) {
34 struct vma_range *v = rbit_entry(n, struct vma_range, mm_node);
35 size_t old_max = n->max, old_min = v->min_low, old_gap = v->max_gap;
36
37 /* node->max: subtree max of interval.high (= max end - 1) */
38 size_t mx = n->interval.high;
39 if (rbit_node_max(n: n->left) > mx)
40 mx = rbit_node_max(n: n->left);
41
42 if (rbit_node_max(n: n->right) > mx)
43 mx = rbit_node_max(n: n->right);
44
45 size_t mn = n->left ? vma_range_node_min_low(n: n->left) : n->interval.low;
46
47 /* max_gap: biggest free run between consecutive VMAs in the subtree */
48 size_t gap = 0;
49 if (n->left) {
50 if (vma_range_node_gap(n: n->left) > gap)
51 gap = vma_range_node_gap(n: n->left);
52
53 size_t pred_end = rbit_node_max(n: n->left) + 1; /* predecessor's end */
54 if (n->interval.low > pred_end && n->interval.low - pred_end > gap)
55 gap = n->interval.low - pred_end;
56 }
57 if (n->right) {
58 if (vma_range_node_gap(n: n->right) > gap)
59 gap = vma_range_node_gap(n: n->right);
60
61 size_t end = n->interval.high + 1; /* this VMA's end */
62 size_t succ_low = vma_range_node_min_low(n: n->right);
63 if (succ_low > end && succ_low - end > gap)
64 gap = succ_low - end;
65 }
66
67 n->max = mx;
68 v->min_low = mn;
69 v->max_gap = gap;
70 return mx != old_max || mn != old_min || gap != old_gap;
71}
72
73enum errno mm_pgtable_init(struct mm *mm) {
74 mm->pml4 = vmm_make_user_pml4();
75 return mm->pml4 ? ERR_OK : ERR_NO_MEM;
76}
77
78void mm_pgtable_free(struct mm *mm) {
79 (void) mm;
80 vmm_unmap_all_user_pages(pml4: vmm_phys_to_pml4(paddr: mm->pml4), vflags: VMM_FLAG_NONE);
81}
82
83void mm_activate(struct mm *mm) {
84 write_cr3(cr3: mm->pml4);
85}
86
87enum errno mm_map_page(struct mm *mm, vaddr_t va, paddr_t pa, uint64_t pflags) {
88 vmm_map_page_user(vmm_phys_to_pml4(mm->pml4), va, pa, pflags,
89 VMM_FLAG_USER);
90 return ERR_OK;
91}
92
93struct mm *mm_alloc(void) {
94 struct mm *mm = FSR_PERDOMAIN_ALLOC(mm);
95 if (!mm)
96 return NULL;
97
98 rbit_init(rbit: &mm->vma_range_tree);
99 mm->vma_range_tree.augment =
100 vma_range_tree_augment; /* O(log n) gap search */
101 /* Anyone can be under this rwlock */
102 rwlock_init(lock: &mm->lock, ceiling: THREAD_PRIO_CLASS_URGENT);
103 mm->vas = NULL;
104 mm->mmap_cursor = MM_USER_MIN;
105
106 refcount_init(rc: &mm->users, val: 1);
107 refcount_init(rc: &mm->refcount, val: 1);
108
109 if (mm_pgtable_init(mm) != ERR_OK) {
110 FSR_PERDOMAIN_FREE(mm, mm);
111 return NULL;
112 }
113 return mm;
114}
115
116void mm_free(struct mm *mm) {
117 kassert(rbit_empty(&mm->vma_range_tree));
118 mm_pgtable_free(mm);
119 FSR_PERDOMAIN_FREE(mm, mm);
120}
121
122void mm_exit(struct mm *mm) {
123 struct rbit_node *n = rbit_first(tree: &mm->vma_range_tree);
124 while (n) {
125 struct rbit_node *next = rbit_next(node: n);
126 struct vma_range *vma_range = rbit_entry(n, struct vma_range, mm_node);
127
128 mm_vma_range_remove(mm, vma_range);
129 vma_range_unlink_anon_vmas(vma_range); /* drops AVCs + anon_vma refs */
130 /* TODO: drop this VMA's PTEs + folio rmap once page
131 * tables can be torn down. For now the mappings die with the pml4 */
132 vma_range_free(vma_range);
133
134 n = next;
135 }
136}
137
138struct vma_range *mm_vma_range_find(struct mm *mm, vaddr_t addr) {
139 struct rbit_node *node = mm->vma_range_tree.root;
140 struct vma_range *result = NULL;
141
142 while (node) {
143 struct vma_range *v = rbit_entry(node, struct vma_range, mm_node);
144 if (vma_range_end(vma_range: v) > addr) {
145 result = v; /* candidate; look left for an earlier one */
146 node = node->left;
147 } else {
148 node = node->right;
149 }
150 }
151 return result;
152}
153
154struct vma_range *mm_vma_range_find_intersection(struct mm *mm, vaddr_t s,
155 vaddr_t e) {
156 /* Inclusive interval [s, e-1] matches the half-open [s, e) reservation */
157 struct interval iv = {.low = s, .high = e - 1};
158 struct rbit_node *n = rbit_overlap_search(root: mm->vma_range_tree.root, iv);
159 return n ? rbit_entry(n, struct vma_range, mm_node) : NULL;
160}
161
162void mm_vma_range_insert(struct mm *mm, struct vma_range *vma_range) {
163 /* vma_range_init() already set the interval; rbit_insert() does the rest */
164 kassert(!mm_vma_range_find_intersection(mm, vma_range_start(vma_range),
165 vma_range_end(vma_range)));
166 rbit_insert(tree: &mm->vma_range_tree, new_node: &vma_range->mm_node);
167}
168
169void mm_vma_range_remove(struct mm *mm, struct vma_range *vma_range) {
170 rbit_delete(tree: &mm->vma_range_tree, z: &vma_range->mm_node);
171}
172
173struct gap_ctx {
174 size_t len;
175 size_t align;
176 vaddr_t high;
177};
178
179static vaddr_t gap_fits(vaddr_t floor, vaddr_t ceil, const struct gap_ctx *c) {
180 if (ceil <= floor)
181 return 0;
182 vaddr_t a = ALIGN_UP(floor, c->align);
183 if (a < floor) /* ALIGN_UP wrapped */
184 return 0;
185 if (a + c->len < a) /* len overflow */
186 return 0;
187 if (a + c->len > ceil || a + c->len > c->high)
188 return 0;
189 return a;
190}
191
192static vaddr_t gap_search(struct rbit_node *n, vaddr_t *floor,
193 const struct gap_ctx *c) {
194 if (!n)
195 return 0;
196 struct vma_range *v = rbit_entry(n, struct vma_range, mm_node);
197
198 if (n->left) {
199 if (vma_range_node_gap(n: n->left) >= c->len ||
200 vma_range_node_min_low(n: n->left) > *floor) {
201 vaddr_t r = gap_search(n: n->left, floor, c);
202 if (r)
203 return r;
204 } else { /* prune: jump floor to the subtree's max end */
205 vaddr_t e = rbit_node_max(n: n->left) + 1;
206 if (e > *floor)
207 *floor = e;
208 }
209 }
210
211 /* gap right before this VMA */
212 vaddr_t r = gap_fits(floor: *floor, ceil: vma_range_start(vma_range: v), c);
213 if (r)
214 return r;
215 if (vma_range_start(vma_range: v) >=
216 c->high) /* this VMA and all to its right are out */
217 return 0;
218 if (vma_range_end(vma_range: v) > *floor)
219 *floor = vma_range_end(vma_range: v);
220
221 if (n->right) {
222 if (vma_range_node_gap(n: n->right) >= c->len ||
223 vma_range_node_min_low(n: n->right) > *floor) {
224 vaddr_t rr = gap_search(n: n->right, floor, c);
225 if (rr)
226 return rr;
227 } else {
228 vaddr_t e = rbit_node_max(n: n->right) + 1;
229 if (e > *floor)
230 *floor = e;
231 }
232 }
233 return 0;
234}
235
236vaddr_t mm_vma_range_find_gap(struct mm *mm, size_t len, size_t align,
237 vaddr_t low, vaddr_t high) {
238 kassert(align != 0);
239 if (len == 0 || high <= low || high - low < len)
240 return 0;
241
242 struct gap_ctx c = {.len = len, .align = align, .high = high};
243 vaddr_t floor = low;
244
245 vaddr_t r = gap_search(n: mm->vma_range_tree.root, floor: &floor, c: &c);
246 if (r)
247 return r;
248
249 /* trailing gap [floor, high) */
250 return gap_fits(floor, ceil: high, c: &c);
251}
252
253vaddr_t mm_map(struct mm *mm, vaddr_t hint, size_t len,
254 enum vma_range_protection prot, enum mm_map_flags flags) {
255 len = PAGE_ALIGN_UP(len);
256 if (len == 0)
257 return 0;
258
259 rwlock_write_lock(lock: &mm->lock);
260
261 vaddr_t addr;
262 if (flags & MM_MAP_FIXED) {
263 addr = PAGE_ALIGN_DOWN(hint);
264 if (mm_vma_range_find_intersection(mm, s: addr, e: addr + len)) {
265 /* TODO: clobber overlap via mm_unmap()
266 * once it can tear down PTEs */
267
268 /* FIXED map on busy range fails */
269 rwlock_unlock(lock: &mm->lock);
270 return 0;
271 }
272 } else {
273 vaddr_t low = hint ? PAGE_ALIGN_UP(hint) : mm->mmap_cursor;
274 addr = mm_vma_range_find_gap(mm, len, PAGE_SIZE, low, MM_USER_MAX);
275 if (!addr)
276 addr = mm_vma_range_find_gap(mm, len, PAGE_SIZE, MM_USER_MIN,
277 MM_USER_MAX);
278
279 if (!addr) {
280 rwlock_unlock(lock: &mm->lock);
281 return 0;
282 }
283 if (!hint) /* advance the cursor past what we just handed out */
284 mm->mmap_cursor = addr + len;
285 }
286
287 /* anon reservation only, pages fault in lazily, so no PTEs */
288 struct vma_range *vma_range = vma_range_alloc(mm, start: addr, end: addr + len, prot);
289 if (!vma_range) {
290 rwlock_unlock(lock: &mm->lock);
291 return 0;
292 }
293 mm_vma_range_insert(mm, vma_range);
294
295 rwlock_unlock(lock: &mm->lock);
296 return addr;
297}
298