1 /* Copyright (c) 2018, Mellanox Technologies All rights reserved.
2 *
3 * This software is available to you under a choice of one of two
4 * licenses. You may choose to be licensed under the terms of the GNU
5 * General Public License (GPL) Version 2, available from the file
6 * COPYING in the main directory of this source tree, or the
7 * OpenIB.org BSD license below:
8 *
9 * Redistribution and use in source and binary forms, with or
10 * without modification, are permitted provided that the following
11 * conditions are met:
12 *
13 * - Redistributions of source code must retain the above
14 * copyright notice, this list of conditions and the following
15 * disclaimer.
16 *
17 * - Redistributions in binary form must reproduce the above
18 * copyright notice, this list of conditions and the following
19 * disclaimer in the documentation and/or other materials
20 * provided with the distribution.
21 *
22 * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND,
23 * EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF
24 * MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND
25 * NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS
26 * BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN
27 * ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN
28 * CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
29 * SOFTWARE.
30 */
31
32 #include <crypto/aead.h>
33 #include <linux/highmem.h>
34 #include <linux/module.h>
35 #include <linux/netdevice.h>
36 #include <net/dst.h>
37 #include <net/inet_connection_sock.h>
38 #include <net/tcp.h>
39 #include <net/tls.h>
40 #include <linux/skbuff_ref.h>
41
42 #include "tls.h"
43 #include "trace.h"
44
45 /* device_offload_lock is used to synchronize tls_dev_add
46 * against NETDEV_DOWN notifications.
47 */
48 static DECLARE_RWSEM(device_offload_lock);
49
50 static struct workqueue_struct *destruct_wq __read_mostly;
51
52 static LIST_HEAD(tls_device_list);
53 static LIST_HEAD(tls_device_down_list);
54 static DEFINE_SPINLOCK(tls_device_lock);
55
56 static struct page *dummy_page;
57
tls_device_free_ctx(struct tls_context * ctx)58 static void tls_device_free_ctx(struct tls_context *ctx)
59 {
60 if (ctx->tx_conf == TLS_HW)
61 kfree(tls_offload_ctx_tx(ctx));
62
63 if (ctx->rx_conf == TLS_HW)
64 kfree(tls_offload_ctx_rx(ctx));
65
66 tls_ctx_free(NULL, ctx);
67 }
68
tls_device_tx_del_task(struct work_struct * work)69 static void tls_device_tx_del_task(struct work_struct *work)
70 {
71 struct tls_offload_context_tx *offload_ctx =
72 container_of(work, struct tls_offload_context_tx, destruct_work);
73 struct tls_context *ctx = offload_ctx->ctx;
74 struct net_device *netdev;
75
76 /* Safe, because this is the destroy flow, refcount is 0, so
77 * tls_device_down can't store this field in parallel.
78 */
79 netdev = rcu_dereference_protected(ctx->netdev,
80 !refcount_read(&ctx->refcount));
81
82 netdev->tlsdev_ops->tls_dev_del(netdev, ctx, TLS_OFFLOAD_CTX_DIR_TX);
83 dev_put(netdev);
84 ctx->netdev = NULL;
85 tls_device_free_ctx(ctx);
86 }
87
tls_device_queue_ctx_destruction(struct tls_context * ctx)88 static void tls_device_queue_ctx_destruction(struct tls_context *ctx)
89 {
90 struct net_device *netdev;
91 unsigned long flags;
92 bool async_cleanup;
93
94 spin_lock_irqsave(&tls_device_lock, flags);
95 if (unlikely(!refcount_dec_and_test(&ctx->refcount))) {
96 spin_unlock_irqrestore(&tls_device_lock, flags);
97 return;
98 }
99
100 list_del(&ctx->list); /* Remove from tls_device_list / tls_device_down_list */
101
102 /* Safe, because this is the destroy flow, refcount is 0, so
103 * tls_device_down can't store this field in parallel.
104 */
105 netdev = rcu_dereference_protected(ctx->netdev,
106 !refcount_read(&ctx->refcount));
107
108 async_cleanup = netdev && ctx->tx_conf == TLS_HW;
109 if (async_cleanup) {
110 struct tls_offload_context_tx *offload_ctx = tls_offload_ctx_tx(ctx);
111
112 /* queue_work inside the spinlock
113 * to make sure tls_device_down waits for that work.
114 */
115 queue_work(destruct_wq, &offload_ctx->destruct_work);
116 }
117 spin_unlock_irqrestore(&tls_device_lock, flags);
118
119 if (!async_cleanup)
120 tls_device_free_ctx(ctx);
121 }
122
123 /* We assume that the socket is already connected */
get_netdev_for_sock(struct sock * sk)124 static struct net_device *get_netdev_for_sock(struct sock *sk)
125 {
126 struct net_device *dev, *lowest_dev = NULL;
127 struct dst_entry *dst;
128
129 rcu_read_lock();
130 dst = __sk_dst_get(sk);
131 dev = dst ? dst_dev_rcu(dst) : NULL;
132 if (likely(dev)) {
133 lowest_dev = netdev_sk_get_lowest_dev(dev, sk);
134 dev_hold(lowest_dev);
135 }
136 rcu_read_unlock();
137
138 return lowest_dev;
139 }
140
destroy_record(struct tls_record_info * record)141 static void destroy_record(struct tls_record_info *record)
142 {
143 int i;
144
145 for (i = 0; i < record->num_frags; i++)
146 __skb_frag_unref(&record->frags[i], false);
147 kfree(record);
148 }
149
delete_all_records(struct tls_offload_context_tx * offload_ctx)150 static void delete_all_records(struct tls_offload_context_tx *offload_ctx)
151 {
152 struct tls_record_info *info, *temp;
153
154 list_for_each_entry_safe(info, temp, &offload_ctx->records_list, list) {
155 list_del(&info->list);
156 destroy_record(info);
157 }
158
159 offload_ctx->retransmit_hint = NULL;
160 }
161
tls_tcp_clean_acked(struct sock * sk,u32 acked_seq)162 static void tls_tcp_clean_acked(struct sock *sk, u32 acked_seq)
163 {
164 struct tls_context *tls_ctx = tls_get_ctx(sk);
165 struct tls_record_info *info, *temp;
166 struct tls_offload_context_tx *ctx;
167 u64 deleted_records = 0;
168 unsigned long flags;
169
170 if (!tls_ctx)
171 return;
172
173 ctx = tls_offload_ctx_tx(tls_ctx);
174
175 spin_lock_irqsave(&ctx->lock, flags);
176 info = ctx->retransmit_hint;
177 if (info && !before(acked_seq, info->end_seq))
178 ctx->retransmit_hint = NULL;
179
180 list_for_each_entry_safe(info, temp, &ctx->records_list, list) {
181 if (before(acked_seq, info->end_seq))
182 break;
183 list_del(&info->list);
184
185 destroy_record(info);
186 deleted_records++;
187 }
188
189 ctx->unacked_record_sn += deleted_records;
190 spin_unlock_irqrestore(&ctx->lock, flags);
191 }
192
193 /* At this point, there should be no references on this
194 * socket and no in-flight SKBs associated with this
195 * socket, so it is safe to free all the resources.
196 */
tls_device_sk_destruct(struct sock * sk)197 void tls_device_sk_destruct(struct sock *sk)
198 {
199 struct tls_context *tls_ctx = tls_get_ctx(sk);
200 struct tls_offload_context_tx *ctx = tls_offload_ctx_tx(tls_ctx);
201
202 tls_ctx->sk_destruct(sk);
203
204 if (tls_ctx->tx_conf == TLS_HW) {
205 if (ctx->open_record)
206 destroy_record(ctx->open_record);
207 delete_all_records(ctx);
208 crypto_free_aead(ctx->aead_send);
209 clean_acked_data_disable(tcp_sk(sk));
210 }
211
212 tls_device_queue_ctx_destruction(tls_ctx);
213 }
214 EXPORT_SYMBOL_GPL(tls_device_sk_destruct);
215
tls_device_free_resources_tx(struct sock * sk)216 void tls_device_free_resources_tx(struct sock *sk)
217 {
218 struct tls_context *tls_ctx = tls_get_ctx(sk);
219
220 tls_free_partial_record(sk, tls_ctx);
221 }
222
tls_offload_tx_resync_request(struct sock * sk,u32 got_seq,u32 exp_seq)223 void tls_offload_tx_resync_request(struct sock *sk, u32 got_seq, u32 exp_seq)
224 {
225 struct tls_context *tls_ctx = tls_get_ctx(sk);
226
227 trace_tls_device_tx_resync_req(sk, got_seq, exp_seq);
228 WARN_ON(test_and_set_bit(TLS_TX_SYNC_SCHED, &tls_ctx->flags));
229 }
230 EXPORT_SYMBOL_GPL(tls_offload_tx_resync_request);
231
tls_device_resync_tx(struct sock * sk,struct tls_context * tls_ctx,u32 seq)232 static void tls_device_resync_tx(struct sock *sk, struct tls_context *tls_ctx,
233 u32 seq)
234 {
235 struct net_device *netdev;
236 int err = 0;
237 u8 *rcd_sn;
238
239 tcp_write_collapse_fence(sk);
240 rcd_sn = tls_ctx->tx.rec_seq;
241
242 trace_tls_device_tx_resync_send(sk, seq, rcd_sn);
243 down_read(&device_offload_lock);
244 netdev = rcu_dereference_protected(tls_ctx->netdev,
245 lockdep_is_held(&device_offload_lock));
246 if (netdev)
247 err = netdev->tlsdev_ops->tls_dev_resync(netdev, sk, seq,
248 rcd_sn,
249 TLS_OFFLOAD_CTX_DIR_TX);
250 up_read(&device_offload_lock);
251 if (err)
252 return;
253
254 clear_bit_unlock(TLS_TX_SYNC_SCHED, &tls_ctx->flags);
255 }
256
tls_append_frag(struct tls_record_info * record,struct page_frag * pfrag,int size)257 static void tls_append_frag(struct tls_record_info *record,
258 struct page_frag *pfrag,
259 int size)
260 {
261 skb_frag_t *frag;
262
263 frag = &record->frags[record->num_frags - 1];
264 if (skb_frag_page(frag) == pfrag->page &&
265 skb_frag_off(frag) + skb_frag_size(frag) == pfrag->offset) {
266 skb_frag_size_add(frag, size);
267 } else {
268 ++frag;
269 skb_frag_fill_page_desc(frag, pfrag->page, pfrag->offset,
270 size);
271 ++record->num_frags;
272 get_page(pfrag->page);
273 }
274
275 pfrag->offset += size;
276 record->len += size;
277 }
278
tls_push_record(struct sock * sk,struct tls_context * ctx,struct tls_offload_context_tx * offload_ctx,struct tls_record_info * record,int flags)279 static int tls_push_record(struct sock *sk,
280 struct tls_context *ctx,
281 struct tls_offload_context_tx *offload_ctx,
282 struct tls_record_info *record,
283 int flags)
284 {
285 struct tls_prot_info *prot = &ctx->prot_info;
286 struct tcp_sock *tp = tcp_sk(sk);
287 skb_frag_t *frag;
288 int i;
289
290 record->end_seq = tp->write_seq + record->len;
291 list_add_tail_rcu(&record->list, &offload_ctx->records_list);
292 offload_ctx->open_record = NULL;
293
294 if (test_bit(TLS_TX_SYNC_SCHED, &ctx->flags))
295 tls_device_resync_tx(sk, ctx, tp->write_seq);
296
297 tls_advance_record_sn(sk, prot, &ctx->tx);
298
299 for (i = 0; i < record->num_frags; i++) {
300 frag = &record->frags[i];
301 sg_unmark_end(&offload_ctx->sg_tx_data[i]);
302 sg_set_page(&offload_ctx->sg_tx_data[i], skb_frag_page(frag),
303 skb_frag_size(frag), skb_frag_off(frag));
304 sk_mem_charge(sk, skb_frag_size(frag));
305 get_page(skb_frag_page(frag));
306 }
307 sg_mark_end(&offload_ctx->sg_tx_data[record->num_frags - 1]);
308
309 /* all ready, send */
310 return tls_push_sg(sk, ctx, offload_ctx->sg_tx_data, 0, flags);
311 }
312
tls_device_record_close(struct sock * sk,struct tls_context * ctx,struct tls_record_info * record,struct page_frag * pfrag,unsigned char record_type)313 static void tls_device_record_close(struct sock *sk,
314 struct tls_context *ctx,
315 struct tls_record_info *record,
316 struct page_frag *pfrag,
317 unsigned char record_type)
318 {
319 struct tls_prot_info *prot = &ctx->prot_info;
320 struct page_frag dummy_tag_frag;
321
322 /* append tag
323 * device will fill in the tag, we just need to append a placeholder
324 * use socket memory to improve coalescing (re-using a single buffer
325 * increases frag count)
326 * if we can't allocate memory now use the dummy page
327 */
328 if (unlikely(pfrag->size - pfrag->offset < prot->tag_size) &&
329 !skb_page_frag_refill(prot->tag_size, pfrag, sk->sk_allocation)) {
330 dummy_tag_frag.page = dummy_page;
331 dummy_tag_frag.offset = 0;
332 pfrag = &dummy_tag_frag;
333 }
334 tls_append_frag(record, pfrag, prot->tag_size);
335
336 /* fill prepend */
337 tls_fill_prepend(ctx, skb_frag_address(&record->frags[0]),
338 record->len - prot->overhead_size,
339 record_type);
340 }
341
tls_create_new_record(struct tls_offload_context_tx * offload_ctx,struct page_frag * pfrag,size_t prepend_size)342 static int tls_create_new_record(struct tls_offload_context_tx *offload_ctx,
343 struct page_frag *pfrag,
344 size_t prepend_size)
345 {
346 struct tls_record_info *record;
347 skb_frag_t *frag;
348
349 record = kmalloc_obj(*record);
350 if (!record)
351 return -ENOMEM;
352
353 frag = &record->frags[0];
354 skb_frag_fill_page_desc(frag, pfrag->page, pfrag->offset,
355 prepend_size);
356
357 get_page(pfrag->page);
358 pfrag->offset += prepend_size;
359
360 record->num_frags = 1;
361 record->len = prepend_size;
362 offload_ctx->open_record = record;
363 return 0;
364 }
365
tls_do_allocation(struct sock * sk,struct tls_offload_context_tx * offload_ctx,struct page_frag * pfrag,size_t prepend_size)366 static int tls_do_allocation(struct sock *sk,
367 struct tls_offload_context_tx *offload_ctx,
368 struct page_frag *pfrag,
369 size_t prepend_size)
370 {
371 int ret;
372
373 if (!offload_ctx->open_record) {
374 if (unlikely(!skb_page_frag_refill(prepend_size, pfrag,
375 sk->sk_allocation))) {
376 if (!sk->sk_bypass_prot_mem)
377 READ_ONCE(sk->sk_prot)->enter_memory_pressure(sk);
378 sk_stream_moderate_sndbuf(sk);
379 return -ENOMEM;
380 }
381
382 ret = tls_create_new_record(offload_ctx, pfrag, prepend_size);
383 if (ret)
384 return ret;
385
386 if (pfrag->size > pfrag->offset)
387 return 0;
388 }
389
390 if (!sk_page_frag_refill(sk, pfrag))
391 return -ENOMEM;
392
393 return 0;
394 }
395
tls_device_copy_data(void * addr,size_t bytes,struct iov_iter * i)396 static int tls_device_copy_data(void *addr, size_t bytes, struct iov_iter *i)
397 {
398 size_t pre_copy, nocache;
399
400 pre_copy = ~((unsigned long)addr - 1) & (SMP_CACHE_BYTES - 1);
401 if (pre_copy) {
402 pre_copy = min(pre_copy, bytes);
403 if (copy_from_iter(addr, pre_copy, i) != pre_copy)
404 return -EFAULT;
405 bytes -= pre_copy;
406 addr += pre_copy;
407 }
408
409 nocache = round_down(bytes, SMP_CACHE_BYTES);
410 if (copy_from_iter_nocache(addr, nocache, i) != nocache)
411 return -EFAULT;
412 bytes -= nocache;
413 addr += nocache;
414
415 if (bytes && copy_from_iter(addr, bytes, i) != bytes)
416 return -EFAULT;
417
418 return 0;
419 }
420
tls_push_data(struct sock * sk,struct iov_iter * iter,size_t size,int flags,unsigned char record_type)421 static int tls_push_data(struct sock *sk,
422 struct iov_iter *iter,
423 size_t size, int flags,
424 unsigned char record_type)
425 {
426 struct tls_context *tls_ctx = tls_get_ctx(sk);
427 struct tls_prot_info *prot = &tls_ctx->prot_info;
428 struct tls_offload_context_tx *ctx = tls_offload_ctx_tx(tls_ctx);
429 struct tls_record_info *record;
430 int tls_push_record_flags;
431 struct page_frag *pfrag;
432 size_t orig_size = size;
433 u32 max_open_record_len;
434 bool more = false;
435 bool done = false;
436 int copy, rc = 0;
437 long timeo;
438
439 if (flags &
440 ~(MSG_MORE | MSG_DONTWAIT | MSG_NOSIGNAL |
441 MSG_SPLICE_PAGES | MSG_EOR))
442 return -EOPNOTSUPP;
443
444 if ((flags & (MSG_MORE | MSG_EOR)) == (MSG_MORE | MSG_EOR))
445 return -EINVAL;
446
447 if (unlikely(sk->sk_err))
448 return -sk->sk_err;
449
450 flags |= MSG_SENDPAGE_DECRYPTED;
451 tls_push_record_flags = flags | MSG_MORE;
452
453 timeo = sock_sndtimeo(sk, flags & MSG_DONTWAIT);
454 if (tls_is_partially_sent_record(tls_ctx)) {
455 rc = tls_push_partial_record(sk, tls_ctx, flags);
456 if (rc < 0)
457 return rc;
458 }
459
460 pfrag = sk_page_frag(sk);
461
462 /* TLS_HEADER_SIZE is not counted as part of the TLS record, and
463 * we need to leave room for an authentication tag.
464 */
465 max_open_record_len = tls_ctx->tx_max_payload_len +
466 prot->prepend_size;
467 do {
468 rc = tls_do_allocation(sk, ctx, pfrag, prot->prepend_size);
469 if (unlikely(rc)) {
470 rc = sk_stream_wait_memory(sk, &timeo);
471 if (!rc)
472 continue;
473
474 record = ctx->open_record;
475 if (!record)
476 break;
477 handle_error:
478 if (record_type != TLS_RECORD_TYPE_DATA) {
479 /* avoid sending partial
480 * record with type !=
481 * application_data
482 */
483 size = orig_size;
484 destroy_record(record);
485 ctx->open_record = NULL;
486 } else if (record->len > prot->prepend_size) {
487 goto last_record;
488 }
489
490 break;
491 }
492
493 record = ctx->open_record;
494
495 copy = min_t(size_t, size, max_open_record_len - record->len);
496 if (copy && (flags & MSG_SPLICE_PAGES)) {
497 struct page_frag zc_pfrag;
498 struct page **pages = &zc_pfrag.page;
499 size_t off;
500
501 rc = iov_iter_extract_pages(iter, &pages,
502 copy, 1, 0, &off);
503 if (rc <= 0) {
504 if (rc == 0)
505 rc = -EIO;
506 goto handle_error;
507 }
508 copy = rc;
509
510 if (WARN_ON_ONCE(!sendpage_ok(zc_pfrag.page))) {
511 iov_iter_revert(iter, copy);
512 rc = -EIO;
513 goto handle_error;
514 }
515
516 zc_pfrag.offset = off;
517 zc_pfrag.size = copy;
518 tls_append_frag(record, &zc_pfrag, copy);
519 } else if (copy) {
520 copy = min_t(size_t, copy, pfrag->size - pfrag->offset);
521
522 rc = tls_device_copy_data(page_address(pfrag->page) +
523 pfrag->offset, copy,
524 iter);
525 if (rc)
526 goto handle_error;
527 tls_append_frag(record, pfrag, copy);
528 }
529
530 size -= copy;
531 if (!size) {
532 last_record:
533 tls_push_record_flags = flags;
534 if (flags & MSG_MORE) {
535 more = true;
536 break;
537 }
538
539 done = true;
540 }
541
542 if (done || record->len >= max_open_record_len ||
543 (record->num_frags >= MAX_SKB_FRAGS - 1)) {
544 tls_device_record_close(sk, tls_ctx, record,
545 pfrag, record_type);
546
547 rc = tls_push_record(sk,
548 tls_ctx,
549 ctx,
550 record,
551 tls_push_record_flags);
552 if (rc < 0)
553 break;
554 }
555 } while (!done);
556
557 tls_ctx->pending_open_record_frags = more;
558
559 if (orig_size - size > 0)
560 rc = orig_size - size;
561
562 return rc;
563 }
564
tls_device_sendmsg(struct sock * sk,struct msghdr * msg,size_t size)565 int tls_device_sendmsg(struct sock *sk, struct msghdr *msg, size_t size)
566 {
567 unsigned char record_type = TLS_RECORD_TYPE_DATA;
568 struct tls_context *tls_ctx = tls_get_ctx(sk);
569 int rc;
570
571 if (!tls_ctx->zerocopy_sendfile)
572 msg->msg_flags &= ~MSG_SPLICE_PAGES;
573
574 mutex_lock(&tls_ctx->tx_lock);
575 lock_sock(sk);
576
577 if (unlikely(msg->msg_controllen)) {
578 rc = tls_process_cmsg(sk, msg, &record_type);
579 if (rc)
580 goto out;
581 }
582
583 rc = tls_push_data(sk, &msg->msg_iter, size, msg->msg_flags,
584 record_type);
585
586 out:
587 release_sock(sk);
588 mutex_unlock(&tls_ctx->tx_lock);
589 return rc;
590 }
591
tls_device_splice_eof(struct socket * sock)592 void tls_device_splice_eof(struct socket *sock)
593 {
594 struct sock *sk = sock->sk;
595 struct tls_context *tls_ctx = tls_get_ctx(sk);
596 struct iov_iter iter = {};
597
598 if (!tls_is_partially_sent_record(tls_ctx) &&
599 !tls_is_pending_open_record(tls_ctx))
600 return;
601
602 mutex_lock(&tls_ctx->tx_lock);
603 lock_sock(sk);
604
605 if (tls_is_partially_sent_record(tls_ctx) ||
606 tls_is_pending_open_record(tls_ctx)) {
607 iov_iter_bvec(&iter, ITER_SOURCE, NULL, 0, 0);
608 tls_push_data(sk, &iter, 0, 0, TLS_RECORD_TYPE_DATA);
609 }
610
611 release_sock(sk);
612 mutex_unlock(&tls_ctx->tx_lock);
613 }
614
tls_get_record(struct tls_offload_context_tx * context,u32 seq,u64 * p_record_sn)615 struct tls_record_info *tls_get_record(struct tls_offload_context_tx *context,
616 u32 seq, u64 *p_record_sn)
617 {
618 u64 record_sn = context->hint_record_sn;
619 struct tls_record_info *info, *last;
620
621 info = context->retransmit_hint;
622 if (!info ||
623 before(seq, info->end_seq - info->len)) {
624 /* if retransmit_hint is irrelevant start
625 * from the beginning of the list
626 */
627 info = list_first_entry_or_null(&context->records_list,
628 struct tls_record_info, list);
629 if (!info)
630 return NULL;
631 /* send the start_marker record if seq number is before the
632 * tls offload start marker sequence number. This record is
633 * required to handle TCP packets which are before TLS offload
634 * started.
635 * And if it's not start marker, look if this seq number
636 * belongs to the list.
637 */
638 if (likely(!tls_record_is_start_marker(info))) {
639 /* we have the first record, get the last record to see
640 * if this seq number belongs to the list.
641 */
642 last = list_last_entry(&context->records_list,
643 struct tls_record_info, list);
644
645 if (!between(seq, tls_record_start_seq(info),
646 last->end_seq))
647 return NULL;
648 }
649 record_sn = context->unacked_record_sn;
650 }
651
652 /* We just need the _rcu for the READ_ONCE() */
653 rcu_read_lock();
654 list_for_each_entry_from_rcu(info, &context->records_list, list) {
655 if (before(seq, info->end_seq)) {
656 if (!context->retransmit_hint ||
657 after(info->end_seq,
658 context->retransmit_hint->end_seq)) {
659 context->hint_record_sn = record_sn;
660 context->retransmit_hint = info;
661 }
662 *p_record_sn = record_sn;
663 goto exit_rcu_unlock;
664 }
665 record_sn++;
666 }
667 info = NULL;
668
669 exit_rcu_unlock:
670 rcu_read_unlock();
671 return info;
672 }
673 EXPORT_SYMBOL(tls_get_record);
674
tls_device_push_pending_record(struct sock * sk,int flags)675 static int tls_device_push_pending_record(struct sock *sk, int flags)
676 {
677 struct iov_iter iter;
678
679 iov_iter_kvec(&iter, ITER_SOURCE, NULL, 0, 0);
680 return tls_push_data(sk, &iter, 0, flags, TLS_RECORD_TYPE_DATA);
681 }
682
tls_device_write_space(struct sock * sk,struct tls_context * ctx)683 void tls_device_write_space(struct sock *sk, struct tls_context *ctx)
684 {
685 if (tls_is_partially_sent_record(ctx)) {
686 gfp_t sk_allocation = sk->sk_allocation;
687
688 WARN_ON_ONCE(sk->sk_write_pending);
689
690 sk->sk_allocation = GFP_ATOMIC;
691 tls_push_partial_record(sk, ctx,
692 MSG_DONTWAIT | MSG_NOSIGNAL |
693 MSG_SENDPAGE_DECRYPTED);
694 sk->sk_allocation = sk_allocation;
695 }
696 }
697
tls_device_resync_rx(struct tls_context * tls_ctx,struct sock * sk,u32 seq,u8 * rcd_sn)698 static void tls_device_resync_rx(struct tls_context *tls_ctx,
699 struct sock *sk, u32 seq, u8 *rcd_sn)
700 {
701 struct tls_offload_context_rx *rx_ctx = tls_offload_ctx_rx(tls_ctx);
702 struct net_device *netdev;
703
704 trace_tls_device_rx_resync_send(sk, seq, rcd_sn, rx_ctx->resync_type);
705 rcu_read_lock();
706 netdev = rcu_dereference(tls_ctx->netdev);
707 if (netdev)
708 netdev->tlsdev_ops->tls_dev_resync(netdev, sk, seq, rcd_sn,
709 TLS_OFFLOAD_CTX_DIR_RX);
710 rcu_read_unlock();
711 TLS_INC_STATS(sock_net(sk), LINUX_MIB_TLSRXDEVICERESYNC);
712 }
713
714 static bool
tls_device_rx_resync_async(struct tls_offload_resync_async * resync_async,s64 resync_req,u32 * seq,u16 * rcd_delta)715 tls_device_rx_resync_async(struct tls_offload_resync_async *resync_async,
716 s64 resync_req, u32 *seq, u16 *rcd_delta)
717 {
718 u32 is_async = resync_req & RESYNC_REQ_ASYNC;
719 u32 req_seq = resync_req >> 32;
720 u32 req_end = req_seq + ((resync_req >> 16) & 0xffff);
721 u16 i;
722
723 *rcd_delta = 0;
724
725 if (is_async) {
726 /* shouldn't get to wraparound:
727 * too long in async stage, something bad happened
728 */
729 if (WARN_ON_ONCE(resync_async->rcd_delta == USHRT_MAX)) {
730 tls_offload_rx_resync_async_request_cancel(resync_async);
731 return false;
732 }
733
734 /* asynchronous stage: log all headers seq such that
735 * req_seq <= seq <= end_seq, and wait for real resync request
736 */
737 if (before(*seq, req_seq))
738 return false;
739 if (!after(*seq, req_end) &&
740 resync_async->loglen < TLS_DEVICE_RESYNC_ASYNC_LOGMAX)
741 resync_async->log[resync_async->loglen++] = *seq;
742
743 resync_async->rcd_delta++;
744
745 return false;
746 }
747
748 /* synchronous stage: check against the logged entries and
749 * proceed to check the next entries if no match was found
750 */
751 for (i = 0; i < resync_async->loglen; i++)
752 if (req_seq == resync_async->log[i] &&
753 atomic64_try_cmpxchg(&resync_async->req, &resync_req, 0)) {
754 *rcd_delta = resync_async->rcd_delta - i;
755 *seq = req_seq;
756 resync_async->loglen = 0;
757 resync_async->rcd_delta = 0;
758 return true;
759 }
760
761 resync_async->loglen = 0;
762 resync_async->rcd_delta = 0;
763
764 if (req_seq == *seq &&
765 atomic64_try_cmpxchg(&resync_async->req,
766 &resync_req, 0))
767 return true;
768
769 return false;
770 }
771
tls_device_rx_resync_new_rec(struct sock * sk,u32 rcd_len,u32 seq)772 void tls_device_rx_resync_new_rec(struct sock *sk, u32 rcd_len, u32 seq)
773 {
774 struct tls_context *tls_ctx = tls_get_ctx(sk);
775 struct tls_offload_context_rx *rx_ctx;
776 u8 rcd_sn[TLS_MAX_REC_SEQ_SIZE];
777 u32 sock_data, is_req_pending;
778 struct tls_prot_info *prot;
779 s64 resync_req;
780 u16 rcd_delta;
781 u32 req_seq;
782
783 if (tls_ctx->rx_conf != TLS_HW)
784 return;
785 if (unlikely(test_bit(TLS_RX_DEV_DEGRADED, &tls_ctx->flags)))
786 return;
787
788 prot = &tls_ctx->prot_info;
789 rx_ctx = tls_offload_ctx_rx(tls_ctx);
790 memcpy(rcd_sn, tls_ctx->rx.rec_seq, prot->rec_seq_size);
791
792 switch (rx_ctx->resync_type) {
793 case TLS_OFFLOAD_SYNC_TYPE_DRIVER_REQ:
794 resync_req = atomic64_read(&rx_ctx->resync_req);
795 req_seq = resync_req >> 32;
796 seq += TLS_HEADER_SIZE - 1;
797 is_req_pending = resync_req;
798
799 if (likely(!is_req_pending) || req_seq != seq ||
800 !atomic64_try_cmpxchg(&rx_ctx->resync_req, &resync_req, 0))
801 return;
802 break;
803 case TLS_OFFLOAD_SYNC_TYPE_CORE_NEXT_HINT:
804 if (likely(!rx_ctx->resync_nh_do_now))
805 return;
806
807 /* head of next rec is already in, note that the sock_inq will
808 * include the currently parsed message when called from parser
809 */
810 sock_data = tcp_inq(sk);
811 if (sock_data > rcd_len) {
812 trace_tls_device_rx_resync_nh_delay(sk, sock_data,
813 rcd_len);
814 return;
815 }
816
817 rx_ctx->resync_nh_do_now = 0;
818 seq += rcd_len;
819 tls_bigint_increment(rcd_sn, prot->rec_seq_size);
820 break;
821 case TLS_OFFLOAD_SYNC_TYPE_DRIVER_REQ_ASYNC:
822 resync_req = atomic64_read(&rx_ctx->resync_async->req);
823 is_req_pending = resync_req;
824 if (likely(!is_req_pending))
825 return;
826
827 if (!tls_device_rx_resync_async(rx_ctx->resync_async,
828 resync_req, &seq, &rcd_delta))
829 return;
830 tls_bigint_subtract(rcd_sn, rcd_delta);
831 break;
832 }
833
834 tls_device_resync_rx(tls_ctx, sk, seq, rcd_sn);
835 }
836
tls_device_core_ctrl_rx_resync(struct tls_context * tls_ctx,struct tls_offload_context_rx * ctx,struct sock * sk,struct sk_buff * skb)837 static void tls_device_core_ctrl_rx_resync(struct tls_context *tls_ctx,
838 struct tls_offload_context_rx *ctx,
839 struct sock *sk, struct sk_buff *skb)
840 {
841 struct strp_msg *rxm;
842
843 /* device will request resyncs by itself based on stream scan */
844 if (ctx->resync_type != TLS_OFFLOAD_SYNC_TYPE_CORE_NEXT_HINT)
845 return;
846 /* already scheduled */
847 if (ctx->resync_nh_do_now)
848 return;
849 /* seen decrypted fragments since last fully-failed record */
850 if (ctx->resync_nh_reset) {
851 ctx->resync_nh_reset = 0;
852 ctx->resync_nh.decrypted_failed = 1;
853 ctx->resync_nh.decrypted_tgt = TLS_DEVICE_RESYNC_NH_START_IVAL;
854 return;
855 }
856
857 if (++ctx->resync_nh.decrypted_failed <= ctx->resync_nh.decrypted_tgt)
858 return;
859
860 /* doing resync, bump the next target in case it fails */
861 if (ctx->resync_nh.decrypted_tgt < TLS_DEVICE_RESYNC_NH_MAX_IVAL)
862 ctx->resync_nh.decrypted_tgt *= 2;
863 else
864 ctx->resync_nh.decrypted_tgt += TLS_DEVICE_RESYNC_NH_MAX_IVAL;
865
866 rxm = strp_msg(skb);
867
868 /* head of next rec is already in, parser will sync for us */
869 if (tcp_inq(sk) > rxm->full_len) {
870 trace_tls_device_rx_resync_nh_schedule(sk);
871 ctx->resync_nh_do_now = 1;
872 } else {
873 struct tls_prot_info *prot = &tls_ctx->prot_info;
874 u8 rcd_sn[TLS_MAX_REC_SEQ_SIZE];
875
876 memcpy(rcd_sn, tls_ctx->rx.rec_seq, prot->rec_seq_size);
877 tls_bigint_increment(rcd_sn, prot->rec_seq_size);
878
879 tls_device_resync_rx(tls_ctx, sk, tcp_sk(sk)->copied_seq,
880 rcd_sn);
881 }
882 }
883
884 static int
tls_device_reencrypt(struct sock * sk,struct tls_context * tls_ctx)885 tls_device_reencrypt(struct sock *sk, struct tls_context *tls_ctx)
886 {
887 struct tls_sw_context_rx *sw_ctx = tls_sw_ctx_rx(tls_ctx);
888 const struct tls_cipher_desc *cipher_desc;
889 int err, offset, copy, data_len, pos;
890 struct sk_buff *skb, *skb_iter;
891 struct scatterlist sg[1];
892 struct strp_msg *rxm;
893 char *orig_buf, *buf;
894
895 cipher_desc = get_cipher_desc(tls_ctx->crypto_recv.info.cipher_type);
896 DEBUG_NET_WARN_ON_ONCE(!cipher_desc || !cipher_desc->offloadable);
897
898 rxm = strp_msg(tls_strp_msg(sw_ctx));
899 orig_buf = kmalloc(rxm->full_len + TLS_HEADER_SIZE + cipher_desc->iv,
900 sk->sk_allocation);
901 if (!orig_buf)
902 return -ENOMEM;
903 buf = orig_buf;
904
905 err = tls_strp_msg_cow(sw_ctx);
906 if (unlikely(err))
907 goto free_buf;
908
909 skb = tls_strp_msg(sw_ctx);
910 rxm = strp_msg(skb);
911 offset = rxm->offset;
912
913 sg_init_table(sg, 1);
914 sg_set_buf(&sg[0], buf,
915 rxm->full_len + TLS_HEADER_SIZE + cipher_desc->iv);
916 err = skb_copy_bits(skb, offset, buf, TLS_HEADER_SIZE + cipher_desc->iv);
917 if (err)
918 goto free_buf;
919
920 /* We are interested only in the decrypted data not the auth */
921 err = decrypt_skb(sk, sg);
922 if (err != -EBADMSG)
923 goto free_buf;
924 else
925 err = 0;
926
927 data_len = rxm->full_len - cipher_desc->tag;
928
929 if (skb_pagelen(skb) > offset) {
930 copy = min_t(int, skb_pagelen(skb) - offset, data_len);
931
932 if (skb->decrypted) {
933 err = skb_store_bits(skb, offset, buf, copy);
934 if (err)
935 goto free_buf;
936 }
937
938 offset += copy;
939 buf += copy;
940 }
941
942 pos = skb_pagelen(skb);
943 skb_walk_frags(skb, skb_iter) {
944 int frag_pos;
945
946 /* Practically all frags must belong to msg if reencrypt
947 * is needed with current strparser and coalescing logic,
948 * but strparser may "get optimized", so let's be safe.
949 */
950 if (pos + skb_iter->len <= offset)
951 goto done_with_frag;
952 if (pos >= data_len + rxm->offset)
953 break;
954
955 frag_pos = offset - pos;
956 copy = min_t(int, skb_iter->len - frag_pos,
957 data_len + rxm->offset - offset);
958
959 if (skb_iter->decrypted) {
960 err = skb_store_bits(skb_iter, frag_pos, buf, copy);
961 if (err)
962 goto free_buf;
963 }
964
965 offset += copy;
966 buf += copy;
967 done_with_frag:
968 pos += skb_iter->len;
969 }
970
971 free_buf:
972 kfree(orig_buf);
973 return err;
974 }
975
tls_device_decrypted(struct sock * sk,struct tls_context * tls_ctx)976 int tls_device_decrypted(struct sock *sk, struct tls_context *tls_ctx)
977 {
978 struct tls_offload_context_rx *ctx = tls_offload_ctx_rx(tls_ctx);
979 struct tls_sw_context_rx *sw_ctx = tls_sw_ctx_rx(tls_ctx);
980 struct sk_buff *skb = tls_strp_msg(sw_ctx);
981 struct strp_msg *rxm = strp_msg(skb);
982 int is_decrypted, is_encrypted;
983
984 if (!tls_strp_msg_mixed_decrypted(sw_ctx)) {
985 is_decrypted = skb->decrypted;
986 is_encrypted = !is_decrypted;
987 } else {
988 is_decrypted = 0;
989 is_encrypted = 0;
990 }
991
992 trace_tls_device_decrypted(sk, tcp_sk(sk)->copied_seq - rxm->full_len,
993 tls_ctx->rx.rec_seq, rxm->full_len,
994 is_encrypted, is_decrypted);
995
996 if (unlikely(test_bit(TLS_RX_DEV_DEGRADED, &tls_ctx->flags))) {
997 if (likely(is_encrypted || is_decrypted))
998 return is_decrypted;
999
1000 /* After tls_device_down disables the offload, the next SKB will
1001 * likely have initial fragments decrypted, and final ones not
1002 * decrypted. We need to reencrypt that single SKB.
1003 */
1004 return tls_device_reencrypt(sk, tls_ctx);
1005 }
1006
1007 /* Return immediately if the record is either entirely plaintext or
1008 * entirely ciphertext. Otherwise handle reencrypt partially decrypted
1009 * record.
1010 */
1011 if (is_decrypted) {
1012 ctx->resync_nh_reset = 1;
1013 return is_decrypted;
1014 }
1015 if (is_encrypted) {
1016 tls_device_core_ctrl_rx_resync(tls_ctx, ctx, sk, skb);
1017 return 0;
1018 }
1019
1020 ctx->resync_nh_reset = 1;
1021 return tls_device_reencrypt(sk, tls_ctx);
1022 }
1023
tls_device_attach(struct tls_context * ctx,struct sock * sk,struct net_device * netdev)1024 static void tls_device_attach(struct tls_context *ctx, struct sock *sk,
1025 struct net_device *netdev)
1026 {
1027 if (sk->sk_destruct != tls_device_sk_destruct) {
1028 refcount_set(&ctx->refcount, 1);
1029 dev_hold(netdev);
1030 RCU_INIT_POINTER(ctx->netdev, netdev);
1031 spin_lock_irq(&tls_device_lock);
1032 list_add_tail(&ctx->list, &tls_device_list);
1033 spin_unlock_irq(&tls_device_lock);
1034
1035 ctx->sk_destruct = sk->sk_destruct;
1036 smp_store_release(&sk->sk_destruct, tls_device_sk_destruct);
1037 }
1038 }
1039
alloc_offload_ctx_tx(struct tls_context * ctx)1040 static struct tls_offload_context_tx *alloc_offload_ctx_tx(struct tls_context *ctx)
1041 {
1042 struct tls_offload_context_tx *offload_ctx;
1043 __be64 rcd_sn;
1044
1045 offload_ctx = kzalloc_obj(*offload_ctx);
1046 if (!offload_ctx)
1047 return NULL;
1048
1049 INIT_WORK(&offload_ctx->destruct_work, tls_device_tx_del_task);
1050 INIT_LIST_HEAD(&offload_ctx->records_list);
1051 spin_lock_init(&offload_ctx->lock);
1052 sg_init_table(offload_ctx->sg_tx_data,
1053 ARRAY_SIZE(offload_ctx->sg_tx_data));
1054
1055 /* start at rec_seq - 1 to account for the start marker record */
1056 memcpy(&rcd_sn, ctx->tx.rec_seq, sizeof(rcd_sn));
1057 offload_ctx->unacked_record_sn = be64_to_cpu(rcd_sn) - 1;
1058
1059 offload_ctx->ctx = ctx;
1060
1061 return offload_ctx;
1062 }
1063
tls_set_device_offload(struct sock * sk)1064 int tls_set_device_offload(struct sock *sk)
1065 {
1066 struct tls_record_info *start_marker_record;
1067 struct tls_offload_context_tx *offload_ctx;
1068 const struct tls_cipher_desc *cipher_desc;
1069 struct tls_crypto_info *crypto_info;
1070 struct tls_prot_info *prot;
1071 struct net_device *netdev;
1072 struct tls_context *ctx;
1073 char *iv, *rec_seq;
1074 int rc;
1075
1076 ctx = tls_get_ctx(sk);
1077 prot = &ctx->prot_info;
1078
1079 if (ctx->priv_ctx_tx)
1080 return -EEXIST;
1081
1082 netdev = get_netdev_for_sock(sk);
1083 if (!netdev) {
1084 pr_err_ratelimited("%s: netdev not found\n", __func__);
1085 return -EINVAL;
1086 }
1087
1088 if (!(netdev->features & NETIF_F_HW_TLS_TX)) {
1089 rc = -EOPNOTSUPP;
1090 goto release_netdev;
1091 }
1092
1093 crypto_info = &ctx->crypto_send.info;
1094 if (crypto_info->version != TLS_1_2_VERSION) {
1095 rc = -EOPNOTSUPP;
1096 goto release_netdev;
1097 }
1098
1099 cipher_desc = get_cipher_desc(crypto_info->cipher_type);
1100 if (!cipher_desc || !cipher_desc->offloadable) {
1101 rc = -EINVAL;
1102 goto release_netdev;
1103 }
1104
1105 rc = init_prot_info(prot, crypto_info, cipher_desc);
1106 if (rc)
1107 goto release_netdev;
1108
1109 iv = crypto_info_iv(crypto_info, cipher_desc);
1110 rec_seq = crypto_info_rec_seq(crypto_info, cipher_desc);
1111
1112 memcpy(ctx->tx.iv + cipher_desc->salt, iv, cipher_desc->iv);
1113 memcpy(ctx->tx.rec_seq, rec_seq, cipher_desc->rec_seq);
1114
1115 start_marker_record = kmalloc_obj(*start_marker_record);
1116 if (!start_marker_record) {
1117 rc = -ENOMEM;
1118 goto release_netdev;
1119 }
1120
1121 offload_ctx = alloc_offload_ctx_tx(ctx);
1122 if (!offload_ctx) {
1123 rc = -ENOMEM;
1124 goto free_marker_record;
1125 }
1126
1127 rc = tls_sw_fallback_init(sk, offload_ctx, crypto_info);
1128 if (rc)
1129 goto free_offload_ctx;
1130
1131 start_marker_record->end_seq = tcp_sk(sk)->write_seq;
1132 start_marker_record->len = 0;
1133 start_marker_record->num_frags = 0;
1134 list_add_tail(&start_marker_record->list, &offload_ctx->records_list);
1135
1136 clean_acked_data_enable(tcp_sk(sk), &tls_tcp_clean_acked);
1137 ctx->push_pending_record = tls_device_push_pending_record;
1138
1139 /* TLS offload is greatly simplified if we don't send
1140 * SKBs where only part of the payload needs to be encrypted.
1141 * So mark the last skb in the write queue as end of record.
1142 */
1143 tcp_write_collapse_fence(sk);
1144
1145 /* Avoid offloading if the device is down
1146 * We don't want to offload new flows after
1147 * the NETDEV_DOWN event
1148 *
1149 * device_offload_lock is taken in tls_devices's NETDEV_DOWN
1150 * handler thus protecting from the device going down before
1151 * ctx was added to tls_device_list.
1152 */
1153 down_read(&device_offload_lock);
1154 if (!(netdev->flags & IFF_UP)) {
1155 rc = -EINVAL;
1156 goto release_lock;
1157 }
1158
1159 ctx->priv_ctx_tx = offload_ctx;
1160 rc = netdev->tlsdev_ops->tls_dev_add(netdev, sk, TLS_OFFLOAD_CTX_DIR_TX,
1161 &ctx->crypto_send.info,
1162 tcp_sk(sk)->write_seq);
1163 trace_tls_device_offload_set(sk, TLS_OFFLOAD_CTX_DIR_TX,
1164 tcp_sk(sk)->write_seq, rec_seq, rc);
1165 if (rc)
1166 goto release_lock;
1167
1168 tls_device_attach(ctx, sk, netdev);
1169 up_read(&device_offload_lock);
1170
1171 /* following this assignment tls_is_skb_tx_device_offloaded
1172 * will return true and the context might be accessed
1173 * by the netdev's xmit function.
1174 */
1175 smp_store_release(&sk->sk_validate_xmit_skb, tls_validate_xmit_skb);
1176 dev_put(netdev);
1177
1178 return 0;
1179
1180 release_lock:
1181 up_read(&device_offload_lock);
1182 clean_acked_data_disable(tcp_sk(sk));
1183 crypto_free_aead(offload_ctx->aead_send);
1184 free_offload_ctx:
1185 kfree(offload_ctx);
1186 ctx->priv_ctx_tx = NULL;
1187 free_marker_record:
1188 kfree(start_marker_record);
1189 release_netdev:
1190 dev_put(netdev);
1191 return rc;
1192 }
1193
tls_set_device_offload_rx(struct sock * sk,struct tls_context * ctx)1194 int tls_set_device_offload_rx(struct sock *sk, struct tls_context *ctx)
1195 {
1196 struct tls12_crypto_info_aes_gcm_128 *info;
1197 struct tls_offload_context_rx *context;
1198 struct net_device *netdev;
1199 int rc = 0;
1200
1201 if (ctx->crypto_recv.info.version != TLS_1_2_VERSION)
1202 return -EOPNOTSUPP;
1203
1204 netdev = get_netdev_for_sock(sk);
1205 if (!netdev) {
1206 pr_err_ratelimited("%s: netdev not found\n", __func__);
1207 return -EINVAL;
1208 }
1209
1210 if (!(netdev->features & NETIF_F_HW_TLS_RX)) {
1211 rc = -EOPNOTSUPP;
1212 goto release_netdev;
1213 }
1214
1215 /* Avoid offloading if the device is down
1216 * We don't want to offload new flows after
1217 * the NETDEV_DOWN event
1218 *
1219 * device_offload_lock is taken in tls_devices's NETDEV_DOWN
1220 * handler thus protecting from the device going down before
1221 * ctx was added to tls_device_list.
1222 */
1223 down_read(&device_offload_lock);
1224 if (!(netdev->flags & IFF_UP)) {
1225 rc = -EINVAL;
1226 goto release_lock;
1227 }
1228
1229 context = kzalloc_obj(*context);
1230 if (!context) {
1231 rc = -ENOMEM;
1232 goto release_lock;
1233 }
1234 context->resync_nh_reset = 1;
1235
1236 ctx->priv_ctx_rx = context;
1237 rc = tls_set_sw_offload(sk, 0, NULL);
1238 if (rc)
1239 goto release_ctx;
1240
1241 rc = netdev->tlsdev_ops->tls_dev_add(netdev, sk, TLS_OFFLOAD_CTX_DIR_RX,
1242 &ctx->crypto_recv.info,
1243 tcp_sk(sk)->copied_seq);
1244 info = (void *)&ctx->crypto_recv.info;
1245 trace_tls_device_offload_set(sk, TLS_OFFLOAD_CTX_DIR_RX,
1246 tcp_sk(sk)->copied_seq, info->rec_seq, rc);
1247 if (rc)
1248 goto free_sw_resources;
1249
1250 tls_device_attach(ctx, sk, netdev);
1251 up_read(&device_offload_lock);
1252
1253 dev_put(netdev);
1254
1255 return 0;
1256
1257 free_sw_resources:
1258 up_read(&device_offload_lock);
1259 tls_sw_free_resources_rx(sk);
1260 down_read(&device_offload_lock);
1261 release_ctx:
1262 ctx->priv_ctx_rx = NULL;
1263 release_lock:
1264 up_read(&device_offload_lock);
1265 release_netdev:
1266 dev_put(netdev);
1267 return rc;
1268 }
1269
tls_device_offload_cleanup_rx(struct sock * sk)1270 void tls_device_offload_cleanup_rx(struct sock *sk)
1271 {
1272 struct tls_context *tls_ctx = tls_get_ctx(sk);
1273 struct net_device *netdev;
1274
1275 down_read(&device_offload_lock);
1276 netdev = rcu_dereference_protected(tls_ctx->netdev,
1277 lockdep_is_held(&device_offload_lock));
1278 if (!netdev)
1279 goto out;
1280
1281 netdev->tlsdev_ops->tls_dev_del(netdev, tls_ctx,
1282 TLS_OFFLOAD_CTX_DIR_RX);
1283
1284 if (tls_ctx->tx_conf != TLS_HW) {
1285 dev_put(netdev);
1286 rcu_assign_pointer(tls_ctx->netdev, NULL);
1287 } else {
1288 set_bit(TLS_RX_DEV_CLOSED, &tls_ctx->flags);
1289 }
1290 out:
1291 up_read(&device_offload_lock);
1292 tls_sw_release_resources_rx(sk);
1293 }
1294
tls_device_down(struct net_device * netdev)1295 static int tls_device_down(struct net_device *netdev)
1296 {
1297 struct tls_context *ctx, *tmp;
1298 unsigned long flags;
1299 LIST_HEAD(list);
1300
1301 /* Request a write lock to block new offload attempts */
1302 down_write(&device_offload_lock);
1303
1304 spin_lock_irqsave(&tls_device_lock, flags);
1305 list_for_each_entry_safe(ctx, tmp, &tls_device_list, list) {
1306 struct net_device *ctx_netdev =
1307 rcu_dereference_protected(ctx->netdev,
1308 lockdep_is_held(&device_offload_lock));
1309
1310 if (ctx_netdev != netdev ||
1311 !refcount_inc_not_zero(&ctx->refcount))
1312 continue;
1313
1314 list_move(&ctx->list, &list);
1315 }
1316 spin_unlock_irqrestore(&tls_device_lock, flags);
1317
1318 list_for_each_entry_safe(ctx, tmp, &list, list) {
1319 /* Stop offloaded TX and switch to the fallback.
1320 * tls_is_skb_tx_device_offloaded will return false.
1321 */
1322 WRITE_ONCE(ctx->sk->sk_validate_xmit_skb, tls_validate_xmit_skb_sw);
1323
1324 /* Stop the RX and TX resync.
1325 * tls_dev_resync must not be called after tls_dev_del.
1326 */
1327 rcu_assign_pointer(ctx->netdev, NULL);
1328
1329 /* Start skipping the RX resync logic completely. */
1330 set_bit(TLS_RX_DEV_DEGRADED, &ctx->flags);
1331
1332 /* Sync with inflight packets. After this point:
1333 * TX: no non-encrypted packets will be passed to the driver.
1334 * RX: resync requests from the driver will be ignored.
1335 */
1336 synchronize_net();
1337
1338 /* Release the offload context on the driver side. */
1339 if (ctx->tx_conf == TLS_HW)
1340 netdev->tlsdev_ops->tls_dev_del(netdev, ctx,
1341 TLS_OFFLOAD_CTX_DIR_TX);
1342 if (ctx->rx_conf == TLS_HW &&
1343 !test_bit(TLS_RX_DEV_CLOSED, &ctx->flags))
1344 netdev->tlsdev_ops->tls_dev_del(netdev, ctx,
1345 TLS_OFFLOAD_CTX_DIR_RX);
1346
1347 dev_put(netdev);
1348
1349 /* Move the context to a separate list for two reasons:
1350 * 1. When the context is deallocated, list_del is called.
1351 * 2. It's no longer an offloaded context, so we don't want to
1352 * run offload-specific code on this context.
1353 */
1354 spin_lock_irqsave(&tls_device_lock, flags);
1355 list_move_tail(&ctx->list, &tls_device_down_list);
1356 spin_unlock_irqrestore(&tls_device_lock, flags);
1357
1358 /* Device contexts for RX and TX will be freed in on sk_destruct
1359 * by tls_device_free_ctx. rx_conf and tx_conf stay in TLS_HW.
1360 * Now release the ref taken above.
1361 */
1362 if (refcount_dec_and_test(&ctx->refcount)) {
1363 /* sk_destruct ran after tls_device_down took a ref, and
1364 * it returned early. Complete the destruction here.
1365 */
1366 list_del(&ctx->list);
1367 tls_device_free_ctx(ctx);
1368 }
1369 }
1370
1371 up_write(&device_offload_lock);
1372
1373 flush_workqueue(destruct_wq);
1374
1375 return NOTIFY_DONE;
1376 }
1377
tls_dev_event(struct notifier_block * this,unsigned long event,void * ptr)1378 static int tls_dev_event(struct notifier_block *this, unsigned long event,
1379 void *ptr)
1380 {
1381 struct net_device *dev = netdev_notifier_info_to_dev(ptr);
1382
1383 if (!dev->tlsdev_ops &&
1384 !(dev->features & (NETIF_F_HW_TLS_RX | NETIF_F_HW_TLS_TX)))
1385 return NOTIFY_DONE;
1386
1387 switch (event) {
1388 case NETDEV_REGISTER:
1389 case NETDEV_FEAT_CHANGE:
1390 if (netif_is_bond_master(dev))
1391 return NOTIFY_DONE;
1392 if (!dev->tlsdev_ops ||
1393 !dev->tlsdev_ops->tls_dev_add ||
1394 !dev->tlsdev_ops->tls_dev_del)
1395 return NOTIFY_BAD;
1396 if ((dev->features & NETIF_F_HW_TLS_RX) &&
1397 !dev->tlsdev_ops->tls_dev_resync)
1398 return NOTIFY_BAD;
1399
1400 return NOTIFY_DONE;
1401 case NETDEV_DOWN:
1402 return tls_device_down(dev);
1403 }
1404 return NOTIFY_DONE;
1405 }
1406
1407 static struct notifier_block tls_dev_notifier = {
1408 .notifier_call = tls_dev_event,
1409 };
1410
tls_device_init(void)1411 int __init tls_device_init(void)
1412 {
1413 int err;
1414
1415 dummy_page = alloc_page(GFP_KERNEL);
1416 if (!dummy_page)
1417 return -ENOMEM;
1418
1419 destruct_wq = alloc_workqueue("ktls_device_destruct", WQ_PERCPU, 0);
1420 if (!destruct_wq) {
1421 err = -ENOMEM;
1422 goto err_free_dummy;
1423 }
1424
1425 err = register_netdevice_notifier(&tls_dev_notifier);
1426 if (err)
1427 goto err_destroy_wq;
1428
1429 return 0;
1430
1431 err_destroy_wq:
1432 destroy_workqueue(destruct_wq);
1433 err_free_dummy:
1434 put_page(dummy_page);
1435 return err;
1436 }
1437
tls_device_cleanup(void)1438 void __exit tls_device_cleanup(void)
1439 {
1440 unregister_netdevice_notifier(&tls_dev_notifier);
1441 destroy_workqueue(destruct_wq);
1442 clean_acked_data_flush();
1443 put_page(dummy_page);
1444 }
1445