xref: /linux/drivers/net/ovpn/tcp.c (revision fab183d632628381b466a41479489541ac0e29a0)
1 // SPDX-License-Identifier: GPL-2.0
2 /*  OpenVPN data channel offload
3  *
4  *  Copyright (C) 2019-2025 OpenVPN, Inc.
5  *
6  *  Author:	Antonio Quartulli <antonio@openvpn.net>
7  */
8 
9 #include <linux/skbuff.h>
10 #include <net/hotdata.h>
11 #include <net/inet_common.h>
12 #include <net/ipv6.h>
13 #include <net/tcp.h>
14 #include <net/transp_v6.h>
15 #include <net/route.h>
16 #include <trace/events/sock.h>
17 
18 #include "ovpnpriv.h"
19 #include "main.h"
20 #include "io.h"
21 #include "peer.h"
22 #include "proto.h"
23 #include "skb.h"
24 #include "tcp.h"
25 
26 #define OVPN_TCP_DEPTH_NESTING	2
27 #if OVPN_TCP_DEPTH_NESTING == SINGLE_DEPTH_NESTING
28 #error "OVPN TCP requires its own lockdep subclass"
29 #endif
30 
31 static struct proto ovpn_tcp_prot __ro_after_init;
32 static struct proto_ops ovpn_tcp_ops __ro_after_init;
33 static struct proto ovpn_tcp6_prot __ro_after_init;
34 static struct proto_ops ovpn_tcp6_ops __ro_after_init;
35 
ovpn_tcp_parse(struct strparser * strp,struct sk_buff * skb)36 static int ovpn_tcp_parse(struct strparser *strp, struct sk_buff *skb)
37 {
38 	struct strp_msg *rxm = strp_msg(skb);
39 	__be16 blen;
40 	u16 len;
41 	int err;
42 
43 	/* when packets are written to the TCP stream, they are prepended with
44 	 * two bytes indicating the actual packet size.
45 	 * Parse accordingly and return the actual size (including the size
46 	 * header)
47 	 */
48 
49 	if (skb->len < rxm->offset + 2)
50 		return 0;
51 
52 	err = skb_copy_bits(skb, rxm->offset, &blen, sizeof(blen));
53 	if (err < 0)
54 		return err;
55 
56 	len = be16_to_cpu(blen);
57 	if (len < 2)
58 		return -EINVAL;
59 
60 	return len + 2;
61 }
62 
63 /* queue skb for sending to userspace via recvmsg on the socket */
ovpn_tcp_to_userspace(struct ovpn_peer * peer,struct sock * sk,struct sk_buff * skb)64 static void ovpn_tcp_to_userspace(struct ovpn_peer *peer, struct sock *sk,
65 				  struct sk_buff *skb)
66 {
67 	skb_set_owner_r(skb, sk);
68 	memset(skb->cb, 0, sizeof(skb->cb));
69 	skb_queue_tail(&peer->tcp.user_queue, skb);
70 	peer->tcp.sk_cb.sk_data_ready(sk);
71 }
72 
ovpn_tcp_skb_packet(const struct ovpn_peer * peer,struct sk_buff * orig_skb,const int pkt_len,const int pkt_off)73 static struct sk_buff *ovpn_tcp_skb_packet(const struct ovpn_peer *peer,
74 					   struct sk_buff *orig_skb,
75 					   const int pkt_len, const int pkt_off)
76 {
77 	struct sk_buff *ovpn_skb;
78 	int err;
79 
80 	/* create a new skb with only the content of the current packet */
81 	ovpn_skb = netdev_alloc_skb(peer->ovpn->dev, pkt_len);
82 	if (unlikely(!ovpn_skb))
83 		goto err;
84 
85 	skb_copy_header(ovpn_skb, orig_skb);
86 	err = skb_copy_bits(orig_skb, pkt_off, skb_put(ovpn_skb, pkt_len),
87 			    pkt_len);
88 	if (unlikely(err)) {
89 		net_warn_ratelimited("%s: skb_copy_bits failed for peer %u\n",
90 				     netdev_name(peer->ovpn->dev), peer->id);
91 		kfree_skb(ovpn_skb);
92 		goto err;
93 	}
94 
95 	consume_skb(orig_skb);
96 	return ovpn_skb;
97 err:
98 	kfree_skb(orig_skb);
99 	return NULL;
100 }
101 
ovpn_tcp_rcv(struct strparser * strp,struct sk_buff * skb)102 static void ovpn_tcp_rcv(struct strparser *strp, struct sk_buff *skb)
103 {
104 	struct ovpn_peer *peer = container_of(strp, struct ovpn_peer, tcp.strp);
105 	struct strp_msg *msg = strp_msg(skb);
106 	int pkt_len = msg->full_len - 2;
107 	u8 opcode;
108 
109 	/* we need at least 4 bytes of data in the packet
110 	 * to extract the opcode and the key ID later on
111 	 */
112 	if (unlikely(pkt_len < OVPN_OPCODE_SIZE)) {
113 		net_warn_ratelimited("%s: packet too small to fetch opcode for peer %u\n",
114 				     netdev_name(peer->ovpn->dev), peer->id);
115 		goto err;
116 	}
117 
118 	/* extract the packet into a new skb */
119 	skb = ovpn_tcp_skb_packet(peer, skb, pkt_len, msg->offset + 2);
120 	if (unlikely(!skb))
121 		goto err;
122 
123 	/* DATA_V2 packets are handled in kernel, the rest goes to user space */
124 	opcode = ovpn_opcode_from_skb(skb, 0);
125 	if (unlikely(opcode != OVPN_DATA_V2)) {
126 		if (opcode == OVPN_DATA_V1) {
127 			net_warn_ratelimited("%s: DATA_V1 detected on the TCP stream\n",
128 					     netdev_name(peer->ovpn->dev));
129 			goto err;
130 		}
131 
132 		/* The packet size header must be there when sending the packet
133 		 * to userspace, therefore we put it back
134 		 */
135 		*(__be16 *)__skb_push(skb, sizeof(u16)) = htons(pkt_len);
136 		ovpn_tcp_to_userspace(peer, strp->sk, skb);
137 		return;
138 	}
139 
140 	/* hold reference to peer as required by ovpn_recv().
141 	 *
142 	 * NOTE: in this context we should already be holding a reference to
143 	 * this peer, therefore ovpn_peer_hold() is not expected to fail
144 	 */
145 	if (WARN_ON(!ovpn_peer_hold(peer)))
146 		goto err_nopeer;
147 
148 	ovpn_recv(peer, skb);
149 	return;
150 err:
151 	/* take reference for deferred peer deletion. should never fail */
152 	if (WARN_ON(!ovpn_peer_hold(peer)))
153 		goto err_nopeer;
154 	if (!queue_work(ovpn_wq, &peer->tcp.defer_del_work))
155 		ovpn_peer_put(peer);
156 	ovpn_dev_dstats_rx_dropped(peer->ovpn->dev);
157 err_nopeer:
158 	kfree_skb(skb);
159 }
160 
ovpn_tcp_recvmsg(struct sock * sk,struct msghdr * msg,size_t len,int flags)161 static int ovpn_tcp_recvmsg(struct sock *sk, struct msghdr *msg, size_t len,
162 			    int flags)
163 {
164 	int err = 0, off, copied = 0, ret;
165 	struct ovpn_socket *sock;
166 	struct ovpn_peer *peer;
167 	struct sk_buff *skb;
168 
169 	rcu_read_lock();
170 	sock = rcu_dereference_sk_user_data(sk);
171 	if (unlikely(!sock || !sock->peer || !ovpn_peer_hold(sock->peer))) {
172 		rcu_read_unlock();
173 		return -EBADF;
174 	}
175 	peer = sock->peer;
176 	rcu_read_unlock();
177 
178 	skb = __skb_recv_datagram(sk, &peer->tcp.user_queue, flags, &off, &err);
179 	if (!skb) {
180 		if (err == -EAGAIN && sk->sk_shutdown & RCV_SHUTDOWN) {
181 			ret = 0;
182 			goto out;
183 		}
184 		ret = err;
185 		goto out;
186 	}
187 
188 	copied = len;
189 	if (copied > skb->len)
190 		copied = skb->len;
191 	else if (copied < skb->len)
192 		msg->msg_flags |= MSG_TRUNC;
193 
194 	err = skb_copy_datagram_msg(skb, 0, msg, copied);
195 	if (unlikely(err)) {
196 		kfree_skb(skb);
197 		ret = err;
198 		goto out;
199 	}
200 
201 	if (flags & MSG_TRUNC)
202 		copied = skb->len;
203 	kfree_skb(skb);
204 	ret = copied;
205 out:
206 	ovpn_peer_put(peer);
207 	return ret;
208 }
209 
ovpn_tcp_socket_detach(struct ovpn_socket * ovpn_sock)210 void ovpn_tcp_socket_detach(struct ovpn_socket *ovpn_sock)
211 {
212 	struct ovpn_peer *peer = ovpn_sock->peer;
213 	struct sock *sk = ovpn_sock->sk;
214 
215 	strp_stop(&peer->tcp.strp);
216 	skb_queue_purge(&peer->tcp.user_queue);
217 
218 	/* restore CBs that were saved in ovpn_sock_set_tcp_cb() */
219 	sk->sk_data_ready = peer->tcp.sk_cb.sk_data_ready;
220 	sk->sk_write_space = peer->tcp.sk_cb.sk_write_space;
221 	sk->sk_prot = peer->tcp.sk_cb.prot;
222 
223 	/* tcp_close() may race this function and could set
224 	 * sk->sk_socket to NULL. It does so by invoking
225 	 * sock_orphan(), which holds sk_callback_lock before
226 	 * doing the assignment.
227 	 *
228 	 * For this reason we acquire the same lock to avoid
229 	 * sk_socket to disappear under our feet
230 	 */
231 	write_lock_bh(&sk->sk_callback_lock);
232 	if (sk->sk_socket)
233 		sk->sk_socket->ops = peer->tcp.sk_cb.ops;
234 	write_unlock_bh(&sk->sk_callback_lock);
235 
236 	rcu_assign_sk_user_data(sk, NULL);
237 }
238 
ovpn_tcp_socket_wait_finish(struct ovpn_socket * sock)239 void ovpn_tcp_socket_wait_finish(struct ovpn_socket *sock)
240 {
241 	struct ovpn_peer *peer = sock->peer;
242 
243 	/* NOTE: we don't wait for peer->tcp.defer_del_work to finish:
244 	 * either the worker is not running or this function
245 	 * was invoked by that worker.
246 	 */
247 
248 	cancel_work_sync(&sock->tcp_tx_work);
249 	strp_done(&peer->tcp.strp);
250 
251 	skb_queue_purge(&peer->tcp.out_queue);
252 	kfree_skb(peer->tcp.out_msg.skb);
253 	peer->tcp.out_msg.skb = NULL;
254 }
255 
ovpn_tcp_send_sock(struct ovpn_peer * peer,struct sock * sk)256 static void ovpn_tcp_send_sock(struct ovpn_peer *peer, struct sock *sk)
257 {
258 	struct sk_buff *skb = peer->tcp.out_msg.skb;
259 	int ret, flags;
260 
261 	if (!skb)
262 		return;
263 
264 	if (peer->tcp.tx_in_progress)
265 		return;
266 
267 	peer->tcp.tx_in_progress = true;
268 
269 	do {
270 		flags = ovpn_skb_cb(skb)->nosignal ? MSG_NOSIGNAL : 0;
271 		ret = skb_send_sock_locked_with_flags(sk, skb,
272 						      peer->tcp.out_msg.offset,
273 						      peer->tcp.out_msg.len,
274 						      flags);
275 		if (unlikely(ret < 0)) {
276 			if (ret == -EAGAIN)
277 				goto out;
278 
279 			net_warn_ratelimited("%s: TCP error to peer %u: %d\n",
280 					     netdev_name(peer->ovpn->dev),
281 					     peer->id, ret);
282 
283 			/* in case of TCP error we can't recover the VPN
284 			 * stream therefore we abort the connection
285 			 */
286 			ovpn_peer_hold(peer);
287 			if (!queue_work(ovpn_wq, &peer->tcp.defer_del_work))
288 				ovpn_peer_put(peer);
289 
290 			/* we bail out immediately and keep tx_in_progress set
291 			 * to true. This way we prevent more TX attempts
292 			 * which would lead to more invocations of queue_work()
293 			 */
294 			return;
295 		}
296 
297 		peer->tcp.out_msg.len -= ret;
298 		peer->tcp.out_msg.offset += ret;
299 	} while (peer->tcp.out_msg.len > 0);
300 
301 	if (!peer->tcp.out_msg.len) {
302 		local_bh_disable();
303 		dev_dstats_tx_add(peer->ovpn->dev, skb->len);
304 		local_bh_enable();
305 	}
306 
307 	kfree_skb(peer->tcp.out_msg.skb);
308 	peer->tcp.out_msg.skb = NULL;
309 	peer->tcp.out_msg.len = 0;
310 	peer->tcp.out_msg.offset = 0;
311 
312 out:
313 	peer->tcp.tx_in_progress = false;
314 }
315 
ovpn_tcp_tx_work(struct work_struct * work)316 void ovpn_tcp_tx_work(struct work_struct *work)
317 {
318 	struct ovpn_socket *sock;
319 
320 	sock = container_of(work, struct ovpn_socket, tcp_tx_work);
321 
322 	lock_sock(sock->sk);
323 	if (sock->peer)
324 		ovpn_tcp_send_sock(sock->peer, sock->sk);
325 	release_sock(sock->sk);
326 }
327 
ovpn_tcp_send_sock_skb(struct ovpn_peer * peer,struct sock * sk,struct sk_buff * skb)328 static void ovpn_tcp_send_sock_skb(struct ovpn_peer *peer, struct sock *sk,
329 				   struct sk_buff *skb)
330 {
331 	if (peer->tcp.out_msg.skb)
332 		ovpn_tcp_send_sock(peer, sk);
333 
334 	if (peer->tcp.out_msg.skb) {
335 		ovpn_dev_dstats_tx_dropped(peer->ovpn->dev);
336 		kfree_skb(skb);
337 		return;
338 	}
339 
340 	peer->tcp.out_msg.skb = skb;
341 	peer->tcp.out_msg.len = skb->len;
342 	peer->tcp.out_msg.offset = 0;
343 	ovpn_tcp_send_sock(peer, sk);
344 }
345 
ovpn_tcp_send_skb(struct ovpn_peer * peer,struct sock * sk,struct sk_buff * skb)346 void ovpn_tcp_send_skb(struct ovpn_peer *peer, struct sock *sk,
347 		       struct sk_buff *skb)
348 {
349 	u16 len = skb->len;
350 
351 	*(__be16 *)__skb_push(skb, sizeof(u16)) = htons(len);
352 
353 	spin_lock_nested(&sk->sk_lock.slock, OVPN_TCP_DEPTH_NESTING);
354 	if (sock_owned_by_user(sk)) {
355 		if (skb_queue_len(&peer->tcp.out_queue) >=
356 		    READ_ONCE(net_hotdata.max_backlog)) {
357 			ovpn_dev_dstats_tx_dropped(peer->ovpn->dev);
358 			kfree_skb(skb);
359 			goto unlock;
360 		}
361 		__skb_queue_tail(&peer->tcp.out_queue, skb);
362 	} else {
363 		ovpn_tcp_send_sock_skb(peer, sk, skb);
364 	}
365 unlock:
366 	spin_unlock(&sk->sk_lock.slock);
367 }
368 
ovpn_tcp_release(struct sock * sk)369 static void ovpn_tcp_release(struct sock *sk)
370 {
371 	struct sk_buff_head queue;
372 	struct ovpn_socket *sock;
373 	struct ovpn_peer *peer;
374 	struct sk_buff *skb;
375 
376 	rcu_read_lock();
377 	sock = rcu_dereference_sk_user_data(sk);
378 	if (!sock) {
379 		rcu_read_unlock();
380 		return;
381 	}
382 
383 	peer = sock->peer;
384 
385 	/* during initialization this function is called before
386 	 * assigning sock->peer
387 	 */
388 	if (unlikely(!peer || !ovpn_peer_hold(peer))) {
389 		rcu_read_unlock();
390 		return;
391 	}
392 	rcu_read_unlock();
393 
394 	__skb_queue_head_init(&queue);
395 	skb_queue_splice_init(&peer->tcp.out_queue, &queue);
396 
397 	while ((skb = __skb_dequeue(&queue)))
398 		ovpn_tcp_send_sock_skb(peer, sk, skb);
399 
400 	peer->tcp.sk_cb.prot->release_cb(sk);
401 	ovpn_peer_put(peer);
402 }
403 
ovpn_tcp_sendmsg(struct sock * sk,struct msghdr * msg,size_t size)404 static int ovpn_tcp_sendmsg(struct sock *sk, struct msghdr *msg, size_t size)
405 {
406 	struct ovpn_socket *sock;
407 	int ret, linear = PAGE_SIZE;
408 	struct ovpn_peer *peer;
409 	struct sk_buff *skb;
410 
411 	lock_sock(sk);
412 	rcu_read_lock();
413 	sock = rcu_dereference_sk_user_data(sk);
414 	if (unlikely(!sock || !sock->peer || !ovpn_peer_hold(sock->peer))) {
415 		rcu_read_unlock();
416 		release_sock(sk);
417 		return -EIO;
418 	}
419 	rcu_read_unlock();
420 	peer = sock->peer;
421 
422 	if (msg->msg_flags & ~(MSG_DONTWAIT | MSG_NOSIGNAL)) {
423 		ret = -EOPNOTSUPP;
424 		goto peer_free;
425 	}
426 
427 	if (peer->tcp.out_msg.skb) {
428 		ret = -EAGAIN;
429 		goto peer_free;
430 	}
431 
432 	if (size < linear)
433 		linear = size;
434 
435 	skb = sock_alloc_send_pskb(sk, linear, size - linear,
436 				   msg->msg_flags & MSG_DONTWAIT, &ret, 0);
437 	if (!skb) {
438 		net_err_ratelimited("%s: skb alloc failed: %d\n",
439 				    netdev_name(peer->ovpn->dev), ret);
440 		goto peer_free;
441 	}
442 
443 	skb_put(skb, linear);
444 	skb->len = size;
445 	skb->data_len = size - linear;
446 
447 	ret = skb_copy_datagram_from_iter(skb, 0, &msg->msg_iter, size);
448 	if (ret) {
449 		kfree_skb(skb);
450 		net_err_ratelimited("%s: skb copy from iter failed: %d\n",
451 				    netdev_name(peer->ovpn->dev), ret);
452 		goto peer_free;
453 	}
454 
455 	ovpn_skb_cb(skb)->nosignal = msg->msg_flags & MSG_NOSIGNAL;
456 	ovpn_tcp_send_sock_skb(peer, sk, skb);
457 	ret = size;
458 peer_free:
459 	release_sock(sk);
460 	ovpn_peer_put(peer);
461 	return ret;
462 }
463 
ovpn_tcp_disconnect(struct sock * sk,int flags)464 static int ovpn_tcp_disconnect(struct sock *sk, int flags)
465 {
466 	return -EBUSY;
467 }
468 
ovpn_tcp_data_ready(struct sock * sk)469 static void ovpn_tcp_data_ready(struct sock *sk)
470 {
471 	struct ovpn_socket *sock;
472 
473 	trace_sk_data_ready(sk);
474 
475 	rcu_read_lock();
476 	sock = rcu_dereference_sk_user_data(sk);
477 	if (likely(sock && sock->peer))
478 		strp_data_ready(&sock->peer->tcp.strp);
479 	rcu_read_unlock();
480 }
481 
ovpn_tcp_write_space(struct sock * sk)482 static void ovpn_tcp_write_space(struct sock *sk)
483 {
484 	struct ovpn_socket *sock;
485 
486 	rcu_read_lock();
487 	sock = rcu_dereference_sk_user_data(sk);
488 	if (likely(sock && sock->peer)) {
489 		queue_work(ovpn_wq, &sock->tcp_tx_work);
490 		sock->peer->tcp.sk_cb.sk_write_space(sk);
491 	}
492 	rcu_read_unlock();
493 }
494 
495 static void ovpn_tcp_build_protos(struct proto *new_prot,
496 				  struct proto_ops *new_ops,
497 				  const struct proto *orig_prot,
498 				  const struct proto_ops *orig_ops);
499 
ovpn_tcp_peer_del_work(struct work_struct * work)500 static void ovpn_tcp_peer_del_work(struct work_struct *work)
501 {
502 	struct ovpn_peer *peer = container_of(work, struct ovpn_peer,
503 					      tcp.defer_del_work);
504 
505 	ovpn_peer_del(peer, OVPN_DEL_PEER_REASON_TRANSPORT_ERROR);
506 	ovpn_peer_put(peer);
507 }
508 
509 /* Set TCP encapsulation callbacks */
ovpn_tcp_socket_attach(struct ovpn_socket * ovpn_sock,struct ovpn_peer * peer)510 int ovpn_tcp_socket_attach(struct ovpn_socket *ovpn_sock,
511 			   struct ovpn_peer *peer)
512 {
513 	struct strp_callbacks cb = {
514 		.rcv_msg = ovpn_tcp_rcv,
515 		.parse_msg = ovpn_tcp_parse,
516 	};
517 	int ret;
518 
519 	/* make sure no pre-existing encapsulation handler exists */
520 	if (ovpn_sock->sk->sk_user_data)
521 		return -EBUSY;
522 	rcu_assign_sk_user_data(ovpn_sock->sk, ovpn_sock);
523 
524 	/* only a fully connected socket is expected. Connection should be
525 	 * handled in userspace
526 	 */
527 	if (ovpn_sock->sk->sk_state != TCP_ESTABLISHED) {
528 		net_err_ratelimited("%s: provided TCP socket is not in ESTABLISHED state: %d\n",
529 				    netdev_name(peer->ovpn->dev),
530 				    ovpn_sock->sk->sk_state);
531 		ret = -EINVAL;
532 		goto err;
533 	}
534 
535 	ret = strp_init(&peer->tcp.strp, ovpn_sock->sk, &cb);
536 	if (ret < 0) {
537 		DEBUG_NET_WARN_ON_ONCE(1);
538 		goto err;
539 	}
540 
541 	INIT_WORK(&peer->tcp.defer_del_work, ovpn_tcp_peer_del_work);
542 
543 	__sk_dst_reset(ovpn_sock->sk);
544 	skb_queue_head_init(&peer->tcp.user_queue);
545 	skb_queue_head_init(&peer->tcp.out_queue);
546 
547 	/* save current CBs so that they can be restored upon socket release */
548 	peer->tcp.sk_cb.sk_data_ready = ovpn_sock->sk->sk_data_ready;
549 	peer->tcp.sk_cb.sk_write_space = ovpn_sock->sk->sk_write_space;
550 	peer->tcp.sk_cb.prot = ovpn_sock->sk->sk_prot;
551 	peer->tcp.sk_cb.ops = ovpn_sock->sk->sk_socket->ops;
552 
553 	/* assign our static CBs and prot/ops */
554 	ovpn_sock->sk->sk_data_ready = ovpn_tcp_data_ready;
555 	ovpn_sock->sk->sk_write_space = ovpn_tcp_write_space;
556 
557 	if (ovpn_sock->sk->sk_family == AF_INET) {
558 		ovpn_sock->sk->sk_prot = &ovpn_tcp_prot;
559 		ovpn_sock->sk->sk_socket->ops = &ovpn_tcp_ops;
560 	} else {
561 		ovpn_sock->sk->sk_prot = &ovpn_tcp6_prot;
562 		ovpn_sock->sk->sk_socket->ops = &ovpn_tcp6_ops;
563 	}
564 
565 	/* avoid using task_frag */
566 	ovpn_sock->sk->sk_allocation = GFP_ATOMIC;
567 	ovpn_sock->sk->sk_use_task_frag = false;
568 
569 	/* enqueue the RX worker */
570 	strp_check_rcv(&peer->tcp.strp);
571 
572 	return 0;
573 err:
574 	rcu_assign_sk_user_data(ovpn_sock->sk, NULL);
575 	return ret;
576 }
577 
ovpn_tcp_close(struct sock * sk,long timeout)578 static void ovpn_tcp_close(struct sock *sk, long timeout)
579 {
580 	struct ovpn_socket *sock;
581 	struct ovpn_peer *peer;
582 
583 	rcu_read_lock();
584 	sock = rcu_dereference_sk_user_data(sk);
585 	if (!sock) {
586 		rcu_read_unlock();
587 		return;
588 	}
589 
590 	peer = sock->peer;
591 	if (!peer || !ovpn_peer_hold(peer)) {
592 		rcu_read_unlock();
593 		return;
594 	}
595 	rcu_read_unlock();
596 
597 	ovpn_peer_del(peer, OVPN_DEL_PEER_REASON_TRANSPORT_DISCONNECT);
598 	peer->tcp.sk_cb.prot->close(sk, timeout);
599 	ovpn_peer_put(peer);
600 }
601 
ovpn_tcp_poll(struct file * file,struct socket * sock,poll_table * wait)602 static __poll_t ovpn_tcp_poll(struct file *file, struct socket *sock,
603 			      poll_table *wait)
604 {
605 	struct sk_buff_head *queue = &sock->sk->sk_receive_queue;
606 	struct ovpn_socket *ovpn_sock;
607 	struct ovpn_peer *peer = NULL;
608 	__poll_t mask;
609 
610 	rcu_read_lock();
611 	ovpn_sock = rcu_dereference_sk_user_data(sock->sk);
612 	/* if we landed in this callback, we expect to have a
613 	 * meaningful state. The ovpn_socket lifecycle would
614 	 * prevent it otherwise.
615 	 */
616 	if (WARN(!ovpn_sock || !ovpn_sock->peer,
617 		 "ovpn: null state in ovpn_tcp_poll!")) {
618 		rcu_read_unlock();
619 		return 0;
620 	}
621 
622 	if (ovpn_peer_hold(ovpn_sock->peer)) {
623 		peer = ovpn_sock->peer;
624 		queue = &peer->tcp.user_queue;
625 	}
626 	rcu_read_unlock();
627 
628 	mask = datagram_poll_queue(file, sock, wait, queue);
629 
630 	if (peer)
631 		ovpn_peer_put(peer);
632 
633 	return mask;
634 }
635 
ovpn_tcp_build_protos(struct proto * new_prot,struct proto_ops * new_ops,const struct proto * orig_prot,const struct proto_ops * orig_ops)636 static void ovpn_tcp_build_protos(struct proto *new_prot,
637 				  struct proto_ops *new_ops,
638 				  const struct proto *orig_prot,
639 				  const struct proto_ops *orig_ops)
640 {
641 	memcpy(new_prot, orig_prot, sizeof(*new_prot));
642 	memcpy(new_ops, orig_ops, sizeof(*new_ops));
643 	new_prot->recvmsg = ovpn_tcp_recvmsg;
644 	new_prot->sendmsg = ovpn_tcp_sendmsg;
645 	new_prot->disconnect = ovpn_tcp_disconnect;
646 	new_prot->close = ovpn_tcp_close;
647 	new_prot->release_cb = ovpn_tcp_release;
648 	new_ops->poll = ovpn_tcp_poll;
649 }
650 
651 /* Initialize TCP static objects */
ovpn_tcp_init(void)652 void __init ovpn_tcp_init(void)
653 {
654 	ovpn_tcp_build_protos(&ovpn_tcp_prot, &ovpn_tcp_ops, &tcp_prot,
655 			      &inet_stream_ops);
656 
657 #if IS_ENABLED(CONFIG_IPV6)
658 	ovpn_tcp_build_protos(&ovpn_tcp6_prot, &ovpn_tcp6_ops, &tcpv6_prot,
659 			      &inet6_stream_ops);
660 #endif
661 }
662