xref: /linux/net/vmw_vsock/virtio_transport_common.c (revision 7f063b2f17eaba2a35e251aa53627f2a70d536e2)
1 // SPDX-License-Identifier: GPL-2.0-only
2 /*
3  * common code for virtio vsock
4  *
5  * Copyright (C) 2013-2015 Red Hat, Inc.
6  * Author: Asias He <asias@redhat.com>
7  *         Stefan Hajnoczi <stefanha@redhat.com>
8  */
9 #include <linux/spinlock.h>
10 #include <linux/module.h>
11 #include <linux/sched/signal.h>
12 #include <linux/ctype.h>
13 #include <linux/list.h>
14 #include <linux/virtio_vsock.h>
15 #include <uapi/linux/vsockmon.h>
16 
17 #include <net/sock.h>
18 #include <net/af_vsock.h>
19 
20 #define CREATE_TRACE_POINTS
21 #include <trace/events/vsock_virtio_transport_common.h>
22 
23 /* How long to wait for graceful shutdown of a connection */
24 #define VSOCK_CLOSE_TIMEOUT (8 * HZ)
25 
26 /* Threshold for detecting small packets to copy */
27 #define GOOD_COPY_LEN  128
28 
29 /* Max payload that can be collapsed into a single linear skb, using the same
30  * allocation threshold as virtio_vsock_alloc_skb() to avoid adding pressure
31  * on the page allocator.
32  */
33 #define MAX_COLLAPSE_LEN \
34 	SKB_MAX_ORDER(VIRTIO_VSOCK_SKB_HEADROOM, PAGE_ALLOC_COSTLY_ORDER)
35 
36 static void virtio_transport_cancel_close_work(struct vsock_sock *vsk,
37 					       bool cancel_timeout);
38 static s64 virtio_transport_has_space(struct virtio_vsock_sock *vvs);
39 
40 static const struct virtio_transport *
virtio_transport_get_ops(struct vsock_sock * vsk)41 virtio_transport_get_ops(struct vsock_sock *vsk)
42 {
43 	const struct vsock_transport *t = vsock_core_get_transport(vsk);
44 
45 	if (WARN_ON(!t))
46 		return NULL;
47 
48 	return container_of(t, struct virtio_transport, transport);
49 }
50 
virtio_transport_can_zcopy(const struct virtio_transport * t_ops,struct virtio_vsock_pkt_info * info,size_t pkt_len)51 static bool virtio_transport_can_zcopy(const struct virtio_transport *t_ops,
52 				       struct virtio_vsock_pkt_info *info,
53 				       size_t pkt_len)
54 {
55 	struct iov_iter *iov_iter;
56 
57 	if (!info->msg)
58 		return false;
59 
60 	iov_iter = &info->msg->msg_iter;
61 
62 	if (iov_iter->iov_offset)
63 		return false;
64 
65 	/* We can't send whole iov. */
66 	if (iov_iter->count > pkt_len)
67 		return false;
68 
69 	/* Check that transport can send data in zerocopy mode. */
70 	if (t_ops->can_msgzerocopy) {
71 		int pages_to_send = iov_iter_npages(iov_iter, MAX_SKB_FRAGS);
72 
73 		/* +1 is for packet header. */
74 		return t_ops->can_msgzerocopy(pages_to_send + 1);
75 	}
76 
77 	return true;
78 }
79 
virtio_transport_fill_skb(struct sk_buff * skb,struct virtio_vsock_pkt_info * info,size_t len,bool zcopy)80 static int virtio_transport_fill_skb(struct sk_buff *skb,
81 				     struct virtio_vsock_pkt_info *info,
82 				     size_t len,
83 				     bool zcopy)
84 {
85 	struct msghdr *msg = info->msg;
86 
87 	if (zcopy)
88 		return __zerocopy_sg_from_iter(msg, NULL, skb,
89 					       &msg->msg_iter, len, NULL);
90 
91 	virtio_vsock_skb_put(skb, len);
92 	return skb_copy_datagram_from_iter_full(skb, 0, &msg->msg_iter, len);
93 }
94 
virtio_transport_init_hdr(struct sk_buff * skb,struct virtio_vsock_pkt_info * info,size_t payload_len,u32 src_cid,u32 src_port,u32 dst_cid,u32 dst_port)95 static void virtio_transport_init_hdr(struct sk_buff *skb,
96 				      struct virtio_vsock_pkt_info *info,
97 				      size_t payload_len,
98 				      u32 src_cid,
99 				      u32 src_port,
100 				      u32 dst_cid,
101 				      u32 dst_port)
102 {
103 	struct virtio_vsock_hdr *hdr;
104 
105 	hdr = virtio_vsock_hdr(skb);
106 	hdr->type	= cpu_to_le16(info->type);
107 	hdr->op		= cpu_to_le16(info->op);
108 	hdr->src_cid	= cpu_to_le64(src_cid);
109 	hdr->dst_cid	= cpu_to_le64(dst_cid);
110 	hdr->src_port	= cpu_to_le32(src_port);
111 	hdr->dst_port	= cpu_to_le32(dst_port);
112 	hdr->flags	= cpu_to_le32(info->flags);
113 	hdr->len	= cpu_to_le32(payload_len);
114 	hdr->buf_alloc	= cpu_to_le32(0);
115 	hdr->fwd_cnt	= cpu_to_le32(0);
116 }
117 
118 /* Packet capture */
virtio_transport_build_skb(void * opaque)119 static struct sk_buff *virtio_transport_build_skb(void *opaque)
120 {
121 	struct virtio_vsock_hdr *pkt_hdr;
122 	struct sk_buff *pkt = opaque;
123 	struct af_vsockmon_hdr *hdr;
124 	struct sk_buff *skb;
125 	size_t payload_len;
126 
127 	/* A packet could be split to fit the RX buffer, so we use
128 	 * the payload length from the header, which has been updated
129 	 * by the sender to reflect the fragment size.
130 	 */
131 	pkt_hdr = virtio_vsock_hdr(pkt);
132 	payload_len = le32_to_cpu(pkt_hdr->len);
133 
134 	skb = alloc_skb(sizeof(*hdr) + sizeof(*pkt_hdr) + payload_len,
135 			GFP_ATOMIC);
136 	if (!skb)
137 		return NULL;
138 
139 	hdr = skb_put(skb, sizeof(*hdr));
140 
141 	/* pkt->hdr is little-endian so no need to byteswap here */
142 	hdr->src_cid = pkt_hdr->src_cid;
143 	hdr->src_port = pkt_hdr->src_port;
144 	hdr->dst_cid = pkt_hdr->dst_cid;
145 	hdr->dst_port = pkt_hdr->dst_port;
146 
147 	hdr->transport = cpu_to_le16(AF_VSOCK_TRANSPORT_VIRTIO);
148 	hdr->len = cpu_to_le16(sizeof(*pkt_hdr));
149 	memset(hdr->reserved, 0, sizeof(hdr->reserved));
150 
151 	switch (le16_to_cpu(pkt_hdr->op)) {
152 	case VIRTIO_VSOCK_OP_REQUEST:
153 	case VIRTIO_VSOCK_OP_RESPONSE:
154 		hdr->op = cpu_to_le16(AF_VSOCK_OP_CONNECT);
155 		break;
156 	case VIRTIO_VSOCK_OP_RST:
157 	case VIRTIO_VSOCK_OP_SHUTDOWN:
158 		hdr->op = cpu_to_le16(AF_VSOCK_OP_DISCONNECT);
159 		break;
160 	case VIRTIO_VSOCK_OP_RW:
161 		hdr->op = cpu_to_le16(AF_VSOCK_OP_PAYLOAD);
162 		break;
163 	case VIRTIO_VSOCK_OP_CREDIT_UPDATE:
164 	case VIRTIO_VSOCK_OP_CREDIT_REQUEST:
165 		hdr->op = cpu_to_le16(AF_VSOCK_OP_CONTROL);
166 		break;
167 	default:
168 		hdr->op = cpu_to_le16(AF_VSOCK_OP_UNKNOWN);
169 		break;
170 	}
171 
172 	skb_put_data(skb, pkt_hdr, sizeof(*pkt_hdr));
173 
174 	if (payload_len) {
175 		struct iov_iter iov_iter;
176 		struct kvec kvec;
177 		void *data = skb_put(skb, payload_len);
178 
179 		kvec.iov_base = data;
180 		kvec.iov_len = payload_len;
181 		iov_iter_kvec(&iov_iter, ITER_DEST, &kvec, 1, payload_len);
182 
183 		if (skb_copy_datagram_iter(pkt, VIRTIO_VSOCK_SKB_CB(pkt)->offset,
184 					   &iov_iter, payload_len)) {
185 			kfree_skb(skb);
186 			return NULL;
187 		}
188 	}
189 
190 	return skb;
191 }
192 
virtio_transport_deliver_tap_pkt(struct sk_buff * skb)193 void virtio_transport_deliver_tap_pkt(struct sk_buff *skb)
194 {
195 	if (virtio_vsock_skb_tap_delivered(skb))
196 		return;
197 
198 	vsock_deliver_tap(virtio_transport_build_skb, skb);
199 	virtio_vsock_skb_set_tap_delivered(skb);
200 }
201 EXPORT_SYMBOL_GPL(virtio_transport_deliver_tap_pkt);
202 
virtio_transport_get_type(struct sock * sk)203 static u16 virtio_transport_get_type(struct sock *sk)
204 {
205 	if (sk->sk_type == SOCK_STREAM)
206 		return VIRTIO_VSOCK_TYPE_STREAM;
207 	else
208 		return VIRTIO_VSOCK_TYPE_SEQPACKET;
209 }
210 
211 /* Returns new sk_buff on success, otherwise returns NULL. */
virtio_transport_alloc_skb(struct virtio_vsock_pkt_info * info,size_t payload_len,bool zcopy,struct ubuf_info * uarg,u32 src_cid,u32 src_port,u32 dst_cid,u32 dst_port)212 static struct sk_buff *virtio_transport_alloc_skb(struct virtio_vsock_pkt_info *info,
213 						  size_t payload_len,
214 						  bool zcopy,
215 						  struct ubuf_info *uarg,
216 						  u32 src_cid,
217 						  u32 src_port,
218 						  u32 dst_cid,
219 						  u32 dst_port)
220 {
221 	struct vsock_sock *vsk;
222 	struct sk_buff *skb;
223 	size_t skb_len;
224 
225 	skb_len = VIRTIO_VSOCK_SKB_HEADROOM;
226 
227 	if (!zcopy)
228 		skb_len += payload_len;
229 
230 	skb = virtio_vsock_alloc_skb(skb_len, GFP_KERNEL);
231 	if (!skb)
232 		return NULL;
233 
234 	virtio_transport_init_hdr(skb, info, payload_len, src_cid, src_port,
235 				  dst_cid, dst_port);
236 
237 	vsk = info->vsk;
238 
239 	/* If 'vsk' != NULL then payload is always present, so we
240 	 * will never call '__zerocopy_sg_from_iter()' below without
241 	 * setting skb owner in 'skb_set_owner_w()'. The only case
242 	 * when 'vsk' == NULL is VIRTIO_VSOCK_OP_RST control message
243 	 * without payload.
244 	 */
245 	WARN_ON_ONCE(!(vsk && (info->msg && payload_len)) && zcopy);
246 
247 	/* Set owner here, because '__zerocopy_sg_from_iter()' uses
248 	 * owner of skb without check to update 'sk_wmem_alloc'.
249 	 */
250 	if (vsk)
251 		skb_set_owner_w(skb, sk_vsock(vsk));
252 
253 	if (info->msg && payload_len > 0) {
254 		int err;
255 
256 		/* Bind the zerocopy lifetime before filling frags so error
257 		 * rollback frees managed fixed-buffer pages through
258 		 * the uarg-aware path.
259 		 */
260 		skb_zcopy_set(skb, uarg, NULL);
261 
262 		err = virtio_transport_fill_skb(skb, info, payload_len, zcopy);
263 		if (err)
264 			goto out;
265 
266 		if (msg_data_left(info->msg) == 0 &&
267 		    info->type == VIRTIO_VSOCK_TYPE_SEQPACKET) {
268 			struct virtio_vsock_hdr *hdr = virtio_vsock_hdr(skb);
269 
270 			hdr->flags |= cpu_to_le32(VIRTIO_VSOCK_SEQ_EOM);
271 
272 			if (info->msg->msg_flags & MSG_EOR)
273 				hdr->flags |= cpu_to_le32(VIRTIO_VSOCK_SEQ_EOR);
274 		}
275 	}
276 
277 	if (info->reply)
278 		virtio_vsock_skb_set_reply(skb);
279 
280 	trace_virtio_transport_alloc_pkt(src_cid, src_port,
281 					 dst_cid, dst_port,
282 					 payload_len,
283 					 info->type,
284 					 info->op,
285 					 info->flags,
286 					 zcopy);
287 
288 	return skb;
289 out:
290 	kfree_skb(skb);
291 	return NULL;
292 }
293 
294 /* This function can only be used on connecting/connected sockets,
295  * since a socket assigned to a transport is required.
296  *
297  * Do not use on listener sockets!
298  */
virtio_transport_send_pkt_info(struct vsock_sock * vsk,struct virtio_vsock_pkt_info * info)299 static int virtio_transport_send_pkt_info(struct vsock_sock *vsk,
300 					  struct virtio_vsock_pkt_info *info)
301 {
302 	u32 max_skb_len = VIRTIO_VSOCK_MAX_PKT_BUF_SIZE;
303 	u32 src_cid, src_port, dst_cid, dst_port;
304 	const struct virtio_transport *t_ops;
305 	struct iov_iter_state msg_iter_state;
306 	struct virtio_vsock_sock *vvs;
307 	struct ubuf_info *uarg = NULL;
308 	u32 pkt_len = info->pkt_len;
309 	bool can_zcopy = false;
310 	bool have_uref = false;
311 	u32 rest_len;
312 	int ret;
313 
314 	info->type = virtio_transport_get_type(sk_vsock(vsk));
315 
316 	t_ops = virtio_transport_get_ops(vsk);
317 	if (unlikely(!t_ops))
318 		return -EFAULT;
319 
320 	src_cid = t_ops->transport.get_local_cid();
321 	src_port = vsk->local_addr.svm_port;
322 	if (!info->remote_cid) {
323 		dst_cid	= vsk->remote_addr.svm_cid;
324 		dst_port = vsk->remote_addr.svm_port;
325 	} else {
326 		dst_cid = info->remote_cid;
327 		dst_port = info->remote_port;
328 	}
329 
330 	vvs = vsk->trans;
331 
332 	/* virtio_transport_get_credit might return less than pkt_len credit */
333 	pkt_len = virtio_transport_get_credit(vvs, pkt_len);
334 
335 	/* Do not send zero length OP_RW pkt */
336 	if (pkt_len == 0 && info->op == VIRTIO_VSOCK_OP_RW)
337 		return pkt_len;
338 
339 	if (info->msg && (info->msg->msg_flags & MSG_ZEROCOPY)) {
340 		/* If 'info->msg' is not NULL, this is only VIRTIO_VSOCK_OP_RW.
341 		 * 'MSG_ZEROCOPY' flag handling here is based on the same flag
342 		 * handling from 'tcp_sendmsg_locked()'.
343 		 */
344 		if (info->msg->msg_ubuf) {
345 			uarg = info->msg->msg_ubuf;
346 			can_zcopy = virtio_transport_can_zcopy(t_ops, info, pkt_len);
347 		} else if (sock_flag(sk_vsock(vsk), SOCK_ZEROCOPY)) {
348 			uarg = msg_zerocopy_realloc(sk_vsock(vsk), pkt_len,
349 						    NULL, false);
350 			if (!uarg) {
351 				virtio_transport_put_credit(vvs, pkt_len);
352 				return -ENOMEM;
353 			}
354 
355 			can_zcopy = virtio_transport_can_zcopy(t_ops, info, pkt_len);
356 			if (!can_zcopy)
357 				uarg_to_msgzc(uarg)->zerocopy = 0;
358 
359 			have_uref = true;
360 		}
361 
362 		/* 'can_zcopy' means that this transmission will be
363 		 * in zerocopy way (e.g. using 'frags' array).
364 		 */
365 		if (can_zcopy)
366 			max_skb_len = min_t(u32, VIRTIO_VSOCK_MAX_PKT_BUF_SIZE,
367 					    (MAX_SKB_FRAGS * PAGE_SIZE));
368 	}
369 
370 	rest_len = pkt_len;
371 
372 	do {
373 		struct sk_buff *skb;
374 		size_t skb_len;
375 
376 		/* Save iterator state in case allocation or transmission fails
377 		 * so we can restore it and retry.
378 		 */
379 		if (info->msg)
380 			iov_iter_save_state(&info->msg->msg_iter, &msg_iter_state);
381 
382 		skb_len = min(max_skb_len, rest_len);
383 
384 		/* Note: virtio_transport_alloc_skb() can advance info->msg->msg_iter
385 		 * even if it fails (e.g. partial GUP success).
386 		 */
387 		skb = virtio_transport_alloc_skb(info, skb_len, can_zcopy,
388 						 uarg,
389 						 src_cid, src_port,
390 						 dst_cid, dst_port);
391 		if (!skb) {
392 			ret = -ENOMEM;
393 			break;
394 		}
395 
396 		virtio_transport_inc_tx_pkt(vvs, skb);
397 
398 		ret = t_ops->send_pkt(skb, info->net);
399 		if (ret < 0)
400 			break;
401 
402 		/* Both virtio and vhost 'send_pkt()' returns 'skb_len',
403 		 * but for reliability use 'ret' instead of 'skb_len'.
404 		 * Also if partial send happens (e.g. 'ret' != 'skb_len')
405 		 * somehow, we break this loop, but account such returned
406 		 * value in 'virtio_transport_put_credit()'.
407 		 */
408 		rest_len -= ret;
409 
410 		if (WARN_ONCE(ret != skb_len,
411 			      "'send_pkt()' returns %i, but %zu expected\n",
412 			      ret, skb_len))
413 			break;
414 	} while (rest_len);
415 
416 	if (info->msg && ret < 0)
417 		iov_iter_restore(&info->msg->msg_iter, &msg_iter_state);
418 
419 	virtio_transport_put_credit(vvs, rest_len);
420 
421 	/* msg_zerocopy_realloc() initializes the ubuf_info refcnt to 1.
422 	 * skb_zcopy_set() increases it for each skb, so we can drop that
423 	 * initial reference to keep it balanced.
424 	 */
425 	if (have_uref) {
426 		if (rest_len == pkt_len)
427 			/* No data sent, abort the notification. */
428 			net_zcopy_put_abort(uarg, true);
429 		else
430 			net_zcopy_put(uarg);
431 	}
432 
433 	/* Return number of bytes, if any data has been sent. */
434 	if (rest_len != pkt_len)
435 		ret = pkt_len - rest_len;
436 
437 	return ret;
438 }
439 
virtio_transport_can_collapse(struct sk_buff * skb)440 static bool virtio_transport_can_collapse(struct sk_buff *skb)
441 {
442 	/* skbs that are partially consumed, mark a SEQPACKET message boundary,
443 	 * or are already large enough should not be collapsed: they either
444 	 * need special accounting, carry protocol state, or already have a
445 	 * good data-to-overhead ratio.
446 	 */
447 	if (VIRTIO_VSOCK_SKB_CB(skb)->offset)
448 		return false;
449 	if (le32_to_cpu(virtio_vsock_hdr(skb)->flags) & VIRTIO_VSOCK_SEQ_EOM)
450 		return false;
451 	if (skb->len >= MAX_COLLAPSE_LEN)
452 		return false;
453 	return true;
454 }
455 
456 /* Iterate through the packets in the queue starting from the current skb to
457  * count the number of bytes we can collapse.
458  */
459 static unsigned int
virtio_transport_collapse_size(struct sk_buff * skb,struct sk_buff_head * queue)460 virtio_transport_collapse_size(struct sk_buff *skb, struct sk_buff_head *queue)
461 {
462 	unsigned int target = skb->len - VIRTIO_VSOCK_SKB_CB(skb)->offset;
463 
464 	while ((skb = skb_peek_next(skb, queue)) &&
465 	       virtio_transport_can_collapse(skb)) {
466 		unsigned int len = skb->len - VIRTIO_VSOCK_SKB_CB(skb)->offset;
467 
468 		if (len > MAX_COLLAPSE_LEN - target)
469 			return target;
470 
471 		target += len;
472 	}
473 
474 	return target;
475 }
476 
477 /* Called under lock_sock to compact the receive queue by merging small skbs.
478  * @min_to_free: minimum number of skbs to eliminate from the queue. May free
479  *               more to fill each collapsed skb to capacity.
480  */
481 static void
virtio_transport_collapse_rx_queue(struct virtio_vsock_sock * vvs,u32 min_to_free)482 virtio_transport_collapse_rx_queue(struct virtio_vsock_sock *vvs,
483 				   u32 min_to_free)
484 {
485 	struct sk_buff *skb, *next_skb, *new_skb = NULL;
486 	struct sk_buff_head new_queue;
487 	u32 saved = 0;
488 
489 	__skb_queue_head_init(&new_queue);
490 
491 	skb_queue_walk_safe(&vvs->rx_queue, skb, next_skb) {
492 		struct virtio_vsock_hdr *hdr = virtio_vsock_hdr(skb);
493 		u32 src_off = VIRTIO_VSOCK_SKB_CB(skb)->offset;
494 		u32 src_len = skb->len - src_off;
495 		bool keep;
496 
497 		keep = !virtio_transport_can_collapse(skb);
498 		if (keep) {
499 			/* Finalize pending collapsed skb to preserve packet
500 			 * ordering.
501 			 */
502 			if (new_skb) {
503 				__skb_queue_tail(&new_queue, new_skb);
504 				new_skb = NULL;
505 				saved--;
506 			}
507 			goto next;
508 		}
509 
510 		/* Finalize if this packet won't fit in the remaining tailroom,
511 		 * so we can allocate a right-sized new_skb.
512 		 */
513 		if (new_skb && src_len > skb_tailroom(new_skb)) {
514 			__skb_queue_tail(&new_queue, new_skb);
515 			new_skb = NULL;
516 			saved--;
517 		}
518 
519 		if (!new_skb) {
520 			unsigned int alloc_size;
521 
522 			/* Check after finalizing to opportunistically fill
523 			 * each collapsed skb to capacity, merging more skbs
524 			 * than strictly required.
525 			 */
526 			if (saved >= min_to_free)
527 				break;
528 
529 			alloc_size = virtio_transport_collapse_size(skb, &vvs->rx_queue);
530 
531 			/* Only this skb's data is eligible, nothing to merge
532 			 * with. Keep as-is.
533 			 */
534 			if (alloc_size <= src_len) {
535 				keep = true;
536 				goto next;
537 			}
538 
539 			new_skb = virtio_vsock_alloc_linear_skb(alloc_size +
540 					VIRTIO_VSOCK_SKB_HEADROOM, GFP_KERNEL);
541 			if (!new_skb)
542 				break;
543 
544 			memcpy(virtio_vsock_hdr(new_skb), hdr,
545 			       sizeof(struct virtio_vsock_hdr));
546 			virtio_vsock_hdr(new_skb)->len = 0;
547 		}
548 
549 		/* Cannot fail since src_off/src_len are within bounds, but if
550 		 * it does, discard new_skb to avoid queuing corrupted data.
551 		 */
552 		if (WARN_ON_ONCE(skb_copy_bits(skb, src_off,
553 					       skb_put(new_skb, src_len),
554 					       src_len))) {
555 			kfree_skb(new_skb);
556 			new_skb = NULL;
557 			break;
558 		}
559 
560 		le32_add_cpu(&virtio_vsock_hdr(new_skb)->len, src_len);
561 		virtio_vsock_hdr(new_skb)->flags |= hdr->flags;
562 
563 next:
564 		__skb_unlink(skb, &vvs->rx_queue);
565 		if (keep) {
566 			__skb_queue_tail(&new_queue, skb);
567 		} else {
568 			consume_skb(skb);
569 			saved++;
570 		}
571 	}
572 
573 	if (new_skb)
574 		__skb_queue_tail(&new_queue, new_skb);
575 
576 	skb_queue_splice(&new_queue, &vvs->rx_queue);
577 }
578 
virtio_transport_inc_rx_pkt(struct virtio_vsock_sock * vvs,u32 len)579 static bool virtio_transport_inc_rx_pkt(struct virtio_vsock_sock *vvs,
580 					u32 len)
581 {
582 	u64 skb_overhead = (skb_queue_len(&vvs->rx_queue) + 1) * SKB_TRUESIZE(0);
583 
584 	/* Allow at most buf_alloc * 2 total budget (payload + overhead),
585 	 * similar to how SO_RCVBUF is doubled to reserve space for sk_buff
586 	 * metadata. Check payload against buf_alloc to be sure the other
587 	 * peer is respecting the credit, and sk_buff overhead to bound
588 	 * queue growth.
589 	 */
590 	if ((u64)vvs->buf_used + len > vvs->buf_alloc ||
591 	    skb_overhead > vvs->buf_alloc)
592 		return false;
593 
594 	vvs->rx_bytes += len;
595 	vvs->buf_used += len;
596 	return true;
597 }
598 
virtio_transport_dec_rx_pkt(struct virtio_vsock_sock * vvs,u32 bytes_read,u32 bytes_dequeued)599 static void virtio_transport_dec_rx_pkt(struct virtio_vsock_sock *vvs,
600 					u32 bytes_read, u32 bytes_dequeued)
601 {
602 	vvs->rx_bytes -= bytes_read;
603 	vvs->buf_used -= bytes_dequeued;
604 	vvs->fwd_cnt += bytes_dequeued;
605 }
606 
virtio_transport_inc_tx_pkt(struct virtio_vsock_sock * vvs,struct sk_buff * skb)607 void virtio_transport_inc_tx_pkt(struct virtio_vsock_sock *vvs, struct sk_buff *skb)
608 {
609 	struct virtio_vsock_hdr *hdr = virtio_vsock_hdr(skb);
610 
611 	spin_lock_bh(&vvs->rx_lock);
612 	vvs->last_fwd_cnt = vvs->fwd_cnt;
613 	hdr->fwd_cnt = cpu_to_le32(vvs->fwd_cnt);
614 	hdr->buf_alloc = cpu_to_le32(vvs->buf_alloc);
615 	spin_unlock_bh(&vvs->rx_lock);
616 }
617 EXPORT_SYMBOL_GPL(virtio_transport_inc_tx_pkt);
618 
virtio_transport_consume_skb_sent(struct sk_buff * skb,bool consume)619 void virtio_transport_consume_skb_sent(struct sk_buff *skb, bool consume)
620 {
621 	struct sock *s = skb->sk;
622 
623 	if (s && skb->len) {
624 		struct vsock_sock *vs = vsock_sk(s);
625 		struct virtio_vsock_sock *vvs;
626 
627 		vvs = vs->trans;
628 
629 		spin_lock_bh(&vvs->tx_lock);
630 		vvs->bytes_unsent -= skb->len;
631 		spin_unlock_bh(&vvs->tx_lock);
632 	}
633 
634 	if (consume)
635 		consume_skb(skb);
636 }
637 EXPORT_SYMBOL_GPL(virtio_transport_consume_skb_sent);
638 
virtio_transport_get_credit(struct virtio_vsock_sock * vvs,u32 credit)639 u32 virtio_transport_get_credit(struct virtio_vsock_sock *vvs, u32 credit)
640 {
641 	u32 ret;
642 
643 	if (!credit)
644 		return 0;
645 
646 	spin_lock_bh(&vvs->tx_lock);
647 	ret = min_t(u32, credit, virtio_transport_has_space(vvs));
648 	vvs->tx_cnt += ret;
649 	vvs->bytes_unsent += ret;
650 	spin_unlock_bh(&vvs->tx_lock);
651 
652 	return ret;
653 }
654 EXPORT_SYMBOL_GPL(virtio_transport_get_credit);
655 
virtio_transport_put_credit(struct virtio_vsock_sock * vvs,u32 credit)656 void virtio_transport_put_credit(struct virtio_vsock_sock *vvs, u32 credit)
657 {
658 	if (!credit)
659 		return;
660 
661 	spin_lock_bh(&vvs->tx_lock);
662 	vvs->tx_cnt -= credit;
663 	vvs->bytes_unsent -= credit;
664 	spin_unlock_bh(&vvs->tx_lock);
665 }
666 EXPORT_SYMBOL_GPL(virtio_transport_put_credit);
667 
virtio_transport_send_credit_update(struct vsock_sock * vsk)668 static int virtio_transport_send_credit_update(struct vsock_sock *vsk)
669 {
670 	struct virtio_vsock_pkt_info info = {
671 		.op = VIRTIO_VSOCK_OP_CREDIT_UPDATE,
672 		.vsk = vsk,
673 		.net = sock_net(sk_vsock(vsk)),
674 	};
675 
676 	return virtio_transport_send_pkt_info(vsk, &info);
677 }
678 
679 static ssize_t
virtio_transport_stream_do_peek(struct vsock_sock * vsk,struct msghdr * msg,size_t len)680 virtio_transport_stream_do_peek(struct vsock_sock *vsk,
681 				struct msghdr *msg,
682 				size_t len)
683 {
684 	struct virtio_vsock_sock *vvs = vsk->trans;
685 	struct sk_buff *skb;
686 	size_t total = 0;
687 	int err;
688 
689 	spin_lock_bh(&vvs->rx_lock);
690 
691 	skb_queue_walk(&vvs->rx_queue, skb) {
692 		size_t bytes;
693 
694 		bytes = min_t(size_t, len - total,
695 			      skb->len - VIRTIO_VSOCK_SKB_CB(skb)->offset);
696 
697 		spin_unlock_bh(&vvs->rx_lock);
698 
699 		/* sk_lock is held by caller so no one else can dequeue.
700 		 * Unlock rx_lock since skb_copy_datagram_iter() may sleep.
701 		 */
702 		err = skb_copy_datagram_iter(skb, VIRTIO_VSOCK_SKB_CB(skb)->offset,
703 					     &msg->msg_iter, bytes);
704 		if (err)
705 			goto out;
706 
707 		total += bytes;
708 
709 		spin_lock_bh(&vvs->rx_lock);
710 
711 		if (total == len)
712 			break;
713 	}
714 
715 	spin_unlock_bh(&vvs->rx_lock);
716 
717 	return total;
718 
719 out:
720 	if (total)
721 		err = total;
722 	return err;
723 }
724 
725 static ssize_t
virtio_transport_stream_do_dequeue(struct vsock_sock * vsk,struct msghdr * msg,size_t len)726 virtio_transport_stream_do_dequeue(struct vsock_sock *vsk,
727 				   struct msghdr *msg,
728 				   size_t len)
729 {
730 	struct virtio_vsock_sock *vvs = vsk->trans;
731 	struct sk_buff *skb;
732 	u32 fwd_cnt_delta;
733 	bool low_rx_bytes;
734 	int err = -EFAULT;
735 	size_t total = 0;
736 	u32 free_space;
737 
738 	spin_lock_bh(&vvs->rx_lock);
739 
740 	if (WARN_ONCE(skb_queue_empty(&vvs->rx_queue) && vvs->rx_bytes,
741 		      "rx_queue is empty, but rx_bytes is non-zero\n")) {
742 		spin_unlock_bh(&vvs->rx_lock);
743 		return err;
744 	}
745 
746 	while (total < len && !skb_queue_empty(&vvs->rx_queue)) {
747 		size_t bytes, dequeued = 0;
748 
749 		skb = skb_peek(&vvs->rx_queue);
750 
751 		bytes = min_t(size_t, len - total,
752 			      skb->len - VIRTIO_VSOCK_SKB_CB(skb)->offset);
753 
754 		/* sk_lock is held by caller so no one else can dequeue.
755 		 * Unlock rx_lock since skb_copy_datagram_iter() may sleep.
756 		 */
757 		spin_unlock_bh(&vvs->rx_lock);
758 
759 		err = skb_copy_datagram_iter(skb,
760 					     VIRTIO_VSOCK_SKB_CB(skb)->offset,
761 					     &msg->msg_iter, bytes);
762 		if (err)
763 			goto out;
764 
765 		spin_lock_bh(&vvs->rx_lock);
766 
767 		total += bytes;
768 
769 		VIRTIO_VSOCK_SKB_CB(skb)->offset += bytes;
770 
771 		if (skb->len == VIRTIO_VSOCK_SKB_CB(skb)->offset) {
772 			dequeued = le32_to_cpu(virtio_vsock_hdr(skb)->len);
773 			__skb_unlink(skb, &vvs->rx_queue);
774 			consume_skb(skb);
775 		}
776 
777 		virtio_transport_dec_rx_pkt(vvs, bytes, dequeued);
778 	}
779 
780 	fwd_cnt_delta = vvs->fwd_cnt - vvs->last_fwd_cnt;
781 	free_space = vvs->buf_alloc - fwd_cnt_delta;
782 	low_rx_bytes = (vvs->rx_bytes <
783 			sock_rcvlowat(sk_vsock(vsk), 0, INT_MAX));
784 
785 	spin_unlock_bh(&vvs->rx_lock);
786 
787 	/* To reduce the number of credit update messages,
788 	 * don't update credits as long as lots of space is available.
789 	 * Note: the limit chosen here is arbitrary. Setting the limit
790 	 * too high causes extra messages. Too low causes transmitter
791 	 * stalls. As stalls are in theory more expensive than extra
792 	 * messages, we set the limit to a high value. TODO: experiment
793 	 * with different values. Also send credit update message when
794 	 * number of bytes in rx queue is not enough to wake up reader.
795 	 */
796 	if (fwd_cnt_delta &&
797 	    (free_space < VIRTIO_VSOCK_MAX_PKT_BUF_SIZE || low_rx_bytes))
798 		virtio_transport_send_credit_update(vsk);
799 
800 	return total;
801 
802 out:
803 	if (total)
804 		err = total;
805 	return err;
806 }
807 
808 static ssize_t
virtio_transport_seqpacket_do_peek(struct vsock_sock * vsk,struct msghdr * msg)809 virtio_transport_seqpacket_do_peek(struct vsock_sock *vsk,
810 				   struct msghdr *msg)
811 {
812 	struct virtio_vsock_sock *vvs = vsk->trans;
813 	struct sk_buff *skb;
814 	size_t total, len;
815 
816 	spin_lock_bh(&vvs->rx_lock);
817 
818 	if (!vvs->msg_count) {
819 		spin_unlock_bh(&vvs->rx_lock);
820 		return 0;
821 	}
822 
823 	total = 0;
824 	len = msg_data_left(msg);
825 
826 	skb_queue_walk(&vvs->rx_queue, skb) {
827 		struct virtio_vsock_hdr *hdr;
828 
829 		if (total < len) {
830 			size_t bytes;
831 			int err;
832 
833 			bytes = len - total;
834 			if (bytes > skb->len)
835 				bytes = skb->len;
836 
837 			spin_unlock_bh(&vvs->rx_lock);
838 
839 			/* sk_lock is held by caller so no one else can dequeue.
840 			 * Unlock rx_lock since skb_copy_datagram_iter() may sleep.
841 			 */
842 			err = skb_copy_datagram_iter(skb, VIRTIO_VSOCK_SKB_CB(skb)->offset,
843 						     &msg->msg_iter, bytes);
844 			if (err)
845 				return err;
846 
847 			spin_lock_bh(&vvs->rx_lock);
848 		}
849 
850 		total += skb->len;
851 		hdr = virtio_vsock_hdr(skb);
852 
853 		if (le32_to_cpu(hdr->flags) & VIRTIO_VSOCK_SEQ_EOM) {
854 			if (le32_to_cpu(hdr->flags) & VIRTIO_VSOCK_SEQ_EOR)
855 				msg->msg_flags |= MSG_EOR;
856 
857 			break;
858 		}
859 	}
860 
861 	spin_unlock_bh(&vvs->rx_lock);
862 
863 	return total;
864 }
865 
virtio_transport_seqpacket_do_dequeue(struct vsock_sock * vsk,struct msghdr * msg,int flags)866 static int virtio_transport_seqpacket_do_dequeue(struct vsock_sock *vsk,
867 						 struct msghdr *msg,
868 						 int flags)
869 {
870 	struct virtio_vsock_sock *vvs = vsk->trans;
871 	int dequeued_len = 0;
872 	size_t user_buf_len = msg_data_left(msg);
873 	bool msg_ready = false;
874 	struct sk_buff *skb;
875 
876 	spin_lock_bh(&vvs->rx_lock);
877 
878 	if (vvs->msg_count == 0) {
879 		spin_unlock_bh(&vvs->rx_lock);
880 		return 0;
881 	}
882 
883 	while (!msg_ready) {
884 		struct virtio_vsock_hdr *hdr;
885 		size_t pkt_len;
886 
887 		skb = __skb_dequeue(&vvs->rx_queue);
888 		if (!skb)
889 			break;
890 		hdr = virtio_vsock_hdr(skb);
891 		pkt_len = (size_t)le32_to_cpu(hdr->len);
892 
893 		if (dequeued_len >= 0) {
894 			size_t bytes_to_copy;
895 
896 			bytes_to_copy = min(user_buf_len, pkt_len);
897 
898 			if (bytes_to_copy) {
899 				int err;
900 
901 				/* sk_lock is held by caller so no one else can dequeue.
902 				 * Unlock rx_lock since skb_copy_datagram_iter() may sleep.
903 				 */
904 				spin_unlock_bh(&vvs->rx_lock);
905 
906 				err = skb_copy_datagram_iter(skb, 0,
907 							     &msg->msg_iter,
908 							     bytes_to_copy);
909 				if (err) {
910 					/* Copy of message failed. Rest of
911 					 * fragments will be freed without copy.
912 					 */
913 					dequeued_len = err;
914 				} else {
915 					user_buf_len -= bytes_to_copy;
916 				}
917 
918 				spin_lock_bh(&vvs->rx_lock);
919 			}
920 
921 			if (dequeued_len >= 0)
922 				dequeued_len += pkt_len;
923 		}
924 
925 		if (le32_to_cpu(hdr->flags) & VIRTIO_VSOCK_SEQ_EOM) {
926 			msg_ready = true;
927 			vvs->msg_count--;
928 
929 			if (le32_to_cpu(hdr->flags) & VIRTIO_VSOCK_SEQ_EOR)
930 				msg->msg_flags |= MSG_EOR;
931 		}
932 
933 		virtio_transport_dec_rx_pkt(vvs, pkt_len, pkt_len);
934 		kfree_skb(skb);
935 	}
936 
937 	spin_unlock_bh(&vvs->rx_lock);
938 
939 	virtio_transport_send_credit_update(vsk);
940 
941 	return dequeued_len;
942 }
943 
944 ssize_t
virtio_transport_stream_dequeue(struct vsock_sock * vsk,struct msghdr * msg,size_t len,int flags)945 virtio_transport_stream_dequeue(struct vsock_sock *vsk,
946 				struct msghdr *msg,
947 				size_t len, int flags)
948 {
949 	if (flags & MSG_PEEK)
950 		return virtio_transport_stream_do_peek(vsk, msg, len);
951 	else
952 		return virtio_transport_stream_do_dequeue(vsk, msg, len);
953 }
954 EXPORT_SYMBOL_GPL(virtio_transport_stream_dequeue);
955 
956 ssize_t
virtio_transport_seqpacket_dequeue(struct vsock_sock * vsk,struct msghdr * msg,int flags)957 virtio_transport_seqpacket_dequeue(struct vsock_sock *vsk,
958 				   struct msghdr *msg,
959 				   int flags)
960 {
961 	if (flags & MSG_PEEK)
962 		return virtio_transport_seqpacket_do_peek(vsk, msg);
963 	else
964 		return virtio_transport_seqpacket_do_dequeue(vsk, msg, flags);
965 }
966 EXPORT_SYMBOL_GPL(virtio_transport_seqpacket_dequeue);
967 
virtio_transport_tx_buf_size(struct virtio_vsock_sock * vvs)968 static u32 virtio_transport_tx_buf_size(struct virtio_vsock_sock *vvs)
969 {
970 	/* The peer advertises its receive buffer via peer_buf_alloc, but we
971 	 * cap it to our local buf_alloc so a remote peer cannot force us to
972 	 * queue more data than our own buffer configuration allows.
973 	 */
974 	return min(vvs->peer_buf_alloc, vvs->buf_alloc);
975 }
976 
977 int
virtio_transport_seqpacket_enqueue(struct vsock_sock * vsk,struct msghdr * msg,size_t len)978 virtio_transport_seqpacket_enqueue(struct vsock_sock *vsk,
979 				   struct msghdr *msg,
980 				   size_t len)
981 {
982 	struct virtio_vsock_sock *vvs = vsk->trans;
983 
984 	spin_lock_bh(&vvs->tx_lock);
985 
986 	if (len > virtio_transport_tx_buf_size(vvs)) {
987 		spin_unlock_bh(&vvs->tx_lock);
988 		return -EMSGSIZE;
989 	}
990 
991 	spin_unlock_bh(&vvs->tx_lock);
992 
993 	return virtio_transport_stream_enqueue(vsk, msg, len);
994 }
995 EXPORT_SYMBOL_GPL(virtio_transport_seqpacket_enqueue);
996 
997 int
virtio_transport_dgram_dequeue(struct vsock_sock * vsk,struct msghdr * msg,size_t len,int flags)998 virtio_transport_dgram_dequeue(struct vsock_sock *vsk,
999 			       struct msghdr *msg,
1000 			       size_t len, int flags)
1001 {
1002 	return -EOPNOTSUPP;
1003 }
1004 EXPORT_SYMBOL_GPL(virtio_transport_dgram_dequeue);
1005 
virtio_transport_stream_has_data(struct vsock_sock * vsk)1006 s64 virtio_transport_stream_has_data(struct vsock_sock *vsk)
1007 {
1008 	struct virtio_vsock_sock *vvs = vsk->trans;
1009 	s64 bytes;
1010 
1011 	spin_lock_bh(&vvs->rx_lock);
1012 	bytes = vvs->rx_bytes;
1013 	spin_unlock_bh(&vvs->rx_lock);
1014 
1015 	return bytes;
1016 }
1017 EXPORT_SYMBOL_GPL(virtio_transport_stream_has_data);
1018 
virtio_transport_seqpacket_has_data(struct vsock_sock * vsk)1019 u32 virtio_transport_seqpacket_has_data(struct vsock_sock *vsk)
1020 {
1021 	struct virtio_vsock_sock *vvs = vsk->trans;
1022 	u32 msg_count;
1023 
1024 	spin_lock_bh(&vvs->rx_lock);
1025 	msg_count = vvs->msg_count;
1026 	spin_unlock_bh(&vvs->rx_lock);
1027 
1028 	return msg_count;
1029 }
1030 EXPORT_SYMBOL_GPL(virtio_transport_seqpacket_has_data);
1031 
virtio_transport_has_space(struct virtio_vsock_sock * vvs)1032 static s64 virtio_transport_has_space(struct virtio_vsock_sock *vvs)
1033 {
1034 	s64 bytes;
1035 
1036 	/* Use s64 arithmetic so if the peer shrinks peer_buf_alloc while
1037 	 * we have bytes in flight (tx_cnt - peer_fwd_cnt), the subtraction
1038 	 * does not underflow.
1039 	 */
1040 	bytes = (s64)virtio_transport_tx_buf_size(vvs) -
1041 		(vvs->tx_cnt - vvs->peer_fwd_cnt);
1042 	if (bytes < 0)
1043 		bytes = 0;
1044 
1045 	return bytes;
1046 }
1047 
virtio_transport_stream_has_space(struct vsock_sock * vsk)1048 s64 virtio_transport_stream_has_space(struct vsock_sock *vsk)
1049 {
1050 	struct virtio_vsock_sock *vvs = vsk->trans;
1051 	s64 bytes;
1052 
1053 	spin_lock_bh(&vvs->tx_lock);
1054 	bytes = virtio_transport_has_space(vvs);
1055 	spin_unlock_bh(&vvs->tx_lock);
1056 
1057 	return bytes;
1058 }
1059 EXPORT_SYMBOL_GPL(virtio_transport_stream_has_space);
1060 
virtio_transport_do_socket_init(struct vsock_sock * vsk,struct vsock_sock * psk)1061 int virtio_transport_do_socket_init(struct vsock_sock *vsk,
1062 				    struct vsock_sock *psk)
1063 {
1064 	struct virtio_vsock_sock *vvs;
1065 
1066 	vvs = kzalloc_obj(*vvs);
1067 	if (!vvs)
1068 		return -ENOMEM;
1069 
1070 	vsk->trans = vvs;
1071 	vvs->vsk = vsk;
1072 	if (psk && psk->trans) {
1073 		struct virtio_vsock_sock *ptrans = psk->trans;
1074 
1075 		vvs->peer_buf_alloc = ptrans->peer_buf_alloc;
1076 	}
1077 
1078 	if (vsk->buffer_size > VIRTIO_VSOCK_MAX_BUF_SIZE)
1079 		vsk->buffer_size = VIRTIO_VSOCK_MAX_BUF_SIZE;
1080 
1081 	vvs->buf_alloc = vsk->buffer_size;
1082 
1083 	spin_lock_init(&vvs->rx_lock);
1084 	spin_lock_init(&vvs->tx_lock);
1085 	skb_queue_head_init(&vvs->rx_queue);
1086 
1087 	return 0;
1088 }
1089 EXPORT_SYMBOL_GPL(virtio_transport_do_socket_init);
1090 
1091 /* sk_lock held by the caller */
virtio_transport_notify_buffer_size(struct vsock_sock * vsk,u64 * val)1092 void virtio_transport_notify_buffer_size(struct vsock_sock *vsk, u64 *val)
1093 {
1094 	struct virtio_vsock_sock *vvs = vsk->trans;
1095 
1096 	if (*val > VIRTIO_VSOCK_MAX_BUF_SIZE)
1097 		*val = VIRTIO_VSOCK_MAX_BUF_SIZE;
1098 
1099 	vvs->buf_alloc = *val;
1100 
1101 	virtio_transport_send_credit_update(vsk);
1102 }
1103 EXPORT_SYMBOL_GPL(virtio_transport_notify_buffer_size);
1104 
1105 int
virtio_transport_notify_poll_in(struct vsock_sock * vsk,size_t target,bool * data_ready_now)1106 virtio_transport_notify_poll_in(struct vsock_sock *vsk,
1107 				size_t target,
1108 				bool *data_ready_now)
1109 {
1110 	*data_ready_now = vsock_stream_has_data(vsk) >= target;
1111 
1112 	return 0;
1113 }
1114 EXPORT_SYMBOL_GPL(virtio_transport_notify_poll_in);
1115 
1116 int
virtio_transport_notify_poll_out(struct vsock_sock * vsk,size_t target,bool * space_avail_now)1117 virtio_transport_notify_poll_out(struct vsock_sock *vsk,
1118 				 size_t target,
1119 				 bool *space_avail_now)
1120 {
1121 	s64 free_space;
1122 
1123 	free_space = vsock_stream_has_space(vsk);
1124 	if (free_space > 0)
1125 		*space_avail_now = true;
1126 	else if (free_space == 0)
1127 		*space_avail_now = false;
1128 
1129 	return 0;
1130 }
1131 EXPORT_SYMBOL_GPL(virtio_transport_notify_poll_out);
1132 
virtio_transport_notify_recv_init(struct vsock_sock * vsk,size_t target,struct vsock_transport_recv_notify_data * data)1133 int virtio_transport_notify_recv_init(struct vsock_sock *vsk,
1134 	size_t target, struct vsock_transport_recv_notify_data *data)
1135 {
1136 	return 0;
1137 }
1138 EXPORT_SYMBOL_GPL(virtio_transport_notify_recv_init);
1139 
virtio_transport_notify_recv_pre_block(struct vsock_sock * vsk,size_t target,struct vsock_transport_recv_notify_data * data)1140 int virtio_transport_notify_recv_pre_block(struct vsock_sock *vsk,
1141 	size_t target, struct vsock_transport_recv_notify_data *data)
1142 {
1143 	return 0;
1144 }
1145 EXPORT_SYMBOL_GPL(virtio_transport_notify_recv_pre_block);
1146 
virtio_transport_notify_recv_pre_dequeue(struct vsock_sock * vsk,size_t target,struct vsock_transport_recv_notify_data * data)1147 int virtio_transport_notify_recv_pre_dequeue(struct vsock_sock *vsk,
1148 	size_t target, struct vsock_transport_recv_notify_data *data)
1149 {
1150 	return 0;
1151 }
1152 EXPORT_SYMBOL_GPL(virtio_transport_notify_recv_pre_dequeue);
1153 
virtio_transport_notify_recv_post_dequeue(struct vsock_sock * vsk,size_t target,ssize_t copied,bool data_read,struct vsock_transport_recv_notify_data * data)1154 int virtio_transport_notify_recv_post_dequeue(struct vsock_sock *vsk,
1155 	size_t target, ssize_t copied, bool data_read,
1156 	struct vsock_transport_recv_notify_data *data)
1157 {
1158 	return 0;
1159 }
1160 EXPORT_SYMBOL_GPL(virtio_transport_notify_recv_post_dequeue);
1161 
virtio_transport_notify_send_init(struct vsock_sock * vsk,struct vsock_transport_send_notify_data * data)1162 int virtio_transport_notify_send_init(struct vsock_sock *vsk,
1163 	struct vsock_transport_send_notify_data *data)
1164 {
1165 	return 0;
1166 }
1167 EXPORT_SYMBOL_GPL(virtio_transport_notify_send_init);
1168 
virtio_transport_notify_send_pre_block(struct vsock_sock * vsk,struct vsock_transport_send_notify_data * data)1169 int virtio_transport_notify_send_pre_block(struct vsock_sock *vsk,
1170 	struct vsock_transport_send_notify_data *data)
1171 {
1172 	return 0;
1173 }
1174 EXPORT_SYMBOL_GPL(virtio_transport_notify_send_pre_block);
1175 
virtio_transport_notify_send_pre_enqueue(struct vsock_sock * vsk,struct vsock_transport_send_notify_data * data)1176 int virtio_transport_notify_send_pre_enqueue(struct vsock_sock *vsk,
1177 	struct vsock_transport_send_notify_data *data)
1178 {
1179 	return 0;
1180 }
1181 EXPORT_SYMBOL_GPL(virtio_transport_notify_send_pre_enqueue);
1182 
virtio_transport_notify_send_post_enqueue(struct vsock_sock * vsk,ssize_t written,struct vsock_transport_send_notify_data * data)1183 int virtio_transport_notify_send_post_enqueue(struct vsock_sock *vsk,
1184 	ssize_t written, struct vsock_transport_send_notify_data *data)
1185 {
1186 	return 0;
1187 }
1188 EXPORT_SYMBOL_GPL(virtio_transport_notify_send_post_enqueue);
1189 
virtio_transport_stream_rcvhiwat(struct vsock_sock * vsk)1190 u64 virtio_transport_stream_rcvhiwat(struct vsock_sock *vsk)
1191 {
1192 	return vsk->buffer_size;
1193 }
1194 EXPORT_SYMBOL_GPL(virtio_transport_stream_rcvhiwat);
1195 
virtio_transport_stream_is_active(struct vsock_sock * vsk)1196 bool virtio_transport_stream_is_active(struct vsock_sock *vsk)
1197 {
1198 	return true;
1199 }
1200 EXPORT_SYMBOL_GPL(virtio_transport_stream_is_active);
1201 
virtio_transport_dgram_bind(struct vsock_sock * vsk,struct sockaddr_vm * addr)1202 int virtio_transport_dgram_bind(struct vsock_sock *vsk,
1203 				struct sockaddr_vm *addr)
1204 {
1205 	return -EOPNOTSUPP;
1206 }
1207 EXPORT_SYMBOL_GPL(virtio_transport_dgram_bind);
1208 
virtio_transport_dgram_allow(struct vsock_sock * vsk,u32 cid,u32 port)1209 bool virtio_transport_dgram_allow(struct vsock_sock *vsk, u32 cid, u32 port)
1210 {
1211 	return false;
1212 }
1213 EXPORT_SYMBOL_GPL(virtio_transport_dgram_allow);
1214 
virtio_transport_connect(struct vsock_sock * vsk)1215 int virtio_transport_connect(struct vsock_sock *vsk)
1216 {
1217 	struct virtio_vsock_pkt_info info = {
1218 		.op = VIRTIO_VSOCK_OP_REQUEST,
1219 		.vsk = vsk,
1220 		.net = sock_net(sk_vsock(vsk)),
1221 	};
1222 
1223 	return virtio_transport_send_pkt_info(vsk, &info);
1224 }
1225 EXPORT_SYMBOL_GPL(virtio_transport_connect);
1226 
virtio_transport_shutdown(struct vsock_sock * vsk,int mode)1227 int virtio_transport_shutdown(struct vsock_sock *vsk, int mode)
1228 {
1229 	struct virtio_vsock_pkt_info info = {
1230 		.op = VIRTIO_VSOCK_OP_SHUTDOWN,
1231 		.flags = (mode & RCV_SHUTDOWN ?
1232 			  VIRTIO_VSOCK_SHUTDOWN_RCV : 0) |
1233 			 (mode & SEND_SHUTDOWN ?
1234 			  VIRTIO_VSOCK_SHUTDOWN_SEND : 0),
1235 		.vsk = vsk,
1236 		.net = sock_net(sk_vsock(vsk)),
1237 	};
1238 
1239 	return virtio_transport_send_pkt_info(vsk, &info);
1240 }
1241 EXPORT_SYMBOL_GPL(virtio_transport_shutdown);
1242 
1243 int
virtio_transport_dgram_enqueue(struct vsock_sock * vsk,struct sockaddr_vm * remote_addr,struct msghdr * msg,size_t dgram_len)1244 virtio_transport_dgram_enqueue(struct vsock_sock *vsk,
1245 			       struct sockaddr_vm *remote_addr,
1246 			       struct msghdr *msg,
1247 			       size_t dgram_len)
1248 {
1249 	return -EOPNOTSUPP;
1250 }
1251 EXPORT_SYMBOL_GPL(virtio_transport_dgram_enqueue);
1252 
1253 ssize_t
virtio_transport_stream_enqueue(struct vsock_sock * vsk,struct msghdr * msg,size_t len)1254 virtio_transport_stream_enqueue(struct vsock_sock *vsk,
1255 				struct msghdr *msg,
1256 				size_t len)
1257 {
1258 	struct virtio_vsock_pkt_info info = {
1259 		.op = VIRTIO_VSOCK_OP_RW,
1260 		.msg = msg,
1261 		.pkt_len = len,
1262 		.vsk = vsk,
1263 		.net = sock_net(sk_vsock(vsk)),
1264 	};
1265 
1266 	return virtio_transport_send_pkt_info(vsk, &info);
1267 }
1268 EXPORT_SYMBOL_GPL(virtio_transport_stream_enqueue);
1269 
virtio_transport_destruct(struct vsock_sock * vsk)1270 void virtio_transport_destruct(struct vsock_sock *vsk)
1271 {
1272 	struct virtio_vsock_sock *vvs = vsk->trans;
1273 
1274 	virtio_transport_cancel_close_work(vsk, true);
1275 
1276 	kfree(vvs);
1277 	vsk->trans = NULL;
1278 }
1279 EXPORT_SYMBOL_GPL(virtio_transport_destruct);
1280 
virtio_transport_unsent_bytes(struct vsock_sock * vsk)1281 ssize_t virtio_transport_unsent_bytes(struct vsock_sock *vsk)
1282 {
1283 	struct virtio_vsock_sock *vvs = vsk->trans;
1284 	size_t ret;
1285 
1286 	spin_lock_bh(&vvs->tx_lock);
1287 	ret = vvs->bytes_unsent;
1288 	spin_unlock_bh(&vvs->tx_lock);
1289 
1290 	return ret;
1291 }
1292 EXPORT_SYMBOL_GPL(virtio_transport_unsent_bytes);
1293 
virtio_transport_reset(struct vsock_sock * vsk,struct sk_buff * skb)1294 static int virtio_transport_reset(struct vsock_sock *vsk,
1295 				  struct sk_buff *skb)
1296 {
1297 	struct virtio_vsock_pkt_info info = {
1298 		.op = VIRTIO_VSOCK_OP_RST,
1299 		.reply = !!skb,
1300 		.vsk = vsk,
1301 		.net = sock_net(sk_vsock(vsk)),
1302 	};
1303 
1304 	/* Send RST only if the original pkt is not a RST pkt */
1305 	if (skb && le16_to_cpu(virtio_vsock_hdr(skb)->op) == VIRTIO_VSOCK_OP_RST)
1306 		return 0;
1307 
1308 	return virtio_transport_send_pkt_info(vsk, &info);
1309 }
1310 
1311 /* Normally packets are associated with a socket.  There may be no socket if an
1312  * attempt was made to connect to a socket that does not exist.
1313  *
1314  * net refers to the namespace of whoever sent the invalid message. For
1315  * loopback, this is the namespace of the socket. For vhost, this is the
1316  * namespace of the VM (i.e., vhost_vsock).
1317  */
virtio_transport_reset_no_sock(const struct virtio_transport * t,struct sk_buff * skb,struct net * net)1318 static int virtio_transport_reset_no_sock(const struct virtio_transport *t,
1319 					  struct sk_buff *skb, struct net *net)
1320 {
1321 	struct virtio_vsock_hdr *hdr = virtio_vsock_hdr(skb);
1322 	struct virtio_vsock_pkt_info info = {
1323 		.op = VIRTIO_VSOCK_OP_RST,
1324 		.type = le16_to_cpu(hdr->type),
1325 		.reply = true,
1326 
1327 		/* Set sk owner to socket we are replying to (may be NULL for
1328 		 * non-loopback). This keeps a reference to the sock and
1329 		 * sock_net(sk) until the reply skb is freed.
1330 		 */
1331 		.vsk = vsock_sk(skb->sk),
1332 
1333 		/* net is not defined here because we pass it directly to
1334 		 * t->send_pkt(), instead of relying on
1335 		 * virtio_transport_send_pkt_info() to pass it. It is not needed
1336 		 * by virtio_transport_alloc_skb().
1337 		 */
1338 	};
1339 	struct sk_buff *reply;
1340 
1341 	/* Send RST only if the original pkt is not a RST pkt */
1342 	if (le16_to_cpu(hdr->op) == VIRTIO_VSOCK_OP_RST)
1343 		return 0;
1344 
1345 	if (!t)
1346 		return -ENOTCONN;
1347 
1348 	reply = virtio_transport_alloc_skb(&info, 0, false, NULL,
1349 					   le64_to_cpu(hdr->dst_cid),
1350 					   le32_to_cpu(hdr->dst_port),
1351 					   le64_to_cpu(hdr->src_cid),
1352 					   le32_to_cpu(hdr->src_port));
1353 	if (!reply)
1354 		return -ENOMEM;
1355 
1356 	return t->send_pkt(reply, net);
1357 }
1358 
1359 /* This function should be called with sk_lock held and SOCK_DONE set */
virtio_transport_remove_sock(struct vsock_sock * vsk)1360 static void virtio_transport_remove_sock(struct vsock_sock *vsk)
1361 {
1362 	struct virtio_vsock_sock *vvs = vsk->trans;
1363 
1364 	/* We don't need to take rx_lock, as the socket is closing and we are
1365 	 * removing it.
1366 	 */
1367 	__skb_queue_purge(&vvs->rx_queue);
1368 	vsock_remove_sock(vsk);
1369 }
1370 
virtio_transport_cancel_close_work(struct vsock_sock * vsk,bool cancel_timeout)1371 static void virtio_transport_cancel_close_work(struct vsock_sock *vsk,
1372 					       bool cancel_timeout)
1373 {
1374 	struct sock *sk = sk_vsock(vsk);
1375 
1376 	if (vsk->close_work_scheduled &&
1377 	    (!cancel_timeout || cancel_delayed_work(&vsk->close_work))) {
1378 		vsk->close_work_scheduled = false;
1379 
1380 		virtio_transport_remove_sock(vsk);
1381 
1382 		/* Release refcnt obtained when we scheduled the timeout */
1383 		sock_put(sk);
1384 	}
1385 }
1386 
virtio_transport_do_close(struct vsock_sock * vsk,bool cancel_timeout)1387 static void virtio_transport_do_close(struct vsock_sock *vsk,
1388 				      bool cancel_timeout)
1389 {
1390 	struct sock *sk = sk_vsock(vsk);
1391 
1392 	sock_set_flag(sk, SOCK_DONE);
1393 	WRITE_ONCE(vsk->peer_shutdown, SHUTDOWN_MASK);
1394 	if (vsock_stream_has_data(vsk) <= 0)
1395 		sk->sk_state = TCP_CLOSING;
1396 	sk->sk_state_change(sk);
1397 
1398 	virtio_transport_cancel_close_work(vsk, cancel_timeout);
1399 }
1400 
virtio_transport_close_timeout(struct work_struct * work)1401 static void virtio_transport_close_timeout(struct work_struct *work)
1402 {
1403 	struct vsock_sock *vsk =
1404 		container_of(work, struct vsock_sock, close_work.work);
1405 	struct sock *sk = sk_vsock(vsk);
1406 
1407 	sock_hold(sk);
1408 	lock_sock(sk);
1409 
1410 	if (!sock_flag(sk, SOCK_DONE)) {
1411 		(void)virtio_transport_reset(vsk, NULL);
1412 
1413 		virtio_transport_do_close(vsk, false);
1414 	}
1415 
1416 	vsk->close_work_scheduled = false;
1417 
1418 	release_sock(sk);
1419 	sock_put(sk);
1420 }
1421 
1422 /* User context, vsk->sk is locked */
virtio_transport_close(struct vsock_sock * vsk)1423 static bool virtio_transport_close(struct vsock_sock *vsk)
1424 {
1425 	struct sock *sk = &vsk->sk;
1426 
1427 	if (!(sk->sk_state == TCP_ESTABLISHED ||
1428 	      sk->sk_state == TCP_CLOSING))
1429 		return true;
1430 
1431 	/* Already received SHUTDOWN from peer, reply with RST */
1432 	if ((vsk->peer_shutdown & SHUTDOWN_MASK) == SHUTDOWN_MASK) {
1433 		(void)virtio_transport_reset(vsk, NULL);
1434 		return true;
1435 	}
1436 
1437 	if ((sk->sk_shutdown & SHUTDOWN_MASK) != SHUTDOWN_MASK)
1438 		(void)virtio_transport_shutdown(vsk, SHUTDOWN_MASK);
1439 
1440 	if (!(current->flags & PF_EXITING))
1441 		vsock_linger(sk);
1442 
1443 	if (sock_flag(sk, SOCK_DONE)) {
1444 		return true;
1445 	}
1446 
1447 	sock_hold(sk);
1448 	INIT_DELAYED_WORK(&vsk->close_work,
1449 			  virtio_transport_close_timeout);
1450 	vsk->close_work_scheduled = true;
1451 	schedule_delayed_work(&vsk->close_work, VSOCK_CLOSE_TIMEOUT);
1452 	return false;
1453 }
1454 
virtio_transport_release(struct vsock_sock * vsk)1455 void virtio_transport_release(struct vsock_sock *vsk)
1456 {
1457 	struct sock *sk = &vsk->sk;
1458 	bool remove_sock = true;
1459 
1460 	if (sk->sk_type == SOCK_STREAM || sk->sk_type == SOCK_SEQPACKET)
1461 		remove_sock = virtio_transport_close(vsk);
1462 
1463 	if (remove_sock) {
1464 		sock_set_flag(sk, SOCK_DONE);
1465 		virtio_transport_remove_sock(vsk);
1466 	}
1467 }
1468 EXPORT_SYMBOL_GPL(virtio_transport_release);
1469 
1470 static int
virtio_transport_recv_connecting(struct sock * sk,struct sk_buff * skb)1471 virtio_transport_recv_connecting(struct sock *sk,
1472 				 struct sk_buff *skb)
1473 {
1474 	struct virtio_vsock_hdr *hdr = virtio_vsock_hdr(skb);
1475 	struct vsock_sock *vsk = vsock_sk(sk);
1476 	int skerr;
1477 	int err;
1478 
1479 	switch (le16_to_cpu(hdr->op)) {
1480 	case VIRTIO_VSOCK_OP_RESPONSE:
1481 		sk->sk_state = TCP_ESTABLISHED;
1482 		sk->sk_socket->state = SS_CONNECTED;
1483 		vsock_insert_connected(vsk);
1484 		sk->sk_state_change(sk);
1485 		break;
1486 	case VIRTIO_VSOCK_OP_INVALID:
1487 		break;
1488 	case VIRTIO_VSOCK_OP_RST:
1489 		skerr = ECONNRESET;
1490 		err = 0;
1491 		goto destroy;
1492 	default:
1493 		skerr = EPROTO;
1494 		err = -EINVAL;
1495 		goto destroy;
1496 	}
1497 	return 0;
1498 
1499 destroy:
1500 	virtio_transport_reset(vsk, skb);
1501 	sk->sk_state = TCP_CLOSE;
1502 	sk->sk_err = skerr;
1503 	sk_error_report(sk);
1504 	return err;
1505 }
1506 
1507 static bool
virtio_transport_recv_enqueue(struct vsock_sock * vsk,struct sk_buff * skb)1508 virtio_transport_recv_enqueue(struct vsock_sock *vsk,
1509 			      struct sk_buff *skb)
1510 {
1511 	struct virtio_vsock_sock *vvs = vsk->trans;
1512 	bool can_enqueue, free_pkt = false;
1513 	u32 len, queue_max, queue_len;
1514 	struct virtio_vsock_hdr *hdr;
1515 
1516 	hdr = virtio_vsock_hdr(skb);
1517 	len = le32_to_cpu(hdr->len);
1518 
1519 	/* virtio_transport_inc_rx_pkt() rejects packets when the per-skb
1520 	 * overhead (skb_queue_len * SKB_TRUESIZE(0)) exceeds buf_alloc.
1521 	 * Proactively collapse the queue before that happens.
1522 	 * No rx_lock needed: lock_sock is held by caller, preventing
1523 	 * concurrent enqueue or dequeue.
1524 	 */
1525 	queue_max = vvs->buf_alloc / SKB_TRUESIZE(0);
1526 	queue_len = skb_queue_len(&vvs->rx_queue);
1527 	if (queue_len >= queue_max) {
1528 		/* Walking a large queue may take a significant amount of time
1529 		 * and cache misses, causing traffic burstiness. Limit the
1530 		 * collapse to freeing room for this packet and the next one.
1531 		 * It may free more to fill each collapsed skb to capacity.
1532 		 */
1533 		virtio_transport_collapse_rx_queue(vvs, queue_len + 2 - queue_max);
1534 	}
1535 
1536 	spin_lock_bh(&vvs->rx_lock);
1537 
1538 	can_enqueue = virtio_transport_inc_rx_pkt(vvs, len);
1539 	if (!can_enqueue)
1540 		goto out;
1541 
1542 	if (le32_to_cpu(hdr->flags) & VIRTIO_VSOCK_SEQ_EOM)
1543 		vvs->msg_count++;
1544 
1545 	/* Try to copy small packets into the buffer of last packet queued,
1546 	 * to avoid wasting memory queueing the entire buffer with a small
1547 	 * payload. Skip non-linear (e.g. zerocopy) skbs; these carry payload
1548 	 * in skb_shinfo.
1549 	 */
1550 	if (len <= GOOD_COPY_LEN && !skb_queue_empty(&vvs->rx_queue) &&
1551 	    !skb_is_nonlinear(skb)) {
1552 		struct virtio_vsock_hdr *last_hdr;
1553 		struct sk_buff *last_skb;
1554 
1555 		last_skb = skb_peek_tail(&vvs->rx_queue);
1556 		last_hdr = virtio_vsock_hdr(last_skb);
1557 
1558 		/* If there is space in the last packet queued, we copy the
1559 		 * new packet in its buffer. We avoid this if the last packet
1560 		 * queued has VIRTIO_VSOCK_SEQ_EOM set, because this is
1561 		 * delimiter of SEQPACKET message, so 'pkt' is the first packet
1562 		 * of a new message.
1563 		 */
1564 		if (skb->len < skb_tailroom(last_skb) &&
1565 		    !(le32_to_cpu(last_hdr->flags) & VIRTIO_VSOCK_SEQ_EOM)) {
1566 			memcpy(skb_put(last_skb, skb->len), skb->data, skb->len);
1567 			free_pkt = true;
1568 			last_hdr->flags |= hdr->flags;
1569 			le32_add_cpu(&last_hdr->len, len);
1570 			goto out;
1571 		}
1572 	}
1573 
1574 	__skb_queue_tail(&vvs->rx_queue, skb);
1575 
1576 out:
1577 	spin_unlock_bh(&vvs->rx_lock);
1578 	if (free_pkt)
1579 		kfree_skb(skb);
1580 
1581 	return can_enqueue;
1582 }
1583 
1584 static int
virtio_transport_recv_connected(struct sock * sk,struct sk_buff * skb)1585 virtio_transport_recv_connected(struct sock *sk,
1586 				struct sk_buff *skb)
1587 {
1588 	struct virtio_vsock_hdr *hdr = virtio_vsock_hdr(skb);
1589 	struct vsock_sock *vsk = vsock_sk(sk);
1590 	int err = 0;
1591 
1592 	switch (le16_to_cpu(hdr->op)) {
1593 	case VIRTIO_VSOCK_OP_RW:
1594 		if (!virtio_transport_recv_enqueue(vsk, skb)) {
1595 			/* There is no more space to queue the packet, so let's
1596 			 * close the connection; otherwise, we'll lose data.
1597 			 */
1598 			(void)virtio_transport_reset(vsk, skb);
1599 			virtio_transport_do_close(vsk, true);
1600 			sk->sk_err = ENOBUFS;
1601 			sk_error_report(sk);
1602 			vsock_remove_sock(vsk);
1603 			break;
1604 		}
1605 		vsock_data_ready(sk);
1606 		return err;
1607 	case VIRTIO_VSOCK_OP_CREDIT_REQUEST:
1608 		virtio_transport_send_credit_update(vsk);
1609 		break;
1610 	case VIRTIO_VSOCK_OP_CREDIT_UPDATE:
1611 		sk->sk_write_space(sk);
1612 		break;
1613 	case VIRTIO_VSOCK_OP_SHUTDOWN: {
1614 		u32 peer_shutdown = READ_ONCE(vsk->peer_shutdown);
1615 
1616 		if (le32_to_cpu(hdr->flags) & VIRTIO_VSOCK_SHUTDOWN_RCV)
1617 			peer_shutdown |= RCV_SHUTDOWN;
1618 		if (le32_to_cpu(hdr->flags) & VIRTIO_VSOCK_SHUTDOWN_SEND)
1619 			peer_shutdown |= SEND_SHUTDOWN;
1620 		WRITE_ONCE(vsk->peer_shutdown, peer_shutdown);
1621 		if (peer_shutdown == SHUTDOWN_MASK) {
1622 			if (vsock_stream_has_data(vsk) <= 0 && !sock_flag(sk, SOCK_DONE)) {
1623 				(void)virtio_transport_reset(vsk, NULL);
1624 				virtio_transport_do_close(vsk, true);
1625 			}
1626 			/* Remove this socket anyway because the remote peer sent
1627 			 * the shutdown. This way a new connection will succeed
1628 			 * if the remote peer uses the same source port,
1629 			 * even if the old socket is still unreleased, but now disconnected.
1630 			 */
1631 			vsock_remove_sock(vsk);
1632 		}
1633 		if (le32_to_cpu(virtio_vsock_hdr(skb)->flags))
1634 			sk->sk_state_change(sk);
1635 		break;
1636 	}
1637 	case VIRTIO_VSOCK_OP_RST:
1638 		virtio_transport_do_close(vsk, true);
1639 		break;
1640 	default:
1641 		err = -EINVAL;
1642 		break;
1643 	}
1644 
1645 	kfree_skb(skb);
1646 	return err;
1647 }
1648 
1649 static void
virtio_transport_recv_disconnecting(struct sock * sk,struct sk_buff * skb)1650 virtio_transport_recv_disconnecting(struct sock *sk,
1651 				    struct sk_buff *skb)
1652 {
1653 	struct virtio_vsock_hdr *hdr = virtio_vsock_hdr(skb);
1654 	struct vsock_sock *vsk = vsock_sk(sk);
1655 
1656 	if (le16_to_cpu(hdr->op) == VIRTIO_VSOCK_OP_RST)
1657 		virtio_transport_do_close(vsk, true);
1658 }
1659 
1660 static int
virtio_transport_send_response(struct vsock_sock * vsk,struct sk_buff * skb)1661 virtio_transport_send_response(struct vsock_sock *vsk,
1662 			       struct sk_buff *skb)
1663 {
1664 	struct virtio_vsock_hdr *hdr = virtio_vsock_hdr(skb);
1665 	struct virtio_vsock_pkt_info info = {
1666 		.op = VIRTIO_VSOCK_OP_RESPONSE,
1667 		.remote_cid = le64_to_cpu(hdr->src_cid),
1668 		.remote_port = le32_to_cpu(hdr->src_port),
1669 		.reply = true,
1670 		.vsk = vsk,
1671 		.net = sock_net(sk_vsock(vsk)),
1672 	};
1673 
1674 	return virtio_transport_send_pkt_info(vsk, &info);
1675 }
1676 
virtio_transport_space_update(struct sock * sk,struct sk_buff * skb)1677 static bool virtio_transport_space_update(struct sock *sk,
1678 					  struct sk_buff *skb)
1679 {
1680 	struct virtio_vsock_hdr *hdr = virtio_vsock_hdr(skb);
1681 	struct vsock_sock *vsk = vsock_sk(sk);
1682 	struct virtio_vsock_sock *vvs = vsk->trans;
1683 	bool space_available;
1684 
1685 	/* Listener sockets are not associated with any transport, so we are
1686 	 * not able to take the state to see if there is space available in the
1687 	 * remote peer, but since they are only used to receive requests, we
1688 	 * can assume that there is always space available in the other peer.
1689 	 */
1690 	if (!vvs)
1691 		return true;
1692 
1693 	/* buf_alloc and fwd_cnt is always included in the hdr */
1694 	spin_lock_bh(&vvs->tx_lock);
1695 	vvs->peer_buf_alloc = le32_to_cpu(hdr->buf_alloc);
1696 	vvs->peer_fwd_cnt = le32_to_cpu(hdr->fwd_cnt);
1697 	space_available = virtio_transport_has_space(vvs);
1698 	spin_unlock_bh(&vvs->tx_lock);
1699 	return space_available;
1700 }
1701 
1702 /* Handle server socket */
1703 static int
virtio_transport_recv_listen(struct sock * sk,struct sk_buff * skb,struct virtio_transport * t)1704 virtio_transport_recv_listen(struct sock *sk, struct sk_buff *skb,
1705 			     struct virtio_transport *t)
1706 {
1707 	struct virtio_vsock_hdr *hdr = virtio_vsock_hdr(skb);
1708 	struct vsock_sock *vsk = vsock_sk(sk);
1709 	struct vsock_sock *vchild;
1710 	struct sock *child;
1711 	int ret;
1712 
1713 	if (le16_to_cpu(hdr->op) != VIRTIO_VSOCK_OP_REQUEST) {
1714 		virtio_transport_reset_no_sock(t, skb, sock_net(sk));
1715 		return -EINVAL;
1716 	}
1717 
1718 	if (sk_acceptq_is_full(sk)) {
1719 		virtio_transport_reset_no_sock(t, skb, sock_net(sk));
1720 		return -ENOMEM;
1721 	}
1722 
1723 	/* __vsock_release() might have already flushed accept_queue.
1724 	 * Subsequent enqueues would lead to a memory leak.
1725 	 */
1726 	if (sk->sk_shutdown == SHUTDOWN_MASK) {
1727 		virtio_transport_reset_no_sock(t, skb, sock_net(sk));
1728 		return -ESHUTDOWN;
1729 	}
1730 
1731 	child = vsock_create_connected(sk);
1732 	if (!child) {
1733 		virtio_transport_reset_no_sock(t, skb, sock_net(sk));
1734 		return -ENOMEM;
1735 	}
1736 
1737 	lock_sock_nested(child, SINGLE_DEPTH_NESTING);
1738 
1739 	child->sk_state = TCP_ESTABLISHED;
1740 
1741 	vchild = vsock_sk(child);
1742 	vsock_addr_init(&vchild->local_addr, le64_to_cpu(hdr->dst_cid),
1743 			le32_to_cpu(hdr->dst_port));
1744 	vsock_addr_init(&vchild->remote_addr, le64_to_cpu(hdr->src_cid),
1745 			le32_to_cpu(hdr->src_port));
1746 
1747 	ret = vsock_assign_transport(vchild, vsk);
1748 	/* Transport assigned (looking at remote_addr) must be the same
1749 	 * where we received the request.
1750 	 */
1751 	if (ret || vchild->transport != &t->transport) {
1752 		release_sock(child);
1753 		virtio_transport_reset_no_sock(t, skb, sock_net(sk));
1754 		sock_put(child);
1755 		return ret;
1756 	}
1757 
1758 	if (virtio_transport_space_update(child, skb))
1759 		child->sk_write_space(child);
1760 
1761 	vsock_insert_connected(vchild);
1762 	vsock_enqueue_accept(sk, child);
1763 	virtio_transport_send_response(vchild, skb);
1764 
1765 	release_sock(child);
1766 
1767 	sk->sk_data_ready(sk);
1768 	return 0;
1769 }
1770 
virtio_transport_valid_type(u16 type)1771 static bool virtio_transport_valid_type(u16 type)
1772 {
1773 	return (type == VIRTIO_VSOCK_TYPE_STREAM) ||
1774 	       (type == VIRTIO_VSOCK_TYPE_SEQPACKET);
1775 }
1776 
1777 /* We are under the virtio-vsock's vsock->rx_lock or vhost-vsock's vq->mutex
1778  * lock.
1779  */
virtio_transport_recv_pkt(struct virtio_transport * t,struct sk_buff * skb,struct net * net)1780 void virtio_transport_recv_pkt(struct virtio_transport *t,
1781 			       struct sk_buff *skb, struct net *net)
1782 {
1783 	struct virtio_vsock_hdr *hdr = virtio_vsock_hdr(skb);
1784 	struct sockaddr_vm src, dst;
1785 	struct vsock_sock *vsk;
1786 	struct sock *sk;
1787 	bool space_available;
1788 
1789 	vsock_addr_init(&src, le64_to_cpu(hdr->src_cid),
1790 			le32_to_cpu(hdr->src_port));
1791 	vsock_addr_init(&dst, le64_to_cpu(hdr->dst_cid),
1792 			le32_to_cpu(hdr->dst_port));
1793 
1794 	trace_virtio_transport_recv_pkt(src.svm_cid, src.svm_port,
1795 					dst.svm_cid, dst.svm_port,
1796 					le32_to_cpu(hdr->len),
1797 					le16_to_cpu(hdr->type),
1798 					le16_to_cpu(hdr->op),
1799 					le32_to_cpu(hdr->flags),
1800 					le32_to_cpu(hdr->buf_alloc),
1801 					le32_to_cpu(hdr->fwd_cnt));
1802 
1803 	if (!virtio_transport_valid_type(le16_to_cpu(hdr->type))) {
1804 		(void)virtio_transport_reset_no_sock(t, skb, net);
1805 		goto free_pkt;
1806 	}
1807 
1808 	/* The socket must be in connected or bound table
1809 	 * otherwise send reset back
1810 	 */
1811 	sk = vsock_find_connected_socket_net(&src, &dst, net);
1812 	if (!sk) {
1813 		sk = vsock_find_bound_socket_net(&dst, net);
1814 		if (!sk) {
1815 			(void)virtio_transport_reset_no_sock(t, skb, net);
1816 			goto free_pkt;
1817 		}
1818 	}
1819 
1820 	if (virtio_transport_get_type(sk) != le16_to_cpu(hdr->type)) {
1821 		(void)virtio_transport_reset_no_sock(t, skb, net);
1822 		sock_put(sk);
1823 		goto free_pkt;
1824 	}
1825 
1826 	if (!skb_set_owner_sk_safe(skb, sk)) {
1827 		WARN_ONCE(1, "receiving vsock socket has sk_refcnt == 0\n");
1828 		goto free_pkt;
1829 	}
1830 
1831 	vsk = vsock_sk(sk);
1832 
1833 	lock_sock(sk);
1834 
1835 	/* Check if sk has been closed or assigned to another transport before
1836 	 * lock_sock (note: listener sockets are not assigned to any transport)
1837 	 */
1838 	if (sock_flag(sk, SOCK_DONE) ||
1839 	    (sk->sk_state != TCP_LISTEN && vsk->transport != &t->transport)) {
1840 		(void)virtio_transport_reset_no_sock(t, skb, net);
1841 		release_sock(sk);
1842 		sock_put(sk);
1843 		goto free_pkt;
1844 	}
1845 
1846 	space_available = virtio_transport_space_update(sk, skb);
1847 
1848 	/* Update CID in case it has changed after a transport reset event */
1849 	if (vsk->local_addr.svm_cid != VMADDR_CID_ANY)
1850 		vsk->local_addr.svm_cid = dst.svm_cid;
1851 
1852 	if (space_available)
1853 		sk->sk_write_space(sk);
1854 
1855 	switch (sk->sk_state) {
1856 	case TCP_LISTEN:
1857 		virtio_transport_recv_listen(sk, skb, t);
1858 		kfree_skb(skb);
1859 		break;
1860 	case TCP_SYN_SENT:
1861 		virtio_transport_recv_connecting(sk, skb);
1862 		kfree_skb(skb);
1863 		break;
1864 	case TCP_ESTABLISHED:
1865 		virtio_transport_recv_connected(sk, skb);
1866 		break;
1867 	case TCP_CLOSING:
1868 		virtio_transport_recv_disconnecting(sk, skb);
1869 		kfree_skb(skb);
1870 		break;
1871 	default:
1872 		(void)virtio_transport_reset_no_sock(t, skb, net);
1873 		kfree_skb(skb);
1874 		break;
1875 	}
1876 
1877 	release_sock(sk);
1878 
1879 	/* Release refcnt obtained when we fetched this socket out of the
1880 	 * bound or connected list.
1881 	 */
1882 	sock_put(sk);
1883 	return;
1884 
1885 free_pkt:
1886 	kfree_skb(skb);
1887 }
1888 EXPORT_SYMBOL_GPL(virtio_transport_recv_pkt);
1889 
1890 /* Remove skbs found in a queue that have a vsk that matches.
1891  *
1892  * Each skb is freed.
1893  *
1894  * Returns the count of skbs that were reply packets.
1895  */
virtio_transport_purge_skbs(void * vsk,struct sk_buff_head * queue)1896 int virtio_transport_purge_skbs(void *vsk, struct sk_buff_head *queue)
1897 {
1898 	struct sk_buff_head freeme;
1899 	struct sk_buff *skb, *tmp;
1900 	int cnt = 0;
1901 
1902 	skb_queue_head_init(&freeme);
1903 
1904 	spin_lock_bh(&queue->lock);
1905 	skb_queue_walk_safe(queue, skb, tmp) {
1906 		if (vsock_sk(skb->sk) != vsk)
1907 			continue;
1908 
1909 		__skb_unlink(skb, queue);
1910 		__skb_queue_tail(&freeme, skb);
1911 
1912 		if (virtio_vsock_skb_reply(skb))
1913 			cnt++;
1914 	}
1915 	spin_unlock_bh(&queue->lock);
1916 
1917 	__skb_queue_purge(&freeme);
1918 
1919 	return cnt;
1920 }
1921 EXPORT_SYMBOL_GPL(virtio_transport_purge_skbs);
1922 
virtio_transport_read_skb(struct vsock_sock * vsk,skb_read_actor_t recv_actor)1923 int virtio_transport_read_skb(struct vsock_sock *vsk, skb_read_actor_t recv_actor)
1924 {
1925 	struct virtio_vsock_sock *vvs = vsk->trans;
1926 	struct sock *sk = sk_vsock(vsk);
1927 	struct virtio_vsock_hdr *hdr;
1928 	struct sk_buff *skb;
1929 	u32 pkt_len;
1930 	int off = 0;
1931 	int err;
1932 
1933 	spin_lock_bh(&vvs->rx_lock);
1934 	/* Use __skb_recv_datagram() for race-free handling of the receive. It
1935 	 * works for types other than dgrams.
1936 	 */
1937 	skb = __skb_recv_datagram(sk, &vvs->rx_queue, MSG_DONTWAIT, &off, &err);
1938 	if (!skb) {
1939 		spin_unlock_bh(&vvs->rx_lock);
1940 		return err;
1941 	}
1942 
1943 	hdr = virtio_vsock_hdr(skb);
1944 	if (le32_to_cpu(hdr->flags) & VIRTIO_VSOCK_SEQ_EOM)
1945 		vvs->msg_count--;
1946 
1947 	pkt_len = le32_to_cpu(hdr->len);
1948 	virtio_transport_dec_rx_pkt(vvs, pkt_len, pkt_len);
1949 	spin_unlock_bh(&vvs->rx_lock);
1950 
1951 	virtio_transport_send_credit_update(vsk);
1952 
1953 	return recv_actor(sk, skb);
1954 }
1955 EXPORT_SYMBOL_GPL(virtio_transport_read_skb);
1956 
virtio_transport_notify_set_rcvlowat(struct vsock_sock * vsk,int val)1957 int virtio_transport_notify_set_rcvlowat(struct vsock_sock *vsk, int val)
1958 {
1959 	struct virtio_vsock_sock *vvs = vsk->trans;
1960 	bool send_update;
1961 
1962 	spin_lock_bh(&vvs->rx_lock);
1963 
1964 	/* If number of available bytes is less than new SO_RCVLOWAT value,
1965 	 * kick sender to send more data, because sender may sleep in its
1966 	 * 'send()' syscall waiting for enough space at our side. Also
1967 	 * don't send credit update when peer already knows actual value -
1968 	 * such transmission will be useless.
1969 	 */
1970 	send_update = (vvs->rx_bytes < val) &&
1971 		      (vvs->fwd_cnt != vvs->last_fwd_cnt);
1972 
1973 	spin_unlock_bh(&vvs->rx_lock);
1974 
1975 	if (send_update) {
1976 		int err;
1977 
1978 		err = virtio_transport_send_credit_update(vsk);
1979 		if (err < 0)
1980 			return err;
1981 	}
1982 
1983 	return 0;
1984 }
1985 EXPORT_SYMBOL_GPL(virtio_transport_notify_set_rcvlowat);
1986 
1987 MODULE_LICENSE("GPL v2");
1988 MODULE_AUTHOR("Asias He");
1989 MODULE_DESCRIPTION("common code for virtio vsock");
1990