1#include <asm.h>
2#include <console/printf.h>
3#include <drivers/iommu/vt_d.h>
4#include <drivers/mmio.h>
5#include <log.h>
6#include <mem/hhdm.h>
7#include <mem/page.h>
8#include <mem/pmm.h>
9#include <mem/vmm.h>
10#include <string.h>
11
12#include "internal.h"
13
14static bool vtd_pt_empty(uint64_t *pt) {
15 for (size_t i = 0; i < SL_ENTRY_COUNT; i++) {
16 if (pt[i] & SL_PTE_READ)
17 return false;
18 }
19 return true;
20}
21
22static uint64_t *vtd_sl_ensure_child(sl_pte_atomic_t *entry) {
23 uint64_t val = atomic_load_explicit(entry, memory_order_relaxed);
24 val &= ~SL_PTE_LOCK_BIT;
25
26 if (val & SL_PTE_READ)
27 return hhdm_paddr_to_ptr(SL_PTE_ADDR(val));
28
29 paddr_t phys = pmm_alloc_page();
30 if (!phys)
31 return NULL;
32
33 void *virt = hhdm_paddr_to_ptr(p: phys);
34 memset(virt, 0, PAGE_SIZE);
35
36 atomic_store_explicit(entry, SL_TABLE_ENTRY(phys) | SL_PTE_LOCK_BIT,
37 memory_order_release);
38 return virt;
39}
40
41static enum iommu_error vtd_sl_map_page(uint64_t *sl_pgd, iova_t iova,
42 paddr_t pa, uint32_t perm) {
43 sl_pte_atomic_t *locked[PT_LEVELS] = {NULL, NULL, NULL, NULL};
44 enum irql saved[PT_LEVELS];
45
46 sl_pte_atomic_t *e4 = (sl_pte_atomic_t *) &sl_pgd[SL_PML4_INDEX(iova)];
47 saved[0] = vtd_pt_lock(pte: e4);
48 locked[0] = e4;
49
50 uint64_t *pdp = vtd_sl_ensure_child(entry: e4);
51 if (!pdp)
52 goto fail;
53
54 sl_pte_atomic_t *e3 = (sl_pte_atomic_t *) &pdp[SL_PDPT_INDEX(iova)];
55 saved[1] = vtd_pt_lock(pte: e3);
56 locked[1] = e3;
57
58 uint64_t *pd = vtd_sl_ensure_child(entry: e3);
59 if (!pd)
60 goto fail;
61
62 sl_pte_atomic_t *e2 = (sl_pte_atomic_t *) &pd[SL_PD_INDEX(iova)];
63 saved[2] = vtd_pt_lock(pte: e2);
64 locked[2] = e2;
65
66 uint64_t *pt = vtd_sl_ensure_child(entry: e2);
67 if (!pt)
68 goto fail;
69
70 sl_pte_atomic_t *e1 = (sl_pte_atomic_t *) &pt[SL_PT_INDEX(iova)];
71 saved[3] = vtd_pt_lock(pte: e1);
72 locked[3] = e1;
73
74 atomic_store_explicit(e1, SL_PAGE_ENTRY(pa, perm), memory_order_release);
75
76 /* Unlock in reverse: leaf → root */
77 for (int i = PT_LEVELS - 1; i >= 0; i--)
78 vtd_pt_unlock(pte: locked[i], old_irql: saved[i]);
79
80 return IOMMU_ERR_OK;
81
82fail:
83 for (int i = PT_LEVELS - 1; i >= 0; i--) {
84 if (locked[i])
85 vtd_pt_unlock(pte: locked[i], old_irql: saved[i]);
86 }
87 return IOMMU_ERR_NO_MEM;
88}
89
90static bool vtd_sl_unmap_page(uint64_t *sl_pgd, iova_t iova) {
91 pte_t pte;
92
93 pte = sl_pgd[SL_PML4_INDEX(iova)];
94 if (!(pte & SL_PTE_READ))
95 return false;
96
97 uint64_t *pdp = hhdm_paddr_to_ptr(SL_PTE_ADDR(pte));
98 pte = pdp[SL_PDPT_INDEX(iova)];
99 if (!(pte & SL_PTE_READ))
100 return false;
101
102 uint64_t *pd = hhdm_paddr_to_ptr(SL_PTE_ADDR(pte));
103 pte = pd[SL_PD_INDEX(iova)];
104 if (!(pte & SL_PTE_READ))
105 return false;
106
107 uint64_t *pt = hhdm_paddr_to_ptr(SL_PTE_ADDR(pte));
108 if (!(pt[SL_PT_INDEX(iova)] & SL_PTE_READ))
109 return false;
110
111 pt[SL_PT_INDEX(iova)] = 0;
112 return true;
113}
114
115static paddr_t vtd_sl_translate(uint64_t *sl_pgd, iova_t iova) {
116 pte_t pte;
117
118 pte = sl_pgd[SL_PML4_INDEX(iova)];
119 if (!(pte & SL_PTE_READ))
120 return 0;
121
122 uint64_t *pdp = hhdm_paddr_to_ptr(SL_PTE_ADDR(pte));
123 pte = pdp[SL_PDPT_INDEX(iova)];
124 if (!(pte & SL_PTE_READ))
125 return 0;
126
127 uint64_t *pd = hhdm_paddr_to_ptr(SL_PTE_ADDR(pte));
128 pte = pd[SL_PD_INDEX(iova)];
129 if (!(pte & SL_PTE_READ))
130 return 0;
131
132 uint64_t *pt = hhdm_paddr_to_ptr(SL_PTE_ADDR(pte));
133 pte = pt[SL_PT_INDEX(iova)];
134 if (!(pte & SL_PTE_READ))
135 return 0;
136
137 return SL_PTE_ADDR(pte) | SL_PAGE_OFFSET(iova);
138}
139
140static enum iommu_error vtd_map(struct iommu_domain *domain, iova_t iova,
141 paddr_t pa, size_t size, uint32_t perm) {
142 struct vtd_unit *u = domain->unit->private;
143 struct vtd_domain *vd = domain->priv;
144
145 if (!IS_PAGE_ALIGNED(iova) || !IS_PAGE_ALIGNED(pa) ||
146 !IS_PAGE_ALIGNED(size))
147 return IOMMU_ERR_INVALID;
148
149 for (size_t offset = 0; offset < size; offset += PAGE_SIZE) {
150 enum iommu_error err =
151 vtd_sl_map_page(sl_pgd: vd->sl_pgd, iova: iova + offset, pa: pa + offset, perm);
152 if (err != IOMMU_ERR_OK)
153 return err;
154 }
155
156 vtd_iq_submit(u, IOTLB_INVAL_DESC_DOMAIN(vd->domain_id));
157 vtd_iq_flush(u);
158
159 return IOMMU_ERR_OK;
160}
161
162static void vtd_flush_iotlb_domain(struct iommu_domain *domain) {
163 struct vtd_unit *u = domain->unit->private;
164 struct vtd_domain *vd = domain->priv;
165
166 vtd_iq_submit(u, IOTLB_INVAL_DESC_DOMAIN(vd->domain_id));
167 vtd_iq_flush(u);
168}
169
170static void vtd_flush_iotlb_range(struct iommu_domain *domain, iova_t iova,
171 size_t size) {
172 struct vtd_unit *u = domain->unit->private;
173 struct vtd_domain *vd = domain->priv;
174
175 if (CAP_PAGE_SELECTIVE_INVALIDATION(u->cap) && size <= 32 * PAGE_SIZE) {
176 for (size_t off = 0; off < size; off += PAGE_SIZE)
177 vtd_iq_submit(u,
178 IOTLB_INVAL_DESC_PAGE(vd->domain_id, iova + off, 0));
179 } else {
180 vtd_iq_submit(u, IOTLB_INVAL_DESC_DOMAIN(vd->domain_id));
181 }
182
183 vtd_iq_flush(u);
184}
185
186struct vtd_walk_state {
187 uint64_t *tables[PT_LEVELS];
188 size_t indices[PT_LEVELS];
189};
190
191static bool vtd_sl_unmap_locked(uint64_t *sl_pgd, iova_t iova,
192 struct vtd_walk_state *ws) {
193 sl_pte_atomic_t *locked[PT_LEVELS];
194 enum irql saved[PT_LEVELS];
195
196 ws->tables[0] = sl_pgd;
197 ws->indices[0] = SL_PML4_INDEX(iova);
198 ws->indices[1] = SL_PDPT_INDEX(iova);
199 ws->indices[2] = SL_PD_INDEX(iova);
200 ws->indices[3] = SL_PT_INDEX(iova);
201
202 uint64_t *cur = sl_pgd;
203
204 for (int lvl = 0; lvl < PT_LEVELS; lvl++) {
205 sl_pte_atomic_t *entry = (sl_pte_atomic_t *) &cur[ws->indices[lvl]];
206 saved[lvl] = vtd_pt_lock(pte: entry);
207 locked[lvl] = entry;
208
209 uint64_t val = atomic_load_explicit(entry, memory_order_relaxed);
210 val &= ~SL_PTE_LOCK_BIT;
211
212 if (!(val & SL_PTE_READ)) {
213 /* not mapped, unlock everything and bail */
214 for (int i = lvl; i >= 0; i--)
215 vtd_pt_unlock(pte: locked[i], old_irql: saved[i]);
216 return false;
217 }
218
219 if (lvl == PT_LEVELS - 1) {
220 /* leaf, clear it, keep lock bit until unlock */
221 atomic_store_explicit(entry, SL_PTE_LOCK_BIT, memory_order_release);
222 } else {
223 ws->tables[lvl + 1] = hhdm_paddr_to_ptr(SL_PTE_ADDR(val));
224 cur = ws->tables[lvl + 1];
225 }
226 }
227
228 for (int i = PT_LEVELS - 1; i >= 0; i--)
229 vtd_pt_unlock(pte: locked[i], old_irql: saved[i]);
230
231 return true;
232}
233
234/* bottom-up reclaim of empty intermediate tables */
235static void vtd_sl_reclaim_walk(struct vtd_walk_state *ws) {
236 /* lvl 3 is the leaf (PT entry) which is already cleared */
237 for (int lvl = PT_LEVELS - 2; lvl >= 0; lvl--) {
238 sl_pte_atomic_t *parent =
239 (sl_pte_atomic_t *) &ws->tables[lvl][ws->indices[lvl]];
240
241 enum irql old_irql = vtd_pt_lock(pte: parent);
242
243 uint64_t val = atomic_load_explicit(parent, memory_order_relaxed);
244 val &= ~SL_PTE_LOCK_BIT;
245
246 if (!(val & SL_PTE_READ)) {
247 /* someone else already freed it */
248 vtd_pt_unlock(pte: parent, old_irql);
249 break;
250 }
251
252 uint64_t *child = ws->tables[lvl + 1];
253
254 if (!vtd_pt_empty(pt: child)) {
255 vtd_pt_unlock(pte: parent, old_irql);
256 break;
257 }
258
259 /* child is empty, clear parent entry and free child */
260 atomic_store_explicit(parent, SL_PTE_LOCK_BIT, memory_order_release);
261 vtd_pt_mark_dead(pte: parent);
262 vtd_pt_unlock(pte: parent, old_irql);
263
264 pmm_free_page(addr: hhdm_ptr_to_paddr(ptr: child));
265 }
266}
267
268static void vtd_unmap(struct iommu_domain *domain, iova_t iova, size_t size) {
269 struct vtd_unit *u = domain->unit->private;
270 struct vtd_domain *vd = domain->priv;
271
272 if (!IS_PAGE_ALIGNED(iova) || !IS_PAGE_ALIGNED(size))
273 return;
274
275 size_t page_count = size / PAGE_SIZE;
276
277 struct vtd_walk_state *walks = kmalloc(page_count * sizeof(*walks));
278 bool *unmapped = kmalloc(page_count * sizeof(bool));
279 if (!walks || !unmapped) {
280 for (size_t off = 0; off < size; off += PAGE_SIZE) {
281 struct vtd_walk_state dummy;
282 vtd_sl_unmap_locked(sl_pgd: vd->sl_pgd, iova: iova + off, ws: &dummy);
283 }
284 vtd_iotlb_flush_range_batched(u, domain_id: vd->domain_id, iova, size);
285 kfree(walks);
286 kfree(unmapped);
287 return;
288 }
289
290 /* clear all leaves */
291 for (size_t i = 0; i < page_count; i++)
292 unmapped[i] =
293 vtd_sl_unmap_locked(sl_pgd: vd->sl_pgd, iova: iova + i * PAGE_SIZE, ws: &walks[i]);
294
295 /* send invalidation */
296 vtd_iotlb_flush_range_batched(u, domain_id: vd->domain_id, iova, size);
297
298 /* go and reclaim */
299 for (size_t i = 0; i < page_count; i++) {
300 if (unmapped[i])
301 vtd_sl_reclaim_walk(ws: &walks[i]);
302 }
303
304 kfree(walks);
305 kfree(unmapped);
306}
307