1 // SPDX-License-Identifier: GPL-2.0-only
2
3 #include <linux/file.h>
4 #include <linux/net.h>
5 #include <linux/rcupdate.h>
6 #include <linux/tcp.h>
7
8 #include <net/ip.h>
9 #include <net/psp.h>
10 #include "psp.h"
11
psp_dev_get_for_sock(struct sock * sk)12 struct psp_dev *psp_dev_get_for_sock(struct sock *sk)
13 {
14 struct psp_dev *psd = NULL;
15 struct dst_entry *dst;
16
17 rcu_read_lock();
18 dst = __sk_dst_get(sk);
19 if (dst) {
20 psd = rcu_dereference(dst_dev_rcu(dst)->psp_dev);
21 if (psd && !psp_dev_tryget(psd))
22 psd = NULL;
23 }
24 rcu_read_unlock();
25
26 return psd;
27 }
28
29 static struct sk_buff *
psp_validate_xmit(struct sock * sk,struct net_device * dev,struct sk_buff * skb)30 psp_validate_xmit(struct sock *sk, struct net_device *dev, struct sk_buff *skb)
31 {
32 struct psp_assoc *pas;
33 bool good;
34
35 rcu_read_lock();
36 pas = psp_skb_get_assoc_rcu(skb);
37 good = !pas || rcu_access_pointer(dev->psp_dev) == pas->psd;
38 rcu_read_unlock();
39 if (!good) {
40 sk_skb_reason_drop(sk, skb, SKB_DROP_REASON_PSP_OUTPUT);
41 return NULL;
42 }
43
44 return skb;
45 }
46
psp_assoc_create(struct psp_dev * psd)47 struct psp_assoc *psp_assoc_create(struct psp_dev *psd)
48 {
49 struct psp_assoc *pas;
50
51 lockdep_assert_held(&psd->lock);
52
53 pas = kzalloc_flex(*pas, drv_data, psd->caps->assoc_drv_spc,
54 GFP_KERNEL_ACCOUNT);
55 if (!pas)
56 return NULL;
57
58 pas->psd = psd;
59 pas->dev_id = psd->id;
60 pas->generation = psd->generation;
61 psp_dev_get(psd);
62 refcount_set(&pas->refcnt, 1);
63
64 list_add_tail(&pas->assocs_list, &psd->active_assocs);
65
66 return pas;
67 }
68
psp_assoc_dummy(struct psp_assoc * pas)69 static struct psp_assoc *psp_assoc_dummy(struct psp_assoc *pas)
70 {
71 struct psp_dev *psd = pas->psd;
72 size_t sz;
73
74 lockdep_assert_held(&psd->lock);
75
76 sz = struct_size(pas, drv_data, psd->caps->assoc_drv_spc);
77 return kmemdup(pas, sz, GFP_KERNEL);
78 }
79
psp_dev_tx_key_add(struct psp_dev * psd,struct psp_assoc * pas,struct netlink_ext_ack * extack)80 static int psp_dev_tx_key_add(struct psp_dev *psd, struct psp_assoc *pas,
81 struct netlink_ext_ack *extack)
82 {
83 return psd->ops->tx_key_add(psd, pas, extack);
84 }
85
psp_dev_tx_key_del(struct psp_dev * psd,struct psp_assoc * pas)86 void psp_dev_tx_key_del(struct psp_dev *psd, struct psp_assoc *pas)
87 {
88 if (pas->tx.spi)
89 psd->ops->tx_key_del(psd, pas);
90 list_del(&pas->assocs_list);
91 }
92
psp_assoc_free(struct work_struct * work)93 static void psp_assoc_free(struct work_struct *work)
94 {
95 struct psp_assoc *pas = container_of(work, struct psp_assoc, work);
96 struct psp_dev *psd = pas->psd;
97
98 mutex_lock(&psd->lock);
99 if (psp_dev_is_registered(psd))
100 psp_dev_tx_key_del(psd, pas);
101 mutex_unlock(&psd->lock);
102 psp_dev_put(psd);
103 kfree(pas);
104 }
105
psp_assoc_free_queue(struct rcu_head * head)106 static void psp_assoc_free_queue(struct rcu_head *head)
107 {
108 struct psp_assoc *pas = container_of(head, struct psp_assoc, rcu);
109
110 INIT_WORK(&pas->work, psp_assoc_free);
111 schedule_work(&pas->work);
112 }
113
114 /**
115 * psp_assoc_put() - release a reference on a PSP association
116 * @pas: association to release
117 */
psp_assoc_put(struct psp_assoc * pas)118 void psp_assoc_put(struct psp_assoc *pas)
119 {
120 if (pas && refcount_dec_and_test(&pas->refcnt))
121 call_rcu(&pas->rcu, psp_assoc_free_queue);
122 }
123
psp_sk_assoc_free(struct sock * sk)124 void psp_sk_assoc_free(struct sock *sk)
125 {
126 struct psp_assoc *pas = rcu_dereference_protected(sk->psp_assoc, 1);
127
128 rcu_assign_pointer(sk->psp_assoc, NULL);
129 psp_assoc_put(pas);
130 }
131
psp_sock_assoc_set_rx(struct sock * sk,struct psp_assoc * pas,struct psp_key_parsed * key,struct netlink_ext_ack * extack)132 int psp_sock_assoc_set_rx(struct sock *sk, struct psp_assoc *pas,
133 struct psp_key_parsed *key,
134 struct netlink_ext_ack *extack)
135 {
136 int err;
137
138 memcpy(&pas->rx, key, sizeof(*key));
139
140 lock_sock(sk);
141
142 if (psp_sk_assoc(sk)) {
143 NL_SET_ERR_MSG(extack, "Socket already has PSP state");
144 err = -EBUSY;
145 goto exit_unlock;
146 } else if (sk_has_decrypt_user(sk)) {
147 NL_SET_ERR_MSG(extack, "Socket has incompatible state");
148 err = -EINVAL;
149 goto exit_unlock;
150 }
151
152 refcount_inc(&pas->refcnt);
153 rcu_assign_pointer(sk->psp_assoc, pas);
154 err = 0;
155
156 exit_unlock:
157 release_sock(sk);
158
159 return err;
160 }
161
psp_sock_recv_queue_check(struct sock * sk,struct psp_assoc * pas)162 static int psp_sock_recv_queue_check(struct sock *sk, struct psp_assoc *pas)
163 {
164 struct psp_skb_ext *pse;
165 struct sk_buff *skb;
166
167 skb_rbtree_walk(skb, &tcp_sk(sk)->out_of_order_queue) {
168 pse = skb_ext_find(skb, SKB_EXT_PSP);
169 if (!psp_pse_matches_pas(pse, pas))
170 return -EBUSY;
171 }
172
173 skb_queue_walk(&sk->sk_receive_queue, skb) {
174 pse = skb_ext_find(skb, SKB_EXT_PSP);
175 if (!psp_pse_matches_pas(pse, pas))
176 return -EBUSY;
177 }
178 return 0;
179 }
180
psp_sock_assoc_set_tx(struct sock * sk,struct psp_dev * psd,u32 version,struct psp_key_parsed * key,struct netlink_ext_ack * extack)181 int psp_sock_assoc_set_tx(struct sock *sk, struct psp_dev *psd,
182 u32 version, struct psp_key_parsed *key,
183 struct netlink_ext_ack *extack)
184 {
185 struct inet_connection_sock *icsk;
186 struct psp_assoc *pas, *dummy;
187 int err;
188
189 lock_sock(sk);
190
191 pas = psp_sk_assoc(sk);
192 if (!pas) {
193 NL_SET_ERR_MSG(extack, "Socket has no Rx key");
194 err = -EINVAL;
195 goto exit_unlock;
196 }
197 if (pas->psd != psd) {
198 NL_SET_ERR_MSG(extack, "Rx key from different device");
199 err = -EINVAL;
200 goto exit_unlock;
201 }
202 if (pas->version != version) {
203 NL_SET_ERR_MSG(extack,
204 "PSP version mismatch with existing state");
205 err = -EINVAL;
206 goto exit_unlock;
207 }
208 if (pas->tx.spi) {
209 NL_SET_ERR_MSG(extack, "Tx key already set");
210 err = -EBUSY;
211 goto exit_unlock;
212 }
213
214 err = psp_sock_recv_queue_check(sk, pas);
215 if (err) {
216 NL_SET_ERR_MSG(extack, "Socket has incompatible segments already in the recv queue");
217 goto exit_unlock;
218 }
219
220 /* Pass a fake association to drivers to make sure they don't
221 * try to store pointers to it. For re-keying we'll need to
222 * re-allocate the assoc structures.
223 */
224 dummy = psp_assoc_dummy(pas);
225 if (!dummy) {
226 err = -ENOMEM;
227 goto exit_unlock;
228 }
229
230 memcpy(&dummy->tx, key, sizeof(*key));
231 err = psp_dev_tx_key_add(psd, dummy, extack);
232 if (err)
233 goto exit_free_dummy;
234
235 memcpy(pas->drv_data, dummy->drv_data, psd->caps->assoc_drv_spc);
236 memcpy(&pas->tx, key, sizeof(*key));
237
238 WRITE_ONCE(sk->sk_validate_xmit_skb, psp_validate_xmit);
239 tcp_write_collapse_fence(sk);
240 pas->upgrade_seq = tcp_sk(sk)->rcv_nxt;
241
242 icsk = inet_csk(sk);
243 icsk->icsk_ext_hdr_len += psp_sk_overhead(sk);
244 icsk->icsk_sync_mss(sk, icsk->icsk_pmtu_cookie);
245
246 exit_free_dummy:
247 kfree(dummy);
248 exit_unlock:
249 release_sock(sk);
250 return err;
251 }
252
psp_assocs_key_rotated(struct psp_dev * psd)253 void psp_assocs_key_rotated(struct psp_dev *psd)
254 {
255 struct psp_assoc *pas, *next;
256
257 /* Mark the stale associations as invalid, they will no longer
258 * be able to Rx any traffic.
259 */
260 list_for_each_entry_safe(pas, next, &psd->prev_assocs, assocs_list) {
261 pas->generation |= ~PSP_GEN_VALID_MASK;
262 psd->stats.stales++;
263 }
264 list_splice_init(&psd->prev_assocs, &psd->stale_assocs);
265 list_splice_init(&psd->active_assocs, &psd->prev_assocs);
266
267 /* TODO: we should inform the sockets that got shut down */
268 }
269
psp_twsk_init(struct inet_timewait_sock * tw,const struct sock * sk)270 void psp_twsk_init(struct inet_timewait_sock *tw, const struct sock *sk)
271 {
272 struct psp_assoc *pas = psp_sk_assoc(sk);
273
274 if (pas)
275 refcount_inc(&pas->refcnt);
276 rcu_assign_pointer(tw->psp_assoc, pas);
277 tw->tw_validate_xmit_skb = psp_validate_xmit;
278 }
279
psp_twsk_assoc_free(struct inet_timewait_sock * tw)280 void psp_twsk_assoc_free(struct inet_timewait_sock *tw)
281 {
282 struct psp_assoc *pas = rcu_dereference_protected(tw->psp_assoc, 1);
283
284 rcu_assign_pointer(tw->psp_assoc, NULL);
285 psp_assoc_put(pas);
286 }
287
psp_reply_set_decrypted(const struct sock * sk,struct sk_buff * skb)288 void psp_reply_set_decrypted(const struct sock *sk, struct sk_buff *skb)
289 {
290 struct psp_assoc *pas;
291
292 rcu_read_lock();
293 pas = psp_sk_get_assoc_rcu(sk);
294 if (pas && pas->tx.spi)
295 skb->decrypted = 1;
296 rcu_read_unlock();
297 }
298