xref: /linux/drivers/infiniband/hw/mana/mr.c (revision 3a2c4d55e32ad65efebdb6de44eef3bfa08bb49d)
1 // SPDX-License-Identifier: GPL-2.0-only
2 /*
3  * Copyright (c) 2022, Microsoft Corporation. All rights reserved.
4  */
5 
6 #include "mana_ib.h"
7 
8 #define VALID_MR_FLAGS (IB_ACCESS_LOCAL_WRITE | IB_ACCESS_REMOTE_WRITE | IB_ACCESS_REMOTE_READ |\
9 			IB_ACCESS_REMOTE_ATOMIC | IB_ACCESS_MW_BIND | IB_ZERO_BASED)
10 
11 #define VALID_DMA_MR_FLAGS (IB_ACCESS_LOCAL_WRITE)
12 
13 static enum gdma_mr_access_flags
14 mana_ib_verbs_to_gdma_access_flags(int access_flags)
15 {
16 	enum gdma_mr_access_flags flags = GDMA_ACCESS_FLAG_LOCAL_READ;
17 
18 	if (access_flags & IB_ACCESS_LOCAL_WRITE)
19 		flags |= GDMA_ACCESS_FLAG_LOCAL_WRITE;
20 
21 	if (access_flags & IB_ACCESS_REMOTE_WRITE)
22 		flags |= GDMA_ACCESS_FLAG_REMOTE_WRITE;
23 
24 	if (access_flags & IB_ACCESS_REMOTE_READ)
25 		flags |= GDMA_ACCESS_FLAG_REMOTE_READ;
26 
27 	if (access_flags & IB_ACCESS_REMOTE_ATOMIC)
28 		flags |= GDMA_ACCESS_FLAG_REMOTE_ATOMIC;
29 
30 	if (access_flags & IB_ACCESS_MW_BIND)
31 		flags |= GDMA_ACCESS_FLAG_BIND_MW;
32 
33 	return flags;
34 }
35 
36 static int mana_ib_gd_create_mr(struct mana_ib_dev *dev, struct mana_ib_mr *mr,
37 				struct gdma_create_mr_params *mr_params)
38 {
39 	struct gdma_create_mr_response resp = {};
40 	struct gdma_create_mr_request req = {};
41 	struct gdma_context *gc = mdev_to_gc(dev);
42 	int err;
43 
44 	mana_gd_init_req_hdr(&req.hdr, GDMA_CREATE_MR, sizeof(req),
45 			     sizeof(resp));
46 	req.hdr.req.msg_version = GDMA_MESSAGE_V2;
47 	req.pd_handle = mr_params->pd_handle;
48 	req.mr_type = mr_params->mr_type;
49 
50 	switch (mr_params->mr_type) {
51 	case GDMA_MR_TYPE_GPA:
52 		break;
53 	case GDMA_MR_TYPE_GVA:
54 		req.gva.dma_region_handle = mr_params->gva.dma_region_handle;
55 		req.gva.virtual_address = mr_params->gva.virtual_address;
56 		req.gva.access_flags = mr_params->gva.access_flags;
57 		break;
58 	case GDMA_MR_TYPE_ZBVA:
59 		req.zbva.dma_region_handle = mr_params->zbva.dma_region_handle;
60 		req.zbva.access_flags = mr_params->zbva.access_flags;
61 		break;
62 	case GDMA_MR_TYPE_DM:
63 		req.da_ext.length = mr_params->da.length;
64 		req.da.dm_handle = mr_params->da.dm_handle;
65 		req.da.offset = mr_params->da.offset;
66 		req.da.access_flags = mr_params->da.access_flags;
67 		break;
68 	default:
69 		ibdev_dbg(&dev->ib_dev,
70 			  "invalid param (GDMA_MR_TYPE) passed, type %d\n",
71 			  req.mr_type);
72 		return -EINVAL;
73 	}
74 
75 	err = mana_gd_send_request(gc, sizeof(req), &req, sizeof(resp), &resp);
76 	if (err)
77 		return err;
78 
79 	mr->ibmr.lkey = resp.lkey;
80 	mr->ibmr.rkey = resp.rkey;
81 	mr->mr_handle = resp.mr_handle;
82 
83 	return 0;
84 }
85 
86 static int mana_ib_gd_destroy_mr(struct mana_ib_dev *dev, u64 mr_handle)
87 {
88 	struct gdma_destroy_mr_response resp = {};
89 	struct gdma_destroy_mr_request req = {};
90 	struct gdma_context *gc = mdev_to_gc(dev);
91 
92 	mana_gd_init_req_hdr(&req.hdr, GDMA_DESTROY_MR, sizeof(req),
93 			     sizeof(resp));
94 
95 	req.mr_handle = mr_handle;
96 
97 	return mana_gd_send_request(gc, sizeof(req), &req, sizeof(resp), &resp);
98 }
99 
100 struct ib_mr *mana_ib_reg_user_mr(struct ib_pd *ibpd, u64 start, u64 length,
101 				  u64 iova, int access_flags,
102 				  struct ib_dmah *dmah,
103 				  struct ib_udata *udata)
104 {
105 	struct mana_ib_pd *pd = container_of(ibpd, struct mana_ib_pd, ibpd);
106 	struct gdma_create_mr_params mr_params = {};
107 	struct ib_device *ibdev = ibpd->device;
108 	struct mana_ib_dev *dev;
109 	struct mana_ib_mr *mr;
110 	u64 dma_region_handle;
111 	int err;
112 
113 	if (dmah)
114 		return ERR_PTR(-EOPNOTSUPP);
115 
116 	err = ib_no_udata_io(udata);
117 	if (err)
118 		return ERR_PTR(err);
119 
120 	dev = container_of(ibdev, struct mana_ib_dev, ib_dev);
121 
122 	ibdev_dbg(ibdev,
123 		  "start 0x%llx, iova 0x%llx length 0x%llx access_flags 0x%x",
124 		  start, iova, length, access_flags);
125 
126 	access_flags &= ~IB_ACCESS_OPTIONAL;
127 	if (access_flags & ~VALID_MR_FLAGS)
128 		return ERR_PTR(-EINVAL);
129 
130 	mr = kzalloc_obj(*mr);
131 	if (!mr)
132 		return ERR_PTR(-ENOMEM);
133 
134 	mr->umem = ib_umem_get_va(ibdev, start, length, access_flags);
135 	if (IS_ERR(mr->umem)) {
136 		err = PTR_ERR(mr->umem);
137 		ibdev_dbg(ibdev,
138 			  "Failed to get umem for register user-mr, %pe\n",
139 			  mr->umem);
140 		goto err_free;
141 	}
142 
143 	err = mana_ib_create_dma_region(dev, mr->umem, &dma_region_handle, iova);
144 	if (err) {
145 		ibdev_dbg(ibdev, "Failed create dma region for user-mr, %d\n",
146 			  err);
147 		goto err_umem;
148 	}
149 
150 	ibdev_dbg(ibdev,
151 		  "created dma region for user-mr 0x%llx\n",
152 		  dma_region_handle);
153 
154 	mr_params.pd_handle = pd->pd_handle;
155 	if (access_flags & IB_ZERO_BASED) {
156 		mr_params.mr_type = GDMA_MR_TYPE_ZBVA;
157 		mr_params.zbva.dma_region_handle = dma_region_handle;
158 		mr_params.zbva.access_flags =
159 			mana_ib_verbs_to_gdma_access_flags(access_flags);
160 	} else {
161 		mr_params.mr_type = GDMA_MR_TYPE_GVA;
162 		mr_params.gva.dma_region_handle = dma_region_handle;
163 		mr_params.gva.virtual_address = iova;
164 		mr_params.gva.access_flags =
165 			mana_ib_verbs_to_gdma_access_flags(access_flags);
166 	}
167 
168 	err = mana_ib_gd_create_mr(dev, mr, &mr_params);
169 	if (err)
170 		goto err_dma_region;
171 
172 	/*
173 	 * There is no need to keep track of dma_region_handle after MR is
174 	 * successfully created. The dma_region_handle is tracked in the PF
175 	 * as part of the lifecycle of this MR.
176 	 */
177 
178 	return &mr->ibmr;
179 
180 err_dma_region:
181 	mana_gd_destroy_dma_region(mdev_to_gc(dev), dma_region_handle);
182 
183 err_umem:
184 	ib_umem_release(mr->umem);
185 
186 err_free:
187 	kfree(mr);
188 	return ERR_PTR(err);
189 }
190 
191 struct ib_mr *mana_ib_reg_user_mr_dmabuf(struct ib_pd *ibpd, u64 start, u64 length,
192 					 u64 iova, int fd, int access_flags,
193 					 struct ib_dmah *dmah,
194 					 struct uverbs_attr_bundle *attrs)
195 {
196 	struct mana_ib_pd *pd = container_of(ibpd, struct mana_ib_pd, ibpd);
197 	struct gdma_create_mr_params mr_params = {};
198 	struct ib_device *ibdev = ibpd->device;
199 	struct ib_umem_dmabuf *umem_dmabuf;
200 	struct mana_ib_dev *dev;
201 	struct mana_ib_mr *mr;
202 	u64 dma_region_handle;
203 	int err;
204 
205 	if (dmah)
206 		return ERR_PTR(-EOPNOTSUPP);
207 
208 	dev = container_of(ibdev, struct mana_ib_dev, ib_dev);
209 
210 	access_flags &= ~IB_ACCESS_OPTIONAL;
211 	if (access_flags & ~VALID_MR_FLAGS)
212 		return ERR_PTR(-EOPNOTSUPP);
213 
214 	mr = kzalloc_obj(*mr);
215 	if (!mr)
216 		return ERR_PTR(-ENOMEM);
217 
218 	umem_dmabuf = ib_umem_dmabuf_get_pinned(ibdev, start, length, fd, access_flags);
219 	if (IS_ERR(umem_dmabuf)) {
220 		err = PTR_ERR(umem_dmabuf);
221 		ibdev_dbg(ibdev, "Failed to get dmabuf umem, %pe\n",
222 			  umem_dmabuf);
223 		goto err_free;
224 	}
225 
226 	mr->umem = &umem_dmabuf->umem;
227 
228 	err = mana_ib_create_dma_region(dev, mr->umem, &dma_region_handle, iova);
229 	if (err) {
230 		ibdev_dbg(ibdev, "Failed create dma region for user-mr, %d\n",
231 			  err);
232 		goto err_umem;
233 	}
234 
235 	mr_params.pd_handle = pd->pd_handle;
236 	mr_params.mr_type = GDMA_MR_TYPE_GVA;
237 	mr_params.gva.dma_region_handle = dma_region_handle;
238 	mr_params.gva.virtual_address = iova;
239 	mr_params.gva.access_flags =
240 		mana_ib_verbs_to_gdma_access_flags(access_flags);
241 
242 	err = mana_ib_gd_create_mr(dev, mr, &mr_params);
243 	if (err)
244 		goto err_dma_region;
245 
246 	/*
247 	 * There is no need to keep track of dma_region_handle after MR is
248 	 * successfully created. The dma_region_handle is tracked in the PF
249 	 * as part of the lifecycle of this MR.
250 	 */
251 
252 	return &mr->ibmr;
253 
254 err_dma_region:
255 	mana_gd_destroy_dma_region(mdev_to_gc(dev), dma_region_handle);
256 
257 err_umem:
258 	ib_umem_release(mr->umem);
259 
260 err_free:
261 	kfree(mr);
262 	return ERR_PTR(err);
263 }
264 
265 struct ib_mr *mana_ib_get_dma_mr(struct ib_pd *ibpd, int access_flags)
266 {
267 	struct mana_ib_pd *pd = container_of(ibpd, struct mana_ib_pd, ibpd);
268 	struct gdma_create_mr_params mr_params = {};
269 	struct ib_device *ibdev = ibpd->device;
270 	struct mana_ib_dev *dev;
271 	struct mana_ib_mr *mr;
272 	int err;
273 
274 	dev = container_of(ibdev, struct mana_ib_dev, ib_dev);
275 
276 	if (access_flags & ~VALID_DMA_MR_FLAGS)
277 		return ERR_PTR(-EINVAL);
278 
279 	mr = kzalloc_obj(*mr);
280 	if (!mr)
281 		return ERR_PTR(-ENOMEM);
282 
283 	mr_params.pd_handle = pd->pd_handle;
284 	mr_params.mr_type = GDMA_MR_TYPE_GPA;
285 
286 	err = mana_ib_gd_create_mr(dev, mr, &mr_params);
287 	if (err)
288 		goto err_free;
289 
290 	return &mr->ibmr;
291 
292 err_free:
293 	kfree(mr);
294 	return ERR_PTR(err);
295 }
296 
297 static int mana_ib_gd_create_mw(struct mana_ib_dev *dev, struct mana_ib_pd *pd, struct ib_mw *ibmw)
298 {
299 	struct mana_ib_mw *mw = container_of(ibmw, struct mana_ib_mw, ibmw);
300 	struct gdma_context *gc = mdev_to_gc(dev);
301 	struct gdma_create_mr_response resp = {};
302 	struct gdma_create_mr_request req = {};
303 	int err;
304 
305 	mana_gd_init_req_hdr(&req.hdr, GDMA_CREATE_MR, sizeof(req), sizeof(resp));
306 	req.hdr.req.msg_version = GDMA_MESSAGE_V2;
307 	req.pd_handle = pd->pd_handle;
308 
309 	switch (mw->ibmw.type) {
310 	case IB_MW_TYPE_1:
311 		req.mr_type = GDMA_MR_TYPE_MW1;
312 		break;
313 	case IB_MW_TYPE_2:
314 		req.mr_type = GDMA_MR_TYPE_MW2;
315 		break;
316 	default:
317 		return -EINVAL;
318 	}
319 
320 	err = mana_gd_send_request(gc, sizeof(req), &req, sizeof(resp), &resp);
321 	if (err)
322 		return err;
323 
324 	mw->ibmw.rkey = resp.rkey;
325 	mw->mw_handle = resp.mr_handle;
326 
327 	return 0;
328 }
329 
330 int mana_ib_alloc_mw(struct ib_mw *ibmw, struct ib_udata *udata)
331 {
332 	struct mana_ib_dev *mdev = container_of(ibmw->device, struct mana_ib_dev, ib_dev);
333 	struct mana_ib_pd *pd = container_of(ibmw->pd, struct mana_ib_pd, ibpd);
334 	int err;
335 
336 	err = ib_no_udata_io(udata);
337 	if (err)
338 		return err;
339 
340 	return mana_ib_gd_create_mw(mdev, pd, ibmw);
341 }
342 
343 int mana_ib_dealloc_mw(struct ib_mw *ibmw)
344 {
345 	struct mana_ib_dev *dev = container_of(ibmw->device, struct mana_ib_dev, ib_dev);
346 	struct mana_ib_mw *mw = container_of(ibmw, struct mana_ib_mw, ibmw);
347 
348 	return mana_ib_gd_destroy_mr(dev, mw->mw_handle);
349 }
350 
351 int mana_ib_dereg_mr(struct ib_mr *ibmr, struct ib_udata *udata)
352 {
353 	struct mana_ib_mr *mr = container_of(ibmr, struct mana_ib_mr, ibmr);
354 	struct ib_device *ibdev = ibmr->device;
355 	struct mana_ib_dev *dev;
356 	int err;
357 
358 	err = ib_no_udata_io(udata);
359 	if (err)
360 		return err;
361 
362 	dev = container_of(ibdev, struct mana_ib_dev, ib_dev);
363 
364 	err = mana_ib_gd_destroy_mr(dev, mr->mr_handle);
365 	if (err)
366 		return err;
367 
368 	if (mr->umem)
369 		ib_umem_release(mr->umem);
370 
371 	kfree(mr);
372 
373 	return 0;
374 }
375 
376 static int mana_ib_gd_alloc_dm(struct mana_ib_dev *mdev, struct mana_ib_dm *dm,
377 			       struct ib_dm_alloc_attr *attr)
378 {
379 	struct gdma_context *gc = mdev_to_gc(mdev);
380 	struct gdma_alloc_dm_resp resp = {};
381 	struct gdma_alloc_dm_req req = {};
382 	int err;
383 
384 	mana_gd_init_req_hdr(&req.hdr, GDMA_ALLOC_DM, sizeof(req), sizeof(resp));
385 	req.length = attr->length;
386 	req.alignment = attr->alignment;
387 	req.flags =  attr->flags;
388 
389 	err = mana_gd_send_request(gc, sizeof(req), &req, sizeof(resp), &resp);
390 	if (err)
391 		return err;
392 
393 	dm->dm_handle = resp.dm_handle;
394 
395 	return 0;
396 }
397 
398 struct ib_dm *mana_ib_alloc_dm(struct ib_device *ibdev,
399 			       struct ib_ucontext *context,
400 			       struct ib_dm_alloc_attr *attr,
401 			       struct uverbs_attr_bundle *attrs)
402 {
403 	struct mana_ib_dev *dev = container_of(ibdev, struct mana_ib_dev, ib_dev);
404 	struct mana_ib_dm *dm;
405 	int err;
406 
407 	dm = kzalloc_obj(*dm);
408 	if (!dm)
409 		return ERR_PTR(-ENOMEM);
410 
411 	err = mana_ib_gd_alloc_dm(dev, dm, attr);
412 	if (err)
413 		goto err_free;
414 
415 	return &dm->ibdm;
416 
417 err_free:
418 	kfree(dm);
419 	return ERR_PTR(err);
420 }
421 
422 static int mana_ib_gd_destroy_dm(struct mana_ib_dev *mdev, struct mana_ib_dm *dm)
423 {
424 	struct gdma_context *gc = mdev_to_gc(mdev);
425 	struct gdma_destroy_dm_resp resp = {};
426 	struct gdma_destroy_dm_req req = {};
427 
428 	mana_gd_init_req_hdr(&req.hdr, GDMA_DESTROY_DM, sizeof(req), sizeof(resp));
429 	req.dm_handle = dm->dm_handle;
430 
431 	return mana_gd_send_request(gc, sizeof(req), &req, sizeof(resp), &resp);
432 }
433 
434 int mana_ib_dealloc_dm(struct ib_dm *ibdm, struct uverbs_attr_bundle *attrs)
435 {
436 	struct mana_ib_dev *dev = container_of(ibdm->device, struct mana_ib_dev, ib_dev);
437 	struct mana_ib_dm *dm = container_of(ibdm, struct mana_ib_dm, ibdm);
438 	int err;
439 
440 	err = mana_ib_gd_destroy_dm(dev, dm);
441 	if (err)
442 		return err;
443 
444 	kfree(dm);
445 	return 0;
446 }
447 
448 struct ib_mr *mana_ib_reg_dm_mr(struct ib_pd *ibpd, struct ib_dm *ibdm,
449 				struct ib_dm_mr_attr *attr,
450 				struct uverbs_attr_bundle *attrs)
451 {
452 	struct mana_ib_dev *dev = container_of(ibpd->device, struct mana_ib_dev, ib_dev);
453 	struct mana_ib_dm *mana_dm = container_of(ibdm, struct mana_ib_dm, ibdm);
454 	struct mana_ib_pd *pd = container_of(ibpd, struct mana_ib_pd, ibpd);
455 	struct gdma_create_mr_params mr_params = {};
456 	struct mana_ib_mr *mr;
457 	int err;
458 
459 	attr->access_flags &= ~IB_ACCESS_OPTIONAL;
460 	if (attr->access_flags & ~VALID_MR_FLAGS)
461 		return ERR_PTR(-EOPNOTSUPP);
462 
463 	mr = kzalloc_obj(*mr);
464 	if (!mr)
465 		return ERR_PTR(-ENOMEM);
466 
467 	mr_params.pd_handle = pd->pd_handle;
468 	mr_params.mr_type = GDMA_MR_TYPE_DM;
469 	mr_params.da.dm_handle = mana_dm->dm_handle;
470 	mr_params.da.offset = attr->offset;
471 	mr_params.da.length = attr->length;
472 	mr_params.da.access_flags =
473 		mana_ib_verbs_to_gdma_access_flags(attr->access_flags);
474 
475 	err = mana_ib_gd_create_mr(dev, mr, &mr_params);
476 	if (err)
477 		goto err_free;
478 
479 	return &mr->ibmr;
480 
481 err_free:
482 	kfree(mr);
483 	return ERR_PTR(err);
484 }
485