dibs_loopback.c 8.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361
  1. // SPDX-License-Identifier: GPL-2.0
  2. /*
  3. * Functions for dibs loopback/loopback-ism device.
  4. *
  5. * Copyright (c) 2024, Alibaba Inc.
  6. *
  7. * Author: Wen Gu <guwen@linux.alibaba.com>
  8. * Tony Lu <tonylu@linux.alibaba.com>
  9. *
  10. */
  11. #include <linux/bitops.h>
  12. #include <linux/device.h>
  13. #include <linux/dibs.h>
  14. #include <linux/mm.h>
  15. #include <linux/slab.h>
  16. #include <linux/spinlock.h>
  17. #include <linux/types.h>
  18. #include "dibs_loopback.h"
  19. #define DIBS_LO_SUPPORT_NOCOPY 0x1
  20. #define DIBS_DMA_ADDR_INVALID (~(dma_addr_t)0)
  21. static const char dibs_lo_dev_name[] = "lo";
  22. /* global loopback device */
  23. static struct dibs_lo_dev *lo_dev;
  24. static u16 dibs_lo_get_fabric_id(struct dibs_dev *dibs)
  25. {
  26. return DIBS_LOOPBACK_FABRIC;
  27. }
  28. static int dibs_lo_query_rgid(struct dibs_dev *dibs, const uuid_t *rgid,
  29. u32 vid_valid, u32 vid)
  30. {
  31. /* rgid should be the same as lgid */
  32. if (!uuid_equal(rgid, &dibs->gid))
  33. return -ENETUNREACH;
  34. return 0;
  35. }
  36. static int dibs_lo_max_dmbs(void)
  37. {
  38. return DIBS_LO_MAX_DMBS;
  39. }
  40. static int dibs_lo_register_dmb(struct dibs_dev *dibs, struct dibs_dmb *dmb,
  41. struct dibs_client *client)
  42. {
  43. struct dibs_lo_dmb_node *dmb_node, *tmp_node;
  44. struct dibs_lo_dev *ldev;
  45. struct folio *folio;
  46. unsigned long flags;
  47. int sba_idx, rc;
  48. ldev = dibs->drv_priv;
  49. sba_idx = dmb->idx;
  50. /* check space for new dmb */
  51. for_each_clear_bit(sba_idx, ldev->sba_idx_mask, DIBS_LO_MAX_DMBS) {
  52. if (!test_and_set_bit(sba_idx, ldev->sba_idx_mask))
  53. break;
  54. }
  55. if (sba_idx == DIBS_LO_MAX_DMBS)
  56. return -ENOSPC;
  57. dmb_node = kzalloc_obj(*dmb_node);
  58. if (!dmb_node) {
  59. rc = -ENOMEM;
  60. goto err_bit;
  61. }
  62. dmb_node->sba_idx = sba_idx;
  63. dmb_node->len = dmb->dmb_len;
  64. /* not critical; fail under memory pressure and fallback to TCP */
  65. folio = folio_alloc(GFP_KERNEL | __GFP_NOWARN | __GFP_NOMEMALLOC |
  66. __GFP_NORETRY | __GFP_ZERO,
  67. get_order(dmb_node->len));
  68. if (!folio) {
  69. rc = -ENOMEM;
  70. goto err_node;
  71. }
  72. dmb_node->cpu_addr = folio_address(folio);
  73. dmb_node->dma_addr = DIBS_DMA_ADDR_INVALID;
  74. refcount_set(&dmb_node->refcnt, 1);
  75. again:
  76. /* add new dmb into hash table */
  77. get_random_bytes(&dmb_node->token, sizeof(dmb_node->token));
  78. write_lock_bh(&ldev->dmb_ht_lock);
  79. hash_for_each_possible(ldev->dmb_ht, tmp_node, list, dmb_node->token) {
  80. if (tmp_node->token == dmb_node->token) {
  81. write_unlock_bh(&ldev->dmb_ht_lock);
  82. goto again;
  83. }
  84. }
  85. hash_add(ldev->dmb_ht, &dmb_node->list, dmb_node->token);
  86. write_unlock_bh(&ldev->dmb_ht_lock);
  87. atomic_inc(&ldev->dmb_cnt);
  88. dmb->idx = dmb_node->sba_idx;
  89. dmb->dmb_tok = dmb_node->token;
  90. dmb->cpu_addr = dmb_node->cpu_addr;
  91. dmb->dma_addr = dmb_node->dma_addr;
  92. dmb->dmb_len = dmb_node->len;
  93. spin_lock_irqsave(&dibs->lock, flags);
  94. dibs->dmb_clientid_arr[sba_idx] = client->id;
  95. spin_unlock_irqrestore(&dibs->lock, flags);
  96. return 0;
  97. err_node:
  98. kfree(dmb_node);
  99. err_bit:
  100. clear_bit(sba_idx, ldev->sba_idx_mask);
  101. return rc;
  102. }
  103. static void __dibs_lo_unregister_dmb(struct dibs_lo_dev *ldev,
  104. struct dibs_lo_dmb_node *dmb_node)
  105. {
  106. /* remove dmb from hash table */
  107. write_lock_bh(&ldev->dmb_ht_lock);
  108. hash_del(&dmb_node->list);
  109. write_unlock_bh(&ldev->dmb_ht_lock);
  110. clear_bit(dmb_node->sba_idx, ldev->sba_idx_mask);
  111. folio_put(virt_to_folio(dmb_node->cpu_addr));
  112. kfree(dmb_node);
  113. if (atomic_dec_and_test(&ldev->dmb_cnt))
  114. wake_up(&ldev->ldev_release);
  115. }
  116. static int dibs_lo_unregister_dmb(struct dibs_dev *dibs, struct dibs_dmb *dmb)
  117. {
  118. struct dibs_lo_dmb_node *dmb_node = NULL, *tmp_node;
  119. struct dibs_lo_dev *ldev;
  120. unsigned long flags;
  121. ldev = dibs->drv_priv;
  122. /* find dmb from hash table */
  123. read_lock_bh(&ldev->dmb_ht_lock);
  124. hash_for_each_possible(ldev->dmb_ht, tmp_node, list, dmb->dmb_tok) {
  125. if (tmp_node->token == dmb->dmb_tok) {
  126. dmb_node = tmp_node;
  127. break;
  128. }
  129. }
  130. read_unlock_bh(&ldev->dmb_ht_lock);
  131. if (!dmb_node)
  132. return -EINVAL;
  133. if (refcount_dec_and_test(&dmb_node->refcnt)) {
  134. spin_lock_irqsave(&dibs->lock, flags);
  135. dibs->dmb_clientid_arr[dmb_node->sba_idx] = NO_DIBS_CLIENT;
  136. spin_unlock_irqrestore(&dibs->lock, flags);
  137. __dibs_lo_unregister_dmb(ldev, dmb_node);
  138. }
  139. return 0;
  140. }
  141. static int dibs_lo_support_dmb_nocopy(struct dibs_dev *dibs)
  142. {
  143. return DIBS_LO_SUPPORT_NOCOPY;
  144. }
  145. static int dibs_lo_attach_dmb(struct dibs_dev *dibs, struct dibs_dmb *dmb)
  146. {
  147. struct dibs_lo_dmb_node *dmb_node = NULL, *tmp_node;
  148. struct dibs_lo_dev *ldev;
  149. ldev = dibs->drv_priv;
  150. /* find dmb_node according to dmb->dmb_tok */
  151. read_lock_bh(&ldev->dmb_ht_lock);
  152. hash_for_each_possible(ldev->dmb_ht, tmp_node, list, dmb->dmb_tok) {
  153. if (tmp_node->token == dmb->dmb_tok) {
  154. dmb_node = tmp_node;
  155. break;
  156. }
  157. }
  158. if (!dmb_node) {
  159. read_unlock_bh(&ldev->dmb_ht_lock);
  160. return -EINVAL;
  161. }
  162. read_unlock_bh(&ldev->dmb_ht_lock);
  163. if (!refcount_inc_not_zero(&dmb_node->refcnt))
  164. /* the dmb is being unregistered, but has
  165. * not been removed from the hash table.
  166. */
  167. return -EINVAL;
  168. /* provide dmb information */
  169. dmb->idx = dmb_node->sba_idx;
  170. dmb->dmb_tok = dmb_node->token;
  171. dmb->cpu_addr = dmb_node->cpu_addr;
  172. dmb->dma_addr = dmb_node->dma_addr;
  173. dmb->dmb_len = dmb_node->len;
  174. return 0;
  175. }
  176. static int dibs_lo_detach_dmb(struct dibs_dev *dibs, u64 token)
  177. {
  178. struct dibs_lo_dmb_node *dmb_node = NULL, *tmp_node;
  179. struct dibs_lo_dev *ldev;
  180. ldev = dibs->drv_priv;
  181. /* find dmb_node according to dmb->dmb_tok */
  182. read_lock_bh(&ldev->dmb_ht_lock);
  183. hash_for_each_possible(ldev->dmb_ht, tmp_node, list, token) {
  184. if (tmp_node->token == token) {
  185. dmb_node = tmp_node;
  186. break;
  187. }
  188. }
  189. if (!dmb_node) {
  190. read_unlock_bh(&ldev->dmb_ht_lock);
  191. return -EINVAL;
  192. }
  193. read_unlock_bh(&ldev->dmb_ht_lock);
  194. if (refcount_dec_and_test(&dmb_node->refcnt))
  195. __dibs_lo_unregister_dmb(ldev, dmb_node);
  196. return 0;
  197. }
  198. static int dibs_lo_move_data(struct dibs_dev *dibs, u64 dmb_tok,
  199. unsigned int idx, bool sf, unsigned int offset,
  200. void *data, unsigned int size)
  201. {
  202. struct dibs_lo_dmb_node *rmb_node = NULL, *tmp_node;
  203. struct dibs_lo_dev *ldev;
  204. u16 s_mask;
  205. u8 client_id;
  206. u32 sba_idx;
  207. ldev = dibs->drv_priv;
  208. read_lock_bh(&ldev->dmb_ht_lock);
  209. hash_for_each_possible(ldev->dmb_ht, tmp_node, list, dmb_tok) {
  210. if (tmp_node->token == dmb_tok) {
  211. rmb_node = tmp_node;
  212. break;
  213. }
  214. }
  215. if (!rmb_node) {
  216. read_unlock_bh(&ldev->dmb_ht_lock);
  217. return -EINVAL;
  218. }
  219. memcpy((char *)rmb_node->cpu_addr + offset, data, size);
  220. sba_idx = rmb_node->sba_idx;
  221. read_unlock_bh(&ldev->dmb_ht_lock);
  222. if (!sf)
  223. return 0;
  224. spin_lock(&dibs->lock);
  225. client_id = dibs->dmb_clientid_arr[sba_idx];
  226. s_mask = ror16(0x1000, idx);
  227. if (likely(client_id != NO_DIBS_CLIENT && dibs->subs[client_id]))
  228. dibs->subs[client_id]->ops->handle_irq(dibs, sba_idx, s_mask);
  229. spin_unlock(&dibs->lock);
  230. return 0;
  231. }
  232. static const struct dibs_dev_ops dibs_lo_ops = {
  233. .get_fabric_id = dibs_lo_get_fabric_id,
  234. .query_remote_gid = dibs_lo_query_rgid,
  235. .max_dmbs = dibs_lo_max_dmbs,
  236. .register_dmb = dibs_lo_register_dmb,
  237. .unregister_dmb = dibs_lo_unregister_dmb,
  238. .move_data = dibs_lo_move_data,
  239. .support_mmapped_rdmb = dibs_lo_support_dmb_nocopy,
  240. .attach_dmb = dibs_lo_attach_dmb,
  241. .detach_dmb = dibs_lo_detach_dmb,
  242. };
  243. static void dibs_lo_dev_init(struct dibs_lo_dev *ldev)
  244. {
  245. rwlock_init(&ldev->dmb_ht_lock);
  246. hash_init(ldev->dmb_ht);
  247. atomic_set(&ldev->dmb_cnt, 0);
  248. init_waitqueue_head(&ldev->ldev_release);
  249. }
  250. static void dibs_lo_dev_exit(struct dibs_lo_dev *ldev)
  251. {
  252. if (atomic_read(&ldev->dmb_cnt))
  253. wait_event(ldev->ldev_release, !atomic_read(&ldev->dmb_cnt));
  254. }
  255. static int dibs_lo_dev_probe(void)
  256. {
  257. struct dibs_lo_dev *ldev;
  258. struct dibs_dev *dibs;
  259. int ret;
  260. ldev = kzalloc_obj(*ldev);
  261. if (!ldev)
  262. return -ENOMEM;
  263. dibs = dibs_dev_alloc();
  264. if (!dibs) {
  265. kfree(ldev);
  266. return -ENOMEM;
  267. }
  268. ldev->dibs = dibs;
  269. dibs->drv_priv = ldev;
  270. dibs_lo_dev_init(ldev);
  271. uuid_gen(&dibs->gid);
  272. dibs->ops = &dibs_lo_ops;
  273. dibs->dev.parent = NULL;
  274. dev_set_name(&dibs->dev, "%s", dibs_lo_dev_name);
  275. ret = dibs_dev_add(dibs);
  276. if (ret)
  277. goto err_reg;
  278. lo_dev = ldev;
  279. return 0;
  280. err_reg:
  281. kfree(dibs->dmb_clientid_arr);
  282. /* pairs with dibs_dev_alloc() */
  283. put_device(&dibs->dev);
  284. kfree(ldev);
  285. return ret;
  286. }
  287. static void dibs_lo_dev_remove(void)
  288. {
  289. if (!lo_dev)
  290. return;
  291. dibs_dev_del(lo_dev->dibs);
  292. dibs_lo_dev_exit(lo_dev);
  293. /* pairs with dibs_dev_alloc() */
  294. put_device(&lo_dev->dibs->dev);
  295. kfree(lo_dev);
  296. lo_dev = NULL;
  297. }
  298. int dibs_loopback_init(void)
  299. {
  300. return dibs_lo_dev_probe();
  301. }
  302. void dibs_loopback_exit(void)
  303. {
  304. dibs_lo_dev_remove();
  305. }