1 // SPDX-License-Identifier: GPL-2.0-only
2
3 #include <linux/ethtool.h>
4 #include <linux/net_namespace.h>
5 #include <linux/skbuff.h>
6 #include <linux/xarray.h>
7 #include <net/genetlink.h>
8 #include <net/psp.h>
9 #include <net/sock.h>
10
11 #include "psp-nl-gen.h"
12 #include "psp.h"
13
14 /* Netlink helpers */
15
psp_nl_reply_new(struct genl_info * info)16 static struct sk_buff *psp_nl_reply_new(struct genl_info *info)
17 {
18 struct sk_buff *rsp;
19 void *hdr;
20
21 rsp = genlmsg_new(GENLMSG_DEFAULT_SIZE, GFP_KERNEL);
22 if (!rsp)
23 return NULL;
24
25 hdr = genlmsg_iput(rsp, info);
26 if (!hdr) {
27 nlmsg_free(rsp);
28 return NULL;
29 }
30
31 return rsp;
32 }
33
psp_nl_reply_send(struct sk_buff * rsp,struct genl_info * info)34 static int psp_nl_reply_send(struct sk_buff *rsp, struct genl_info *info)
35 {
36 /* Note that this *only* works with a single message per skb! */
37 nlmsg_end(rsp, (struct nlmsghdr *)rsp->data);
38
39 return genlmsg_reply(rsp, info);
40 }
41
42 /**
43 * psp_nl_multicast_per_ns() - multicast a notification to each unique netns
44 * @psd: PSP device (must be locked)
45 * @group: multicast group
46 * @build_ntf: callback to build an skb for a given netns, or NULL on failure
47 * @ctx: opaque context passed to @build_ntf
48 *
49 * Iterates all unique network namespaces from the associated device list
50 * plus the main device's netns. For each unique netns, calls @build_ntf
51 * to construct a notification skb and multicasts it.
52 */
53 static void
psp_nl_multicast_per_ns(struct psp_dev * psd,unsigned int group,struct sk_buff * (* build_ntf)(struct psp_dev *,struct net *,void *),void * ctx)54 psp_nl_multicast_per_ns(struct psp_dev *psd, unsigned int group,
55 struct sk_buff *(*build_ntf)(struct psp_dev *,
56 struct net *,
57 void *),
58 void *ctx)
59 {
60 struct psp_assoc_dev *entry;
61 struct xarray sent_nets;
62 struct net *main_net;
63 struct sk_buff *ntf;
64
65 /* device may be changing netns in parallel */
66 rcu_read_lock();
67 main_net = maybe_get_net(dev_net_rcu(psd->main_netdev));
68 rcu_read_unlock();
69
70 if (!main_net)
71 return;
72
73 xa_init(&sent_nets);
74
75 list_for_each_entry(entry, &psd->assoc_dev_list, dev_list) {
76 struct net *assoc_net = dev_net(entry->assoc_dev);
77 int ret;
78
79 if (net_eq(assoc_net, main_net))
80 continue;
81
82 ret = xa_insert(&sent_nets, (unsigned long)assoc_net, assoc_net,
83 GFP_KERNEL);
84 if (ret == -EBUSY)
85 continue;
86
87 ntf = build_ntf(psd, assoc_net, ctx);
88 if (!ntf)
89 continue;
90
91 genlmsg_multicast_netns(&psp_nl_family, assoc_net, ntf, 0,
92 group, GFP_KERNEL);
93 }
94 xa_destroy(&sent_nets);
95
96 /* Send to main device netns */
97 ntf = build_ntf(psd, main_net, ctx);
98 if (ntf)
99 genlmsg_multicast_netns(&psp_nl_family, main_net, ntf, 0, group,
100 GFP_KERNEL);
101 put_net(main_net);
102 }
103
psp_nl_clone_ntf(struct psp_dev * psd,struct net * net,void * ctx)104 static struct sk_buff *psp_nl_clone_ntf(struct psp_dev *psd, struct net *net,
105 void *ctx)
106 {
107 return skb_clone(ctx, GFP_KERNEL);
108 }
109
psp_nl_multicast_all_ns(struct psp_dev * psd,struct sk_buff * ntf,unsigned int group)110 static void psp_nl_multicast_all_ns(struct psp_dev *psd, struct sk_buff *ntf,
111 unsigned int group)
112 {
113 psp_nl_multicast_per_ns(psd, group, psp_nl_clone_ntf, ntf);
114 nlmsg_consume(ntf);
115 }
116
117 /* Device stuff */
118
119 static struct psp_dev *
psp_device_get_and_lock(struct net * net,struct nlattr * dev_id,bool admin)120 psp_device_get_and_lock(struct net *net, struct nlattr *dev_id,
121 bool admin)
122 {
123 struct psp_dev *psd;
124 int err;
125
126 mutex_lock(&psp_devs_lock);
127 psd = xa_load(&psp_devs, nla_get_u32(dev_id));
128 if (!psd) {
129 mutex_unlock(&psp_devs_lock);
130 return ERR_PTR(-ENODEV);
131 }
132
133 mutex_lock(&psd->lock);
134 mutex_unlock(&psp_devs_lock);
135
136 err = psp_dev_check_access(psd, net, admin);
137 if (err) {
138 mutex_unlock(&psd->lock);
139 return ERR_PTR(err);
140 }
141
142 return psd;
143 }
144
__psp_device_get_locked(const struct genl_split_ops * ops,struct sk_buff * skb,struct genl_info * info,bool admin)145 static int __psp_device_get_locked(const struct genl_split_ops *ops,
146 struct sk_buff *skb, struct genl_info *info,
147 bool admin)
148 {
149 if (GENL_REQ_ATTR_CHECK(info, PSP_A_DEV_ID))
150 return -EINVAL;
151
152 info->user_ptr[0] = psp_device_get_and_lock(genl_info_net(info),
153 info->attrs[PSP_A_DEV_ID],
154 admin);
155 return PTR_ERR_OR_ZERO(info->user_ptr[0]);
156 }
157
psp_device_get_locked_admin(const struct genl_split_ops * ops,struct sk_buff * skb,struct genl_info * info)158 int psp_device_get_locked_admin(const struct genl_split_ops *ops,
159 struct sk_buff *skb, struct genl_info *info)
160 {
161 return __psp_device_get_locked(ops, skb, info, true);
162 }
163
psp_device_get_locked(const struct genl_split_ops * ops,struct sk_buff * skb,struct genl_info * info)164 int psp_device_get_locked(const struct genl_split_ops *ops,
165 struct sk_buff *skb, struct genl_info *info)
166 {
167 return __psp_device_get_locked(ops, skb, info, false);
168 }
169
170 /*
171 * Non-admin version of psp_device_get_locked() + psp_attach_netdev_notifier()
172 * only used for dev-assoc.
173 */
psp_device_get_locked_dev_assoc(const struct genl_split_ops * ops,struct sk_buff * skb,struct genl_info * info)174 int psp_device_get_locked_dev_assoc(const struct genl_split_ops *ops,
175 struct sk_buff *skb, struct genl_info *info)
176 {
177 int err;
178
179 err = psp_attach_netdev_notifier();
180 if (err)
181 return err;
182
183 return __psp_device_get_locked(ops, skb, info, false);
184 }
185
psp_nl_resolve_assoc_dev_ns(struct psp_dev * psd,struct genl_info * info)186 static struct net *psp_nl_resolve_assoc_dev_ns(struct psp_dev *psd,
187 struct genl_info *info)
188 {
189 struct net *net;
190 int nsid;
191
192 if (GENL_REQ_ATTR_CHECK(info, PSP_A_DEV_IFINDEX))
193 return ERR_PTR(-EINVAL);
194
195 if (info->attrs[PSP_A_DEV_NSID]) {
196 /* Only callers in the main netns may specify nsid */
197 if (dev_net(psd->main_netdev) != genl_info_net(info)) {
198 NL_SET_BAD_ATTR(info->extack,
199 info->attrs[PSP_A_DEV_NSID]);
200 return ERR_PTR(-EPERM);
201 }
202
203 nsid = nla_get_s32(info->attrs[PSP_A_DEV_NSID]);
204
205 net = get_net_ns_by_id(genl_info_net(info), nsid);
206 if (!net) {
207 NL_SET_BAD_ATTR(info->extack,
208 info->attrs[PSP_A_DEV_NSID]);
209 return ERR_PTR(-EINVAL);
210 }
211 } else {
212 net = get_net(genl_info_net(info));
213 }
214
215 return net;
216 }
217
218 void
psp_device_unlock(const struct genl_split_ops * ops,struct sk_buff * skb,struct genl_info * info)219 psp_device_unlock(const struct genl_split_ops *ops, struct sk_buff *skb,
220 struct genl_info *info)
221 {
222 struct socket *socket = info->user_ptr[1];
223 struct psp_dev *psd = info->user_ptr[0];
224
225 mutex_unlock(&psd->lock);
226 if (socket)
227 sockfd_put(socket);
228 }
229
psp_has_assoc_dev_in_ns(struct psp_dev * psd,struct net * net)230 bool psp_has_assoc_dev_in_ns(struct psp_dev *psd, struct net *net)
231 {
232 struct psp_assoc_dev *entry;
233
234 list_for_each_entry(entry, &psd->assoc_dev_list, dev_list) {
235 if (dev_net(entry->assoc_dev) == net)
236 return true;
237 }
238
239 return false;
240 }
241
psp_nl_fill_assoc_dev_list(struct psp_dev * psd,struct sk_buff * rsp,struct net * cur_net,struct net * filter_net)242 static int psp_nl_fill_assoc_dev_list(struct psp_dev *psd, struct sk_buff *rsp,
243 struct net *cur_net,
244 struct net *filter_net)
245 {
246 struct psp_assoc_dev *entry;
247 struct net *dev_net_ns;
248 struct nlattr *nest;
249 int nsid;
250
251 list_for_each_entry(entry, &psd->assoc_dev_list, dev_list) {
252 dev_net_ns = dev_net(entry->assoc_dev);
253
254 if (filter_net && dev_net_ns != filter_net)
255 continue;
256
257 /* When filtering by namespace, all devices are in the caller's
258 * namespace so nsid is always NETNSA_NSID_NOT_ASSIGNED (-1).
259 * Otherwise, calculate the nsid relative to cur_net.
260 */
261 nsid = filter_net ? NETNSA_NSID_NOT_ASSIGNED :
262 peernet2id_alloc(cur_net, dev_net_ns,
263 GFP_KERNEL);
264
265 nest = nla_nest_start(rsp, PSP_A_DEV_ASSOC_LIST);
266 if (!nest)
267 return -EMSGSIZE;
268
269 if (nla_put_u32(rsp, PSP_A_ASSOC_DEV_INFO_IFINDEX,
270 entry->assoc_dev->ifindex) ||
271 nla_put_s32(rsp, PSP_A_ASSOC_DEV_INFO_NSID, nsid)) {
272 nla_nest_cancel(rsp, nest);
273 return -EMSGSIZE;
274 }
275
276 nla_nest_end(rsp, nest);
277 }
278
279 return 0;
280 }
281
282 static int
psp_nl_dev_fill(struct psp_dev * psd,struct sk_buff * rsp,const struct genl_info * info)283 psp_nl_dev_fill(struct psp_dev *psd, struct sk_buff *rsp,
284 const struct genl_info *info)
285 {
286 struct net *cur_net;
287 void *hdr;
288 int err;
289
290 cur_net = genl_info_net(info);
291
292 hdr = genlmsg_iput(rsp, info);
293 if (!hdr)
294 return -EMSGSIZE;
295
296 if (nla_put_u32(rsp, PSP_A_DEV_ID, psd->id) ||
297 nla_put_u32(rsp, PSP_A_DEV_IFINDEX, psd->main_netdev->ifindex) ||
298 nla_put_u32(rsp, PSP_A_DEV_PSP_VERSIONS_CAP, psd->caps->versions) ||
299 nla_put_u32(rsp, PSP_A_DEV_PSP_VERSIONS_ENA, psd->config.versions))
300 goto err_cancel_msg;
301
302 if (cur_net == dev_net(psd->main_netdev)) {
303 /* Primary device - dump assoc list */
304 err = psp_nl_fill_assoc_dev_list(psd, rsp, cur_net, NULL);
305 if (err)
306 goto err_cancel_msg;
307 } else {
308 /* In netns: set by-association flag and dump filtered
309 * assoc list containing only devices in cur_net
310 */
311 if (nla_put_flag(rsp, PSP_A_DEV_BY_ASSOCIATION))
312 goto err_cancel_msg;
313 err = psp_nl_fill_assoc_dev_list(psd, rsp, cur_net, cur_net);
314 if (err)
315 goto err_cancel_msg;
316 }
317
318 genlmsg_end(rsp, hdr);
319 return 0;
320
321 err_cancel_msg:
322 genlmsg_cancel(rsp, hdr);
323 return -EMSGSIZE;
324 }
325
psp_nl_build_dev_ntf(struct psp_dev * psd,struct net * net,void * ctx)326 static struct sk_buff *psp_nl_build_dev_ntf(struct psp_dev *psd,
327 struct net *net, void *ctx)
328 {
329 u32 cmd = *(u32 *)ctx;
330 struct genl_info info;
331 struct sk_buff *ntf;
332
333 if (!genl_has_listeners(&psp_nl_family, net, PSP_NLGRP_MGMT))
334 return NULL;
335
336 ntf = genlmsg_new(GENLMSG_DEFAULT_SIZE, GFP_KERNEL);
337 if (!ntf)
338 return NULL;
339
340 genl_info_init_ntf(&info, &psp_nl_family, cmd);
341 genl_info_net_set(&info, net);
342 if (psp_nl_dev_fill(psd, ntf, &info)) {
343 nlmsg_free(ntf);
344 return NULL;
345 }
346
347 return ntf;
348 }
349
psp_nl_notify_dev(struct psp_dev * psd,u32 cmd)350 void psp_nl_notify_dev(struct psp_dev *psd, u32 cmd)
351 {
352 psp_nl_multicast_per_ns(psd, PSP_NLGRP_MGMT,
353 psp_nl_build_dev_ntf, &cmd);
354 }
355
psp_nl_dev_get_doit(struct sk_buff * req,struct genl_info * info)356 int psp_nl_dev_get_doit(struct sk_buff *req, struct genl_info *info)
357 {
358 struct psp_dev *psd = info->user_ptr[0];
359 struct sk_buff *rsp;
360 int err;
361
362 rsp = genlmsg_new(GENLMSG_DEFAULT_SIZE, GFP_KERNEL);
363 if (!rsp)
364 return -ENOMEM;
365
366 err = psp_nl_dev_fill(psd, rsp, info);
367 if (err)
368 goto err_free_msg;
369
370 return genlmsg_reply(rsp, info);
371
372 err_free_msg:
373 nlmsg_free(rsp);
374 return err;
375 }
376
377 static int
psp_nl_dev_get_dumpit_one(struct sk_buff * rsp,struct netlink_callback * cb,struct psp_dev * psd)378 psp_nl_dev_get_dumpit_one(struct sk_buff *rsp, struct netlink_callback *cb,
379 struct psp_dev *psd)
380 {
381 if (psp_dev_check_access(psd, sock_net(rsp->sk), false))
382 return 0;
383
384 return psp_nl_dev_fill(psd, rsp, genl_info_dump(cb));
385 }
386
psp_nl_dev_get_dumpit(struct sk_buff * rsp,struct netlink_callback * cb)387 int psp_nl_dev_get_dumpit(struct sk_buff *rsp, struct netlink_callback *cb)
388 {
389 struct psp_dev *psd;
390 int err = 0;
391
392 mutex_lock(&psp_devs_lock);
393 xa_for_each_start(&psp_devs, cb->args[0], psd, cb->args[0]) {
394 mutex_lock(&psd->lock);
395 err = psp_nl_dev_get_dumpit_one(rsp, cb, psd);
396 mutex_unlock(&psd->lock);
397 if (err)
398 break;
399 }
400 mutex_unlock(&psp_devs_lock);
401
402 return err;
403 }
404
psp_nl_dev_set_doit(struct sk_buff * skb,struct genl_info * info)405 int psp_nl_dev_set_doit(struct sk_buff *skb, struct genl_info *info)
406 {
407 struct psp_dev *psd = info->user_ptr[0];
408 struct psp_dev_config new_config;
409 struct sk_buff *rsp;
410 int err;
411
412 memcpy(&new_config, &psd->config, sizeof(new_config));
413
414 if (info->attrs[PSP_A_DEV_PSP_VERSIONS_ENA]) {
415 new_config.versions =
416 nla_get_u32(info->attrs[PSP_A_DEV_PSP_VERSIONS_ENA]);
417 if (new_config.versions & ~psd->caps->versions) {
418 NL_SET_ERR_MSG(info->extack, "Requested PSP versions not supported by the device");
419 return -EINVAL;
420 }
421 } else {
422 NL_SET_ERR_MSG(info->extack, "No settings present");
423 return -EINVAL;
424 }
425
426 rsp = psp_nl_reply_new(info);
427 if (!rsp)
428 return -ENOMEM;
429
430 if (memcmp(&new_config, &psd->config, sizeof(new_config))) {
431 err = psd->ops->set_config(psd, &new_config, info->extack);
432 if (err)
433 goto err_free_rsp;
434
435 memcpy(&psd->config, &new_config, sizeof(new_config));
436 }
437
438 psp_nl_notify_dev(psd, PSP_CMD_DEV_CHANGE_NTF);
439
440 return psp_nl_reply_send(rsp, info);
441
442 err_free_rsp:
443 nlmsg_free(rsp);
444 return err;
445 }
446
psp_nl_key_rotate_doit(struct sk_buff * skb,struct genl_info * info)447 int psp_nl_key_rotate_doit(struct sk_buff *skb, struct genl_info *info)
448 {
449 struct psp_dev *psd = info->user_ptr[0];
450 struct genl_info ntf_info;
451 struct sk_buff *ntf, *rsp;
452 u8 prev_gen;
453 int err;
454
455 rsp = psp_nl_reply_new(info);
456 if (!rsp)
457 return -ENOMEM;
458
459 genl_info_init_ntf(&ntf_info, &psp_nl_family, PSP_CMD_KEY_ROTATE_NTF);
460 ntf = psp_nl_reply_new(&ntf_info);
461 if (!ntf) {
462 err = -ENOMEM;
463 goto err_free_rsp;
464 }
465
466 if (nla_put_u32(rsp, PSP_A_DEV_ID, psd->id) ||
467 nla_put_u32(ntf, PSP_A_DEV_ID, psd->id)) {
468 err = -EMSGSIZE;
469 goto err_free_ntf;
470 }
471
472 /* suggest the next gen number, driver can override */
473 prev_gen = psd->generation;
474 psd->generation = (prev_gen + 1) & PSP_GEN_VALID_MASK;
475
476 err = psd->ops->key_rotate(psd, info->extack);
477 if (err)
478 goto err_free_ntf;
479
480 WARN_ON_ONCE((psd->generation && psd->generation == prev_gen) ||
481 psd->generation & ~PSP_GEN_VALID_MASK);
482
483 psp_assocs_key_rotated(psd);
484 psd->stats.rotations++;
485
486 nlmsg_end(ntf, (struct nlmsghdr *)ntf->data);
487
488 psp_nl_multicast_all_ns(psd, ntf, PSP_NLGRP_USE);
489
490 return psp_nl_reply_send(rsp, info);
491
492 err_free_ntf:
493 nlmsg_free(ntf);
494 err_free_rsp:
495 nlmsg_free(rsp);
496 return err;
497 }
498
psp_nl_dev_assoc_doit(struct sk_buff * skb,struct genl_info * info)499 int psp_nl_dev_assoc_doit(struct sk_buff *skb, struct genl_info *info)
500 {
501 struct psp_dev *psd = info->user_ptr[0];
502 struct psp_assoc_dev *psp_assoc_dev;
503 struct net_device *assoc_dev;
504 struct sk_buff *rsp;
505 u32 assoc_ifindex;
506 struct net *net;
507 int err;
508
509 if (psd->assoc_dev_cnt >= PSP_ASSOC_DEV_MAX) {
510 NL_SET_ERR_MSG(info->extack,
511 "Maximum number of associated devices reached");
512 return -ENOSPC;
513 }
514
515 net = psp_nl_resolve_assoc_dev_ns(psd, info);
516 if (IS_ERR(net))
517 return PTR_ERR(net);
518
519 psp_assoc_dev = kzalloc_obj(*psp_assoc_dev);
520 if (!psp_assoc_dev) {
521 err = -ENOMEM;
522 goto err_put_net;
523 }
524
525 assoc_ifindex = nla_get_u32(info->attrs[PSP_A_DEV_IFINDEX]);
526 assoc_dev = netdev_get_by_index(net, assoc_ifindex,
527 &psp_assoc_dev->dev_tracker,
528 GFP_KERNEL);
529 if (!assoc_dev) {
530 NL_SET_BAD_ATTR(info->extack, info->attrs[PSP_A_DEV_IFINDEX]);
531 err = -ENODEV;
532 goto err_free_assoc;
533 }
534
535 /* Check if device is already associated with a PSP device */
536 if (unrcu_pointer(cmpxchg(&assoc_dev->psp_dev, NULL,
537 RCU_INITIALIZER(psd)))) {
538 NL_SET_ERR_MSG(info->extack,
539 "Device already associated with a PSP device");
540 err = -EBUSY;
541 goto err_put_dev;
542 }
543
544 psp_assoc_dev->assoc_dev = assoc_dev;
545
546 /* Check for race with NETDEV_UNREGISTER. The cmpxchg above is a
547 * full barrier, and the unregister path has synchronize_net()
548 * between setting NETREG_UNREGISTERING and reading psp_dev in the
549 * notifier. So at least one side would do the clean-up if we are in
550 * the middle of unregitering assoc_dev.
551 * And the clean-up is serialized by psd->lock.
552 */
553 if (READ_ONCE(assoc_dev->reg_state) != NETREG_REGISTERED) {
554 err = -ENODEV;
555 goto err_clean_ptr;
556 }
557
558 rsp = psp_nl_reply_new(info);
559 if (!rsp) {
560 err = -ENOMEM;
561 goto err_clean_ptr;
562 }
563
564 list_add_tail(&psp_assoc_dev->dev_list, &psd->assoc_dev_list);
565 psd->assoc_dev_cnt++;
566
567 put_net(net);
568
569 psp_nl_notify_dev(psd, PSP_CMD_DEV_CHANGE_NTF);
570
571 return psp_nl_reply_send(rsp, info);
572
573 err_clean_ptr:
574 rcu_assign_pointer(assoc_dev->psp_dev, NULL);
575 err_put_dev:
576 netdev_put(assoc_dev, &psp_assoc_dev->dev_tracker);
577 err_free_assoc:
578 kfree(psp_assoc_dev);
579 err_put_net:
580 put_net(net);
581
582 return err;
583 }
584
psp_nl_dev_disassoc_doit(struct sk_buff * skb,struct genl_info * info)585 int psp_nl_dev_disassoc_doit(struct sk_buff *skb, struct genl_info *info)
586 {
587 struct psp_assoc_dev *entry, *found = NULL;
588 struct psp_dev *psd = info->user_ptr[0];
589 struct sk_buff *rsp;
590 u32 assoc_ifindex;
591 struct net *net;
592
593 net = psp_nl_resolve_assoc_dev_ns(psd, info);
594 if (IS_ERR(net))
595 return PTR_ERR(net);
596
597 assoc_ifindex = nla_get_u32(info->attrs[PSP_A_DEV_IFINDEX]);
598
599 /* Search the association list by ifindex and netns */
600 list_for_each_entry(entry, &psd->assoc_dev_list, dev_list) {
601 if (entry->assoc_dev->ifindex == assoc_ifindex &&
602 dev_net(entry->assoc_dev) == net) {
603 found = entry;
604 break;
605 }
606 }
607
608 if (!found) {
609 put_net(net);
610 NL_SET_BAD_ATTR(info->extack, info->attrs[PSP_A_DEV_IFINDEX]);
611 return -ENODEV;
612 }
613
614 rsp = psp_nl_reply_new(info);
615 if (!rsp) {
616 put_net(net);
617 return -ENOMEM;
618 }
619
620 put_net(net);
621
622 /* Notify before removal so listeners in the disassociated namespace
623 * still receive the notification.
624 */
625 psp_nl_notify_dev(psd, PSP_CMD_DEV_CHANGE_NTF);
626
627 /* Remove from the association list */
628 list_del(&found->dev_list);
629 psd->assoc_dev_cnt--;
630 rcu_assign_pointer(found->assoc_dev->psp_dev, NULL);
631 netdev_put(found->assoc_dev, &found->dev_tracker);
632 kfree(found);
633
634 return psp_nl_reply_send(rsp, info);
635 }
636
637 /* Key etc. */
638
psp_assoc_device_get_locked(const struct genl_split_ops * ops,struct sk_buff * skb,struct genl_info * info)639 int psp_assoc_device_get_locked(const struct genl_split_ops *ops,
640 struct sk_buff *skb, struct genl_info *info)
641 {
642 struct socket *socket;
643 struct psp_dev *psd;
644 struct nlattr *id;
645 int fd, err;
646
647 if (GENL_REQ_ATTR_CHECK(info, PSP_A_ASSOC_SOCK_FD))
648 return -EINVAL;
649
650 fd = nla_get_u32(info->attrs[PSP_A_ASSOC_SOCK_FD]);
651 socket = sockfd_lookup(fd, &err);
652 if (!socket)
653 return err;
654
655 if (!sk_is_tcp(socket->sk)) {
656 NL_SET_ERR_MSG_ATTR(info->extack,
657 info->attrs[PSP_A_ASSOC_SOCK_FD],
658 "Unsupported socket family and type");
659 err = -EOPNOTSUPP;
660 goto err_sock_put;
661 }
662
663 psd = psp_dev_get_for_sock(socket->sk);
664 if (psd) {
665 /* Extra care needed here, psp_dev_get_for_sock() only gives
666 * us access to struct psp_dev's memory, which is quite weak.
667 */
668 mutex_lock(&psd->lock);
669 if (!psp_dev_is_registered(psd) ||
670 psp_dev_check_access(psd, genl_info_net(info), false)) {
671 mutex_unlock(&psd->lock);
672 psp_dev_put(psd);
673 psd = NULL;
674 }
675 }
676
677 if (!psd && GENL_REQ_ATTR_CHECK(info, PSP_A_ASSOC_DEV_ID)) {
678 err = -EINVAL;
679 goto err_sock_put;
680 }
681
682 id = info->attrs[PSP_A_ASSOC_DEV_ID];
683 if (psd) {
684 if (id && psd->id != nla_get_u32(id)) {
685 mutex_unlock(&psd->lock);
686 NL_SET_ERR_MSG_ATTR(info->extack, id,
687 "Device id vs socket mismatch");
688 err = -EINVAL;
689 goto err_psd_put;
690 }
691
692 psp_dev_put(psd);
693 } else {
694 psd = psp_device_get_and_lock(genl_info_net(info), id, false);
695 if (IS_ERR(psd)) {
696 err = PTR_ERR(psd);
697 goto err_sock_put;
698 }
699 }
700
701 info->user_ptr[0] = psd;
702 info->user_ptr[1] = socket;
703
704 return 0;
705
706 err_psd_put:
707 psp_dev_put(psd);
708 err_sock_put:
709 sockfd_put(socket);
710 return err;
711 }
712
713 static int
psp_nl_parse_key(struct genl_info * info,u32 attr,struct psp_key_parsed * key,unsigned int key_sz)714 psp_nl_parse_key(struct genl_info *info, u32 attr, struct psp_key_parsed *key,
715 unsigned int key_sz)
716 {
717 struct nlattr *nest = info->attrs[attr];
718 struct nlattr *tb[PSP_A_KEYS_SPI + 1];
719 u32 spi;
720 int err;
721
722 err = nla_parse_nested(tb, ARRAY_SIZE(tb) - 1, nest,
723 psp_keys_nl_policy, info->extack);
724 if (err)
725 return err;
726
727 if (NL_REQ_ATTR_CHECK(info->extack, nest, tb, PSP_A_KEYS_KEY) ||
728 NL_REQ_ATTR_CHECK(info->extack, nest, tb, PSP_A_KEYS_SPI))
729 return -EINVAL;
730
731 if (nla_len(tb[PSP_A_KEYS_KEY]) != key_sz) {
732 NL_SET_ERR_MSG_ATTR(info->extack, tb[PSP_A_KEYS_KEY],
733 "incorrect key length");
734 return -EINVAL;
735 }
736
737 spi = nla_get_u32(tb[PSP_A_KEYS_SPI]);
738 if (!(spi & PSP_SPI_KEY_ID)) {
739 NL_SET_ERR_MSG_ATTR(info->extack, tb[PSP_A_KEYS_KEY],
740 "invalid SPI: lower 31b must be non-zero");
741 return -EINVAL;
742 }
743
744 key->spi = cpu_to_be32(spi);
745 memcpy(key->key, nla_data(tb[PSP_A_KEYS_KEY]), key_sz);
746
747 return 0;
748 }
749
750 static int
psp_nl_put_key(struct sk_buff * skb,u32 attr,u32 version,struct psp_key_parsed * key)751 psp_nl_put_key(struct sk_buff *skb, u32 attr, u32 version,
752 struct psp_key_parsed *key)
753 {
754 int key_sz = psp_key_size(version);
755 void *nest;
756
757 nest = nla_nest_start(skb, attr);
758
759 if (nla_put_u32(skb, PSP_A_KEYS_SPI, be32_to_cpu(key->spi)) ||
760 nla_put(skb, PSP_A_KEYS_KEY, key_sz, key->key)) {
761 nla_nest_cancel(skb, nest);
762 return -EMSGSIZE;
763 }
764
765 nla_nest_end(skb, nest);
766
767 return 0;
768 }
769
psp_nl_rx_assoc_doit(struct sk_buff * skb,struct genl_info * info)770 int psp_nl_rx_assoc_doit(struct sk_buff *skb, struct genl_info *info)
771 {
772 struct socket *socket = info->user_ptr[1];
773 struct psp_dev *psd = info->user_ptr[0];
774 struct psp_key_parsed key;
775 struct psp_assoc *pas;
776 struct sk_buff *rsp;
777 u32 version;
778 int err;
779
780 if (GENL_REQ_ATTR_CHECK(info, PSP_A_ASSOC_VERSION))
781 return -EINVAL;
782
783 version = nla_get_u32(info->attrs[PSP_A_ASSOC_VERSION]);
784 if (!(psd->caps->versions & (1 << version))) {
785 NL_SET_BAD_ATTR(info->extack, info->attrs[PSP_A_ASSOC_VERSION]);
786 return -EOPNOTSUPP;
787 }
788
789 rsp = psp_nl_reply_new(info);
790 if (!rsp)
791 return -ENOMEM;
792
793 pas = psp_assoc_create(psd);
794 if (!pas) {
795 err = -ENOMEM;
796 goto err_free_rsp;
797 }
798 pas->version = version;
799
800 err = psd->ops->rx_spi_alloc(psd, version, &key, info->extack);
801 if (err)
802 goto err_free_pas;
803
804 if (nla_put_u32(rsp, PSP_A_ASSOC_DEV_ID, psd->id) ||
805 psp_nl_put_key(rsp, PSP_A_ASSOC_RX_KEY, version, &key)) {
806 err = -EMSGSIZE;
807 goto err_free_pas;
808 }
809
810 err = psp_sock_assoc_set_rx(socket->sk, pas, &key, info->extack);
811 if (err) {
812 NL_SET_BAD_ATTR(info->extack, info->attrs[PSP_A_ASSOC_SOCK_FD]);
813 goto err_free_pas;
814 }
815 psp_assoc_put(pas);
816
817 return psp_nl_reply_send(rsp, info);
818
819 err_free_pas:
820 psp_assoc_put(pas);
821 err_free_rsp:
822 nlmsg_free(rsp);
823 return err;
824 }
825
psp_nl_tx_assoc_doit(struct sk_buff * skb,struct genl_info * info)826 int psp_nl_tx_assoc_doit(struct sk_buff *skb, struct genl_info *info)
827 {
828 struct socket *socket = info->user_ptr[1];
829 struct psp_dev *psd = info->user_ptr[0];
830 struct psp_key_parsed key;
831 struct sk_buff *rsp;
832 unsigned int key_sz;
833 u32 version;
834 int err;
835
836 if (GENL_REQ_ATTR_CHECK(info, PSP_A_ASSOC_VERSION) ||
837 GENL_REQ_ATTR_CHECK(info, PSP_A_ASSOC_TX_KEY))
838 return -EINVAL;
839
840 version = nla_get_u32(info->attrs[PSP_A_ASSOC_VERSION]);
841 if (!(psd->caps->versions & (1 << version))) {
842 NL_SET_BAD_ATTR(info->extack, info->attrs[PSP_A_ASSOC_VERSION]);
843 return -EOPNOTSUPP;
844 }
845
846 key_sz = psp_key_size(version);
847 if (!key_sz)
848 return -EINVAL;
849
850 err = psp_nl_parse_key(info, PSP_A_ASSOC_TX_KEY, &key, key_sz);
851 if (err < 0)
852 return err;
853
854 rsp = psp_nl_reply_new(info);
855 if (!rsp)
856 return -ENOMEM;
857
858 err = psp_sock_assoc_set_tx(socket->sk, psd, version, &key,
859 info->extack);
860 if (err)
861 goto err_free_msg;
862
863 return psp_nl_reply_send(rsp, info);
864
865 err_free_msg:
866 nlmsg_free(rsp);
867 return err;
868 }
869
870 static int
psp_nl_stats_fill(struct psp_dev * psd,struct sk_buff * rsp,const struct genl_info * info)871 psp_nl_stats_fill(struct psp_dev *psd, struct sk_buff *rsp,
872 const struct genl_info *info)
873 {
874 unsigned int required_cnt = sizeof(struct psp_dev_stats) / sizeof(u64);
875 struct psp_dev_stats stats;
876 void *hdr;
877 int i;
878
879 memset(&stats, 0xff, sizeof(stats));
880 psd->ops->get_stats(psd, &stats);
881
882 for (i = 0; i < required_cnt; i++)
883 if (WARN_ON_ONCE(stats.required[i] == ETHTOOL_STAT_NOT_SET))
884 return -EOPNOTSUPP;
885
886 hdr = genlmsg_iput(rsp, info);
887 if (!hdr)
888 return -EMSGSIZE;
889
890 if (nla_put_u32(rsp, PSP_A_STATS_DEV_ID, psd->id) ||
891 nla_put_uint(rsp, PSP_A_STATS_KEY_ROTATIONS,
892 psd->stats.rotations) ||
893 nla_put_uint(rsp, PSP_A_STATS_STALE_EVENTS, psd->stats.stales) ||
894 nla_put_uint(rsp, PSP_A_STATS_RX_PACKETS, stats.rx_packets) ||
895 nla_put_uint(rsp, PSP_A_STATS_RX_BYTES, stats.rx_bytes) ||
896 nla_put_uint(rsp, PSP_A_STATS_RX_AUTH_FAIL, stats.rx_auth_fail) ||
897 nla_put_uint(rsp, PSP_A_STATS_RX_ERROR, stats.rx_error) ||
898 nla_put_uint(rsp, PSP_A_STATS_RX_BAD, stats.rx_bad) ||
899 nla_put_uint(rsp, PSP_A_STATS_TX_PACKETS, stats.tx_packets) ||
900 nla_put_uint(rsp, PSP_A_STATS_TX_BYTES, stats.tx_bytes) ||
901 nla_put_uint(rsp, PSP_A_STATS_TX_ERROR, stats.tx_error))
902 goto err_cancel_msg;
903
904 genlmsg_end(rsp, hdr);
905 return 0;
906
907 err_cancel_msg:
908 genlmsg_cancel(rsp, hdr);
909 return -EMSGSIZE;
910 }
911
psp_nl_get_stats_doit(struct sk_buff * skb,struct genl_info * info)912 int psp_nl_get_stats_doit(struct sk_buff *skb, struct genl_info *info)
913 {
914 struct psp_dev *psd = info->user_ptr[0];
915 struct sk_buff *rsp;
916 int err;
917
918 rsp = genlmsg_new(GENLMSG_DEFAULT_SIZE, GFP_KERNEL);
919 if (!rsp)
920 return -ENOMEM;
921
922 err = psp_nl_stats_fill(psd, rsp, info);
923 if (err)
924 goto err_free_msg;
925
926 return genlmsg_reply(rsp, info);
927
928 err_free_msg:
929 nlmsg_free(rsp);
930 return err;
931 }
932
933 static int
psp_nl_stats_get_dumpit_one(struct sk_buff * rsp,struct netlink_callback * cb,struct psp_dev * psd)934 psp_nl_stats_get_dumpit_one(struct sk_buff *rsp, struct netlink_callback *cb,
935 struct psp_dev *psd)
936 {
937 if (psp_dev_check_access(psd, sock_net(rsp->sk), false))
938 return 0;
939
940 return psp_nl_stats_fill(psd, rsp, genl_info_dump(cb));
941 }
942
psp_nl_get_stats_dumpit(struct sk_buff * rsp,struct netlink_callback * cb)943 int psp_nl_get_stats_dumpit(struct sk_buff *rsp, struct netlink_callback *cb)
944 {
945 struct psp_dev *psd;
946 int err = 0;
947
948 mutex_lock(&psp_devs_lock);
949 xa_for_each_start(&psp_devs, cb->args[0], psd, cb->args[0]) {
950 mutex_lock(&psd->lock);
951 err = psp_nl_stats_get_dumpit_one(rsp, cb, psd);
952 mutex_unlock(&psd->lock);
953 if (err)
954 break;
955 }
956 mutex_unlock(&psp_devs_lock);
957
958 return err;
959 }
960