From c6331af2fa0f01784221d8352df17a41cf0eb498 Mon Sep 17 00:00:00 2001 From: xtaci Date: Fri, 3 Jan 2020 15:48:50 +0800 Subject: [PATCH] Squashed commit of the following: commit 68a8b437de6c6f2c93e868d5e4dda7712807b5cf Author: xtaci Date: Fri Jan 3 15:48:21 2020 +0800 add comments commit 37e06370a13afacd7af5a28b97a0160d80d18579 Author: xtaci Date: Fri Jan 3 15:20:48 2020 +0800 add comments commit 7060ce363f93896e01cc2c3af2b79cbd55a5c7e1 Author: xtaci Date: Fri Jan 3 14:32:22 2020 +0800 update --- crypt.go | 69 +++++++++++++++++++++++++++++++------------------------- sess.go | 3 ++- 2 files changed, 40 insertions(+), 32 deletions(-) diff --git a/crypt.go b/crypt.go index 62aaa0d..99c0f1c 100644 --- a/crypt.go +++ b/crypt.go @@ -5,6 +5,7 @@ import ( "crypto/cipher" "crypto/des" "crypto/sha1" + "unsafe" xor "github.com/templexxx/xorsimd" "github.com/tjfoc/gmsm/sm4" @@ -57,7 +58,7 @@ func (c *salsa20BlockCrypt) Decrypt(dst, src []byte) { } type sm4BlockCrypt struct { - encbuf [sm4.BlockSize]byte + encbuf [sm4.BlockSize]byte // 64bit alignment enc/dec buffer decbuf [2 * sm4.BlockSize]byte block cipher.Block } @@ -260,69 +261,71 @@ func encrypt8(block cipher.Block, dst, src, buf []byte) { base := 0 repeat := n / 8 left := n % 8 + ptr_tbl := (*uint64)(unsafe.Pointer(&tbl[0])) + for i := 0; i < repeat; i++ { s := src[base:][0:64] d := dst[base:][0:64] // 1 - xor.Bytes8(d[0:8], s[0:8], tbl) + *(*uint64)(unsafe.Pointer(&d[0])) = *(*uint64)(unsafe.Pointer(&s[0])) ^ *ptr_tbl block.Encrypt(tbl, d[0:8]) // 2 - xor.Bytes8(d[8:16], s[8:16], tbl) + *(*uint64)(unsafe.Pointer(&d[8])) = *(*uint64)(unsafe.Pointer(&s[8])) ^ *ptr_tbl block.Encrypt(tbl, d[8:16]) // 3 - xor.Bytes8(d[16:24], s[16:24], tbl) + *(*uint64)(unsafe.Pointer(&d[16])) = *(*uint64)(unsafe.Pointer(&s[16])) ^ *ptr_tbl block.Encrypt(tbl, d[16:24]) // 4 - xor.Bytes8(d[24:32], s[24:32], tbl) + *(*uint64)(unsafe.Pointer(&d[24])) = *(*uint64)(unsafe.Pointer(&s[24])) ^ *ptr_tbl block.Encrypt(tbl, d[24:32]) // 5 - xor.Bytes8(d[32:40], s[32:40], tbl) + *(*uint64)(unsafe.Pointer(&d[32])) = *(*uint64)(unsafe.Pointer(&s[32])) ^ *ptr_tbl block.Encrypt(tbl, d[32:40]) // 6 - xor.Bytes8(d[40:48], s[40:48], tbl) + *(*uint64)(unsafe.Pointer(&d[40])) = *(*uint64)(unsafe.Pointer(&s[40])) ^ *ptr_tbl block.Encrypt(tbl, d[40:48]) // 7 - xor.Bytes8(d[48:56], s[48:56], tbl) + *(*uint64)(unsafe.Pointer(&d[48])) = *(*uint64)(unsafe.Pointer(&s[48])) ^ *ptr_tbl block.Encrypt(tbl, d[48:56]) // 8 - xor.Bytes8(d[56:64], s[56:64], tbl) + *(*uint64)(unsafe.Pointer(&d[56])) = *(*uint64)(unsafe.Pointer(&s[56])) ^ *ptr_tbl block.Encrypt(tbl, d[56:64]) base += 64 } switch left { case 7: - xor.Bytes8(dst[base:], src[base:], tbl) + *(*uint64)(unsafe.Pointer(&dst[base])) = *(*uint64)(unsafe.Pointer(&src[base])) ^ *ptr_tbl block.Encrypt(tbl, dst[base:]) base += 8 fallthrough case 6: - xor.Bytes8(dst[base:], src[base:], tbl) + *(*uint64)(unsafe.Pointer(&dst[base])) = *(*uint64)(unsafe.Pointer(&src[base])) ^ *ptr_tbl block.Encrypt(tbl, dst[base:]) base += 8 fallthrough case 5: - xor.Bytes8(dst[base:], src[base:], tbl) + *(*uint64)(unsafe.Pointer(&dst[base])) = *(*uint64)(unsafe.Pointer(&src[base])) ^ *ptr_tbl block.Encrypt(tbl, dst[base:]) base += 8 fallthrough case 4: - xor.Bytes8(dst[base:], src[base:], tbl) + *(*uint64)(unsafe.Pointer(&dst[base])) = *(*uint64)(unsafe.Pointer(&src[base])) ^ *ptr_tbl block.Encrypt(tbl, dst[base:]) base += 8 fallthrough case 3: - xor.Bytes8(dst[base:], src[base:], tbl) + *(*uint64)(unsafe.Pointer(&dst[base])) = *(*uint64)(unsafe.Pointer(&src[base])) ^ *ptr_tbl block.Encrypt(tbl, dst[base:]) base += 8 fallthrough case 2: - xor.Bytes8(dst[base:], src[base:], tbl) + *(*uint64)(unsafe.Pointer(&dst[base])) = *(*uint64)(unsafe.Pointer(&src[base])) ^ *ptr_tbl block.Encrypt(tbl, dst[base:]) base += 8 fallthrough case 1: - xor.Bytes8(dst[base:], src[base:], tbl) + *(*uint64)(unsafe.Pointer(&dst[base])) = *(*uint64)(unsafe.Pointer(&src[base])) ^ *ptr_tbl block.Encrypt(tbl, dst[base:]) base += 8 fallthrough @@ -513,6 +516,7 @@ func decrypt(block cipher.Block, dst, src, buf []byte) { } } +// decrypt 8 bytes block, all byte slices are supposed to be 64bit aligned func decrypt8(block cipher.Block, dst, src, buf []byte) { tbl := buf[0:8] next := buf[8:16] @@ -521,76 +525,79 @@ func decrypt8(block cipher.Block, dst, src, buf []byte) { base := 0 repeat := n / 8 left := n % 8 + ptr_tbl := (*uint64)(unsafe.Pointer(&tbl[0])) + ptr_next := (*uint64)(unsafe.Pointer(&next[0])) + for i := 0; i < repeat; i++ { s := src[base:][0:64] d := dst[base:][0:64] // 1 block.Encrypt(next, s[0:8]) - xor.Bytes8(d[0:8], s[0:8], tbl) + *(*uint64)(unsafe.Pointer(&d[0])) = *(*uint64)(unsafe.Pointer(&s[0])) ^ *ptr_tbl // 2 block.Encrypt(tbl, s[8:16]) - xor.Bytes8(d[8:16], s[8:16], next) + *(*uint64)(unsafe.Pointer(&d[8])) = *(*uint64)(unsafe.Pointer(&s[8])) ^ *ptr_next // 3 block.Encrypt(next, s[16:24]) - xor.Bytes8(d[16:24], s[16:24], tbl) + *(*uint64)(unsafe.Pointer(&d[16])) = *(*uint64)(unsafe.Pointer(&s[16])) ^ *ptr_tbl // 4 block.Encrypt(tbl, s[24:32]) - xor.Bytes8(d[24:32], s[24:32], next) + *(*uint64)(unsafe.Pointer(&d[24])) = *(*uint64)(unsafe.Pointer(&s[24])) ^ *ptr_next // 5 block.Encrypt(next, s[32:40]) - xor.Bytes8(d[32:40], s[32:40], tbl) + *(*uint64)(unsafe.Pointer(&d[32])) = *(*uint64)(unsafe.Pointer(&s[32])) ^ *ptr_tbl // 6 block.Encrypt(tbl, s[40:48]) - xor.Bytes8(d[40:48], s[40:48], next) + *(*uint64)(unsafe.Pointer(&d[40])) = *(*uint64)(unsafe.Pointer(&s[40])) ^ *ptr_next // 7 block.Encrypt(next, s[48:56]) - xor.Bytes8(d[48:56], s[48:56], tbl) + *(*uint64)(unsafe.Pointer(&d[48])) = *(*uint64)(unsafe.Pointer(&s[48])) ^ *ptr_tbl // 8 block.Encrypt(tbl, s[56:64]) - xor.Bytes8(d[56:64], s[56:64], next) + *(*uint64)(unsafe.Pointer(&d[56])) = *(*uint64)(unsafe.Pointer(&s[56])) ^ *ptr_next base += 64 } switch left { case 7: block.Encrypt(next, src[base:]) - xor.Bytes8(dst[base:], src[base:], tbl) + *(*uint64)(unsafe.Pointer(&dst[base])) = *(*uint64)(unsafe.Pointer(&src[base])) ^ *(*uint64)(unsafe.Pointer(&tbl[0])) tbl, next = next, tbl base += 8 fallthrough case 6: block.Encrypt(next, src[base:]) - xor.Bytes8(dst[base:], src[base:], tbl) + *(*uint64)(unsafe.Pointer(&dst[base])) = *(*uint64)(unsafe.Pointer(&src[base])) ^ *(*uint64)(unsafe.Pointer(&tbl[0])) tbl, next = next, tbl base += 8 fallthrough case 5: block.Encrypt(next, src[base:]) - xor.Bytes8(dst[base:], src[base:], tbl) + *(*uint64)(unsafe.Pointer(&dst[base])) = *(*uint64)(unsafe.Pointer(&src[base])) ^ *(*uint64)(unsafe.Pointer(&tbl[0])) tbl, next = next, tbl base += 8 fallthrough case 4: block.Encrypt(next, src[base:]) - xor.Bytes8(dst[base:], src[base:], tbl) + *(*uint64)(unsafe.Pointer(&dst[base])) = *(*uint64)(unsafe.Pointer(&src[base])) ^ *(*uint64)(unsafe.Pointer(&tbl[0])) tbl, next = next, tbl base += 8 fallthrough case 3: block.Encrypt(next, src[base:]) - xor.Bytes8(dst[base:], src[base:], tbl) + *(*uint64)(unsafe.Pointer(&dst[base])) = *(*uint64)(unsafe.Pointer(&src[base])) ^ *(*uint64)(unsafe.Pointer(&tbl[0])) tbl, next = next, tbl base += 8 fallthrough case 2: block.Encrypt(next, src[base:]) - xor.Bytes8(dst[base:], src[base:], tbl) + *(*uint64)(unsafe.Pointer(&dst[base])) = *(*uint64)(unsafe.Pointer(&src[base])) ^ *(*uint64)(unsafe.Pointer(&tbl[0])) tbl, next = next, tbl base += 8 fallthrough case 1: block.Encrypt(next, src[base:]) - xor.Bytes8(dst[base:], src[base:], tbl) + *(*uint64)(unsafe.Pointer(&dst[base])) = *(*uint64)(unsafe.Pointer(&src[base])) ^ *(*uint64)(unsafe.Pointer(&tbl[0])) tbl, next = next, tbl base += 8 fallthrough diff --git a/sess.go b/sess.go index 2384816..38b4c9e 100644 --- a/sess.go +++ b/sess.go @@ -49,7 +49,8 @@ var ( var ( // a system-wide packet buffer shared among sending, receiving and FEC - // to mitigate high-frequency memory allocation for packets + // to mitigate high-frequency memory allocation for packets, bytes from xmitBuf + // is aligned to 64bit xmitBuf sync.Pool )