xref: /linux/net/psp/psp_sock.c (revision b5a051f6b840d48f159166ef073d3021989bfb50)
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