mirror of
https://github.com/garethgeorge/backrest.git
synced 2026-09-24 17:05:45 +00:00
progress towards encryption support
This commit is contained in:
@@ -0,0 +1,107 @@
|
||||
package tunnel
|
||||
|
||||
import (
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"crypto/rand"
|
||||
"errors"
|
||||
"io"
|
||||
"sync"
|
||||
)
|
||||
|
||||
type crypt struct {
|
||||
secret []byte
|
||||
block cipher.Block
|
||||
once sync.Once
|
||||
err error
|
||||
}
|
||||
|
||||
func newCrypt(secret []byte) *crypt {
|
||||
return &crypt{secret: secret}
|
||||
}
|
||||
|
||||
func (c *crypt) init() {
|
||||
c.block, c.err = aes.NewCipher(c.secret)
|
||||
}
|
||||
|
||||
// pkcs5Padding adds PKCS#5 padding to the data
|
||||
func pkcs5Padding(data []byte, blockSize int) []byte {
|
||||
padding := blockSize - len(data)%blockSize
|
||||
padtext := make([]byte, padding)
|
||||
for i := range padtext {
|
||||
padtext[i] = byte(padding)
|
||||
}
|
||||
return append(data, padtext...)
|
||||
}
|
||||
|
||||
// pkcs5Unpadding removes PKCS#5 padding from the data
|
||||
func pkcs5Unpadding(data []byte) ([]byte, error) {
|
||||
if len(data) == 0 {
|
||||
return nil, errors.New("data is empty")
|
||||
}
|
||||
|
||||
padding := int(data[len(data)-1])
|
||||
if padding > len(data) || padding == 0 {
|
||||
return nil, errors.New("invalid padding")
|
||||
}
|
||||
|
||||
// Check that all padding bytes are correct
|
||||
for i := len(data) - padding; i < len(data); i++ {
|
||||
if data[i] != byte(padding) {
|
||||
return nil, errors.New("invalid padding")
|
||||
}
|
||||
}
|
||||
|
||||
return data[:len(data)-padding], nil
|
||||
}
|
||||
|
||||
func (c *crypt) Encrypt(data []byte) ([]byte, error) {
|
||||
c.once.Do(c.init)
|
||||
if c.err != nil {
|
||||
return nil, c.err
|
||||
}
|
||||
|
||||
// Add PKCS#5 padding
|
||||
paddedData := pkcs5Padding(data, aes.BlockSize)
|
||||
|
||||
// Create a random IV
|
||||
ciphertext := make([]byte, aes.BlockSize+len(paddedData))
|
||||
iv := ciphertext[:aes.BlockSize]
|
||||
if _, err := io.ReadFull(rand.Reader, iv); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Encrypt the data
|
||||
mode := cipher.NewCBCEncrypter(c.block, iv)
|
||||
mode.CryptBlocks(ciphertext[aes.BlockSize:], paddedData)
|
||||
|
||||
return ciphertext, nil
|
||||
}
|
||||
|
||||
func (c *crypt) Decrypt(data []byte) ([]byte, error) {
|
||||
if len(data) < aes.BlockSize {
|
||||
return nil, errors.New("ciphertext too short")
|
||||
}
|
||||
|
||||
c.once.Do(c.init)
|
||||
if c.err != nil {
|
||||
return nil, c.err
|
||||
}
|
||||
|
||||
// Extract IV and ciphertext
|
||||
iv := data[:aes.BlockSize]
|
||||
ciphertext := data[aes.BlockSize:]
|
||||
|
||||
// Check that ciphertext is a multiple of block size
|
||||
if len(ciphertext)%aes.BlockSize != 0 {
|
||||
return nil, errors.New("ciphertext is not a multiple of the block size")
|
||||
}
|
||||
|
||||
// Decrypt the data
|
||||
mode := cipher.NewCBCDecrypter(c.block, iv)
|
||||
plaintext := make([]byte, len(ciphertext))
|
||||
mode.CryptBlocks(plaintext, ciphertext)
|
||||
|
||||
// Remove PKCS#5 padding
|
||||
return pkcs5Unpadding(plaintext)
|
||||
}
|
||||
@@ -0,0 +1,338 @@
|
||||
package tunnel
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/rand"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestNewCrypt(t *testing.T) {
|
||||
secret := []byte("test-secret-key-")
|
||||
c := newCrypt(secret)
|
||||
|
||||
if c == nil {
|
||||
t.Fatal("newCrypt returned nil")
|
||||
}
|
||||
|
||||
if !bytes.Equal(c.secret, secret) {
|
||||
t.Errorf("expected secret %v, got %v", secret, c.secret)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPKCS5Padding(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
data []byte
|
||||
blockSize int
|
||||
expected int // expected length after padding
|
||||
}{
|
||||
{
|
||||
name: "empty data",
|
||||
data: []byte{},
|
||||
blockSize: 16,
|
||||
expected: 16,
|
||||
},
|
||||
{
|
||||
name: "data length equals block size",
|
||||
data: make([]byte, 16),
|
||||
blockSize: 16,
|
||||
expected: 32,
|
||||
},
|
||||
{
|
||||
name: "data length less than block size",
|
||||
data: []byte("hello"),
|
||||
blockSize: 16,
|
||||
expected: 16,
|
||||
},
|
||||
{
|
||||
name: "data length greater than block size",
|
||||
data: make([]byte, 20),
|
||||
blockSize: 16,
|
||||
expected: 32,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
padded := pkcs5Padding(tt.data, tt.blockSize)
|
||||
|
||||
if len(padded) != tt.expected {
|
||||
t.Errorf("expected length %d, got %d", tt.expected, len(padded))
|
||||
}
|
||||
|
||||
// Check that padding is correct
|
||||
paddingLen := tt.blockSize - (len(tt.data) % tt.blockSize)
|
||||
for i := len(tt.data); i < len(padded); i++ {
|
||||
if padded[i] != byte(paddingLen) {
|
||||
t.Errorf("invalid padding byte at position %d: expected %d, got %d", i, paddingLen, padded[i])
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPKCS5Unpadding(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
data []byte
|
||||
expectErr bool
|
||||
expected []byte
|
||||
}{
|
||||
{
|
||||
name: "valid padding",
|
||||
data: []byte{1, 2, 3, 4, 5, 5, 5, 5, 5},
|
||||
expectErr: false,
|
||||
expected: []byte{1, 2, 3, 4},
|
||||
},
|
||||
{
|
||||
name: "full block padding",
|
||||
data: []byte{16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16},
|
||||
expectErr: false,
|
||||
expected: []byte{},
|
||||
},
|
||||
{
|
||||
name: "empty data",
|
||||
data: []byte{},
|
||||
expectErr: true,
|
||||
expected: nil,
|
||||
},
|
||||
{
|
||||
name: "zero padding",
|
||||
data: []byte{1, 2, 3, 0},
|
||||
expectErr: true,
|
||||
expected: nil,
|
||||
},
|
||||
{
|
||||
name: "invalid padding length",
|
||||
data: []byte{1, 2, 3, 10},
|
||||
expectErr: true,
|
||||
expected: nil,
|
||||
},
|
||||
{
|
||||
name: "inconsistent padding",
|
||||
data: []byte{1, 2, 3, 3, 2},
|
||||
expectErr: true,
|
||||
expected: nil,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result, err := pkcs5Unpadding(tt.data)
|
||||
|
||||
if tt.expectErr {
|
||||
if err == nil {
|
||||
t.Error("expected error but got none")
|
||||
}
|
||||
} else {
|
||||
if err != nil {
|
||||
t.Errorf("unexpected error: %v", err)
|
||||
}
|
||||
if !bytes.Equal(result, tt.expected) {
|
||||
t.Errorf("expected %v, got %v", tt.expected, result)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestEncryptDecrypt(t *testing.T) {
|
||||
secret := []byte("my-secret-key...")
|
||||
c := newCrypt(secret)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
data []byte
|
||||
}{
|
||||
{
|
||||
name: "empty data",
|
||||
data: []byte{},
|
||||
},
|
||||
{
|
||||
name: "short data",
|
||||
data: []byte("hello"),
|
||||
},
|
||||
{
|
||||
name: "data equal to block size",
|
||||
data: make([]byte, 16),
|
||||
},
|
||||
{
|
||||
name: "data larger than block size",
|
||||
data: []byte("this is a longer message that spans multiple blocks"),
|
||||
},
|
||||
{
|
||||
name: "binary data",
|
||||
data: []byte{0, 1, 2, 3, 255, 254, 253, 252},
|
||||
},
|
||||
{
|
||||
name: "large data",
|
||||
data: make([]byte, 1024),
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
// Fill large data with random bytes
|
||||
if len(tt.data) == 1024 {
|
||||
rand.Read(tt.data)
|
||||
}
|
||||
|
||||
// Encrypt
|
||||
encrypted, err := c.Encrypt(tt.data)
|
||||
if err != nil {
|
||||
t.Fatalf("encryption failed: %v", err)
|
||||
}
|
||||
|
||||
// Verify encrypted data is different from original (unless original is empty)
|
||||
if len(tt.data) > 0 && bytes.Equal(encrypted, tt.data) {
|
||||
t.Error("encrypted data is same as original")
|
||||
}
|
||||
|
||||
// Verify encrypted data includes IV (at least block size)
|
||||
if len(encrypted) < 16 {
|
||||
t.Errorf("encrypted data too short: %d bytes", len(encrypted))
|
||||
}
|
||||
|
||||
// Decrypt
|
||||
decrypted, err := c.Decrypt(encrypted)
|
||||
if err != nil {
|
||||
t.Fatalf("decryption failed: %v", err)
|
||||
}
|
||||
|
||||
// Verify decrypted matches original
|
||||
if !bytes.Equal(decrypted, tt.data) {
|
||||
t.Errorf("decrypted data doesn't match original.\nOriginal: %v\nDecrypted: %v", tt.data, decrypted)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestEncryptionIsRandomized(t *testing.T) {
|
||||
secret := []byte("test-secret.....")
|
||||
c := newCrypt(secret)
|
||||
data := []byte("same data every time")
|
||||
|
||||
// Encrypt the same data multiple times
|
||||
encrypted1, err := c.Encrypt(data)
|
||||
if err != nil {
|
||||
t.Fatalf("encryption 1 failed: %v", err)
|
||||
}
|
||||
|
||||
encrypted2, err := c.Encrypt(data)
|
||||
if err != nil {
|
||||
t.Fatalf("encryption 2 failed: %v", err)
|
||||
}
|
||||
|
||||
// Encrypted results should be different due to random IV
|
||||
if bytes.Equal(encrypted1, encrypted2) {
|
||||
t.Error("encryption is not randomized - same input produced same output")
|
||||
}
|
||||
|
||||
// But both should decrypt to the same original data
|
||||
decrypted1, err := c.Decrypt(encrypted1)
|
||||
if err != nil {
|
||||
t.Fatalf("decryption 1 failed: %v", err)
|
||||
}
|
||||
|
||||
decrypted2, err := c.Decrypt(encrypted2)
|
||||
if err != nil {
|
||||
t.Fatalf("decryption 2 failed: %v", err)
|
||||
}
|
||||
|
||||
if !bytes.Equal(decrypted1, data) || !bytes.Equal(decrypted2, data) {
|
||||
t.Error("decryption failed to recover original data")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecryptInvalidData(t *testing.T) {
|
||||
secret := []byte("test-secret.....")
|
||||
c := newCrypt(secret)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
data []byte
|
||||
}{
|
||||
{
|
||||
name: "too short",
|
||||
data: []byte("short"),
|
||||
},
|
||||
{
|
||||
name: "not multiple of block size",
|
||||
data: make([]byte, 17), // 16 + 1
|
||||
},
|
||||
{
|
||||
name: "corrupted padding",
|
||||
data: func() []byte {
|
||||
// Create valid encrypted data then corrupt it
|
||||
original := []byte("test data")
|
||||
encrypted, _ := c.Encrypt(original)
|
||||
// Corrupt the last byte (padding)
|
||||
encrypted[len(encrypted)-1] = 255
|
||||
return encrypted
|
||||
}(),
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
_, err := c.Decrypt(tt.data)
|
||||
if err == nil {
|
||||
t.Error("expected decryption to fail but it succeeded")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDifferentSecretsCannotDecrypt(t *testing.T) {
|
||||
data := []byte("secret message")
|
||||
|
||||
c1 := newCrypt([]byte("secret1........."))
|
||||
c2 := newCrypt([]byte("secret2........."))
|
||||
|
||||
// Encrypt with first secret
|
||||
encrypted, err := c1.Encrypt(data)
|
||||
if err != nil {
|
||||
t.Fatalf("encryption failed: %v", err)
|
||||
}
|
||||
|
||||
// Try to decrypt with second secret
|
||||
_, err = c2.Decrypt(encrypted)
|
||||
if err == nil {
|
||||
t.Error("decryption with wrong secret should fail")
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkEncrypt(b *testing.B) {
|
||||
secret := []byte("benchmark-secret")
|
||||
c := newCrypt(secret)
|
||||
data := make([]byte, 1024)
|
||||
rand.Read(data)
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
_, err := c.Encrypt(data)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkDecrypt(b *testing.B) {
|
||||
secret := []byte("benchmark-secret")
|
||||
c := newCrypt(secret)
|
||||
data := make([]byte, 1024)
|
||||
rand.Read(data)
|
||||
|
||||
encrypted, err := c.Encrypt(data)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
_, err := c.Decrypt(encrypted)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1,6 +1,7 @@
|
||||
package tunnel
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"net"
|
||||
"os"
|
||||
@@ -8,7 +9,7 @@ import (
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
v1 "github.com/garethgeorge/backrest/gen/go/v1"
|
||||
"github.com/garethgeorge/backrest/gen/go/v1sync"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
@@ -35,15 +36,21 @@ type connState struct {
|
||||
|
||||
var _ net.Conn = (*connState)(nil)
|
||||
|
||||
func newConnState(stream stream, connId int64, logger *zap.Logger) *connState {
|
||||
func newConnState(streamm stream, connId int64, secret []byte, logger *zap.Logger) *connState {
|
||||
if logger != nil {
|
||||
logger = logger.Named("connState").With(
|
||||
zap.Int64("connId", connId))
|
||||
}
|
||||
|
||||
var s stream = streamm
|
||||
// if len(secret) > 0 {
|
||||
// s = newCryptedStream(streamm, secret)
|
||||
// }
|
||||
|
||||
return &connState{
|
||||
connId: connId,
|
||||
seqno: 0,
|
||||
stream: stream,
|
||||
stream: s,
|
||||
logger: logger,
|
||||
|
||||
closedCh: make(chan struct{}),
|
||||
@@ -65,7 +72,7 @@ func (c *connState) sendOpenPacket() error {
|
||||
c.logger.Info("sending open packet")
|
||||
}
|
||||
|
||||
return c.stream.Send(&v1.TunnelMessage{
|
||||
return c.stream.Send(&v1sync.TunnelMessage{
|
||||
ConnId: c.connId,
|
||||
Seqno: 0, // Open packet has Seqno 0
|
||||
})
|
||||
@@ -82,9 +89,9 @@ func (c *connState) Write(data []byte) (int, error) {
|
||||
if c.logger != nil {
|
||||
c.logger.Debug("writing data", zap.Int("dataLength", len(data)), zap.Int64("seqno", c.seqno))
|
||||
}
|
||||
err := c.stream.Send(&v1.TunnelMessage{
|
||||
err := c.stream.Send(&v1sync.TunnelMessage{
|
||||
ConnId: c.connId,
|
||||
Data: data,
|
||||
Data: bytes.Clone(data),
|
||||
Seqno: c.nextWriteSeqno,
|
||||
})
|
||||
if err != nil {
|
||||
@@ -142,7 +149,7 @@ func (c *connState) Close() error {
|
||||
c.logger.Info("closing connection")
|
||||
}
|
||||
close(c.closedCh)
|
||||
if err := c.stream.Send(&v1.TunnelMessage{
|
||||
if err := c.stream.Send(&v1sync.TunnelMessage{
|
||||
ConnId: c.connId,
|
||||
Close: true,
|
||||
}); err != nil {
|
||||
|
||||
@@ -2,30 +2,82 @@ package tunnel
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
|
||||
"connectrpc.com/connect"
|
||||
v1 "github.com/garethgeorge/backrest/gen/go/v1"
|
||||
"github.com/garethgeorge/backrest/gen/go/v1sync"
|
||||
"github.com/hashicorp/go-multierror"
|
||||
"google.golang.org/protobuf/proto"
|
||||
)
|
||||
|
||||
var ErrStreamClosed = errors.New("stream closed")
|
||||
|
||||
type stream interface {
|
||||
Send(item *v1.TunnelMessage) error
|
||||
Receive() (*v1.TunnelMessage, error)
|
||||
Send(item *v1sync.TunnelMessage) error
|
||||
Receive() (*v1sync.TunnelMessage, error)
|
||||
Close() error
|
||||
}
|
||||
|
||||
type cryptedStream struct {
|
||||
stream
|
||||
crypt
|
||||
}
|
||||
|
||||
func newCryptedStream(s stream, secret []byte) *cryptedStream {
|
||||
return &cryptedStream{
|
||||
stream: s,
|
||||
crypt: crypt{
|
||||
secret: secret,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (cs *cryptedStream) Send(item *v1sync.TunnelMessage) error {
|
||||
if item.Data == nil {
|
||||
return cs.stream.Send(item)
|
||||
}
|
||||
bytes, err := proto.Marshal(item)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
enc, err := cs.Encrypt(bytes)
|
||||
if err != nil {
|
||||
return fmt.Errorf("encrypt: %w", err)
|
||||
}
|
||||
return cs.stream.Send(&v1sync.TunnelMessage{
|
||||
Encrypted: enc,
|
||||
})
|
||||
}
|
||||
|
||||
func (cs *cryptedStream) Receive() (*v1sync.TunnelMessage, error) {
|
||||
msg, err := cs.stream.Receive()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
dec, err := cs.Decrypt(msg.Encrypted)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("decrypt: %w", err)
|
||||
}
|
||||
var tm v1sync.TunnelMessage
|
||||
if err := proto.Unmarshal(dec, &tm); err != nil {
|
||||
return nil, fmt.Errorf("unmarshal: %w", err)
|
||||
}
|
||||
if len(tm.Encrypted) != 0 {
|
||||
return nil, fmt.Errorf("unexpected encrypted field in decrypted message")
|
||||
}
|
||||
return &tm, nil
|
||||
}
|
||||
|
||||
type clientStream struct {
|
||||
sendMu sync.Mutex
|
||||
receiveMu sync.Mutex
|
||||
stream *connect.BidiStreamForClient[v1.TunnelMessage, v1.TunnelMessage]
|
||||
stream *connect.BidiStreamForClient[v1sync.TunnelMessage, v1sync.TunnelMessage]
|
||||
closed atomic.Bool
|
||||
}
|
||||
|
||||
func (s *clientStream) Send(item *v1.TunnelMessage) error {
|
||||
func (s *clientStream) Send(item *v1sync.TunnelMessage) error {
|
||||
s.sendMu.Lock()
|
||||
defer s.sendMu.Unlock()
|
||||
if s.closed.Load() {
|
||||
@@ -34,7 +86,7 @@ func (s *clientStream) Send(item *v1.TunnelMessage) error {
|
||||
return s.stream.Send(item)
|
||||
}
|
||||
|
||||
func (s *clientStream) Receive() (*v1.TunnelMessage, error) {
|
||||
func (s *clientStream) Receive() (*v1sync.TunnelMessage, error) {
|
||||
s.receiveMu.Lock()
|
||||
defer s.receiveMu.Unlock()
|
||||
if s.closed.Load() {
|
||||
@@ -64,11 +116,11 @@ func (s *clientStream) Close() error {
|
||||
type serverStream struct {
|
||||
sendMu sync.Mutex
|
||||
receiveMu sync.Mutex
|
||||
stream *connect.BidiStream[v1.TunnelMessage, v1.TunnelMessage]
|
||||
stream *connect.BidiStream[v1sync.TunnelMessage, v1sync.TunnelMessage]
|
||||
closed atomic.Bool
|
||||
}
|
||||
|
||||
func (s *serverStream) Send(item *v1.TunnelMessage) error {
|
||||
func (s *serverStream) Send(item *v1sync.TunnelMessage) error {
|
||||
s.sendMu.Lock()
|
||||
defer s.sendMu.Unlock()
|
||||
if s.closed.Load() {
|
||||
@@ -77,7 +129,7 @@ func (s *serverStream) Send(item *v1.TunnelMessage) error {
|
||||
return s.stream.Send(item)
|
||||
}
|
||||
|
||||
func (s *serverStream) Receive() (*v1.TunnelMessage, error) {
|
||||
func (s *serverStream) Receive() (*v1sync.TunnelMessage, error) {
|
||||
s.receiveMu.Lock()
|
||||
defer s.receiveMu.Unlock()
|
||||
if s.closed.Load() {
|
||||
|
||||
@@ -10,8 +10,8 @@ import (
|
||||
"time"
|
||||
|
||||
"connectrpc.com/connect"
|
||||
v1 "github.com/garethgeorge/backrest/gen/go/v1"
|
||||
"github.com/garethgeorge/backrest/gen/go/v1/v1connect"
|
||||
"github.com/garethgeorge/backrest/gen/go/v1sync"
|
||||
"github.com/garethgeorge/backrest/gen/go/v1sync/v1syncconnect"
|
||||
"github.com/garethgeorge/backrest/internal/testutil"
|
||||
"go.uber.org/zap"
|
||||
"golang.org/x/net/http2"
|
||||
@@ -33,9 +33,9 @@ type sampleHandler struct {
|
||||
provider *ConnectionProvider
|
||||
}
|
||||
|
||||
var _ v1connect.TunnelServiceHandler = (*sampleHandler)(nil)
|
||||
var _ v1syncconnect.TunnelServiceHandler = (*sampleHandler)(nil)
|
||||
|
||||
func (sh *sampleHandler) Tunnel(ctx context.Context, stream *connect.BidiStream[v1.TunnelMessage, v1.TunnelMessage]) error {
|
||||
func (sh *sampleHandler) Tunnel(ctx context.Context, stream *connect.BidiStream[v1sync.TunnelMessage, v1sync.TunnelMessage]) error {
|
||||
wrapped := NewWrappedStream(stream, WithLogger(sh.logger))
|
||||
wrapped.ProvideConnectionsTo(sh.provider)
|
||||
sh.streams = append(sh.streams, wrapped)
|
||||
@@ -113,7 +113,7 @@ func TestConnect(t *testing.T) {
|
||||
|
||||
// Serve the sample handler
|
||||
mux := http.NewServeMux()
|
||||
mux.Handle(v1connect.NewTunnelServiceHandler(sampleHandler))
|
||||
mux.Handle(v1syncconnect.NewTunnelServiceHandler(sampleHandler))
|
||||
grpcServer := &http.Server{
|
||||
Addr: testutil.AllocOpenBindAddr(t),
|
||||
Handler: h2c.NewHandler(mux, &http2.Server{}), // h2c is HTTP/2 without TLS for grpc-connect support.
|
||||
@@ -121,7 +121,7 @@ func TestConnect(t *testing.T) {
|
||||
listenAndServeForTest("gRPCServer", t, grpcServer)
|
||||
|
||||
// Create a client and connect to the server
|
||||
client := v1connect.NewTunnelServiceClient(NewInsecureHttpClient(), "http://"+grpcServer.Addr)
|
||||
client := v1syncconnect.NewTunnelServiceClient(NewInsecureHttpClient(), "http://"+grpcServer.Addr)
|
||||
stream := client.Tunnel(ctx)
|
||||
wrapped := NewWrappedStreamFromClient(stream, WithLogger(testutil.NewTestLogger(t).Named("client")))
|
||||
go func() {
|
||||
|
||||
@@ -11,7 +11,7 @@ import (
|
||||
"time"
|
||||
|
||||
"connectrpc.com/connect"
|
||||
v1 "github.com/garethgeorge/backrest/gen/go/v1"
|
||||
"github.com/garethgeorge/backrest/gen/go/v1sync"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
@@ -46,6 +46,9 @@ type WrappedStream struct {
|
||||
|
||||
handlingPackets atomic.Bool
|
||||
streamStopped atomic.Bool
|
||||
|
||||
sharedSecretCh chan struct{}
|
||||
sharedSecret []byte
|
||||
}
|
||||
|
||||
func newWrappedStreamInternal(stream stream, isClient bool, opts ...WrappedStreamOptions) *WrappedStream {
|
||||
@@ -54,6 +57,7 @@ func newWrappedStreamInternal(stream stream, isClient bool, opts ...WrappedStrea
|
||||
stream: stream,
|
||||
heartbeatInterval: 30 * time.Second,
|
||||
conns: make(map[int64]*connState),
|
||||
sharedSecretCh: make(chan struct{}),
|
||||
}
|
||||
for _, opt := range opts {
|
||||
opt(ws)
|
||||
@@ -66,13 +70,13 @@ func newWrappedStreamInternal(stream stream, isClient bool, opts ...WrappedStrea
|
||||
return ws
|
||||
}
|
||||
|
||||
func NewWrappedStream(stream *connect.BidiStream[v1.TunnelMessage, v1.TunnelMessage], opts ...WrappedStreamOptions) *WrappedStream {
|
||||
func NewWrappedStream(stream *connect.BidiStream[v1sync.TunnelMessage, v1sync.TunnelMessage], opts ...WrappedStreamOptions) *WrappedStream {
|
||||
return newWrappedStreamInternal(&serverStream{
|
||||
stream: stream,
|
||||
}, false, opts...)
|
||||
}
|
||||
|
||||
func NewWrappedStreamFromClient(stream *connect.BidiStreamForClient[v1.TunnelMessage, v1.TunnelMessage], opts ...WrappedStreamOptions) *WrappedStream {
|
||||
func NewWrappedStreamFromClient(stream *connect.BidiStreamForClient[v1sync.TunnelMessage, v1sync.TunnelMessage], opts ...WrappedStreamOptions) *WrappedStream {
|
||||
return newWrappedStreamInternal(&clientStream{
|
||||
stream: stream,
|
||||
}, true, opts...)
|
||||
@@ -82,6 +86,14 @@ func (ws *WrappedStream) allocConnID() int64 {
|
||||
return ws.lastConnID.Add(2)
|
||||
}
|
||||
|
||||
func (ws *WrappedStream) getSharedSecret() ([]byte, error) {
|
||||
<-ws.sharedSecretCh
|
||||
if ws.sharedSecret == nil {
|
||||
return nil, fmt.Errorf("shared secret not available")
|
||||
}
|
||||
return ws.sharedSecret, nil
|
||||
}
|
||||
|
||||
func (ws *WrappedStream) IsReady() bool {
|
||||
return ws.handlingPackets.Load() && !ws.streamStopped.Load()
|
||||
}
|
||||
@@ -91,8 +103,13 @@ func (ws *WrappedStream) Dial() (net.Conn, error) {
|
||||
return nil, fmt.Errorf("cannot dial before handling packets")
|
||||
}
|
||||
|
||||
secret, err := ws.getSharedSecret()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get shared secret: %w", err)
|
||||
}
|
||||
|
||||
connID := ws.allocConnID()
|
||||
new := newConnState(ws.stream, connID, ws.logger)
|
||||
new := newConnState(ws.stream, connID, secret, ws.logger)
|
||||
if err := new.sendOpenPacket(); err != nil {
|
||||
return nil, fmt.Errorf("send open packet: %w", err)
|
||||
}
|
||||
@@ -106,7 +123,7 @@ func (ws *WrappedStream) ProvideConnectionsTo(provider *ConnectionProvider) {
|
||||
ws.provider = provider
|
||||
}
|
||||
|
||||
func (ws *WrappedStream) sendHeartbeats(ctx context.Context) {
|
||||
func (ws *WrappedStream) sendHeartbeats(ctx context.Context, stream stream) {
|
||||
if ws.heartbeatInterval <= 0 || !ws.isClient {
|
||||
return
|
||||
}
|
||||
@@ -122,7 +139,7 @@ func (ws *WrappedStream) sendHeartbeats(ctx context.Context) {
|
||||
if ws.logger != nil {
|
||||
ws.logger.Debug("sending heartbeat")
|
||||
}
|
||||
if err := ws.stream.Send(&v1.TunnelMessage{
|
||||
if err := stream.Send(&v1sync.TunnelMessage{
|
||||
ConnId: -1, // handshake packet
|
||||
}); err != nil && ws.logger != nil {
|
||||
ws.logger.Error("failed to send heartbeat", zap.Error(err))
|
||||
@@ -148,13 +165,29 @@ func (ws *WrappedStream) HandlePackets(ctx context.Context) error {
|
||||
return fmt.Errorf("generate key for handshake packet: %w", err)
|
||||
}
|
||||
|
||||
if err := ws.stream.Send(&v1.TunnelMessage{
|
||||
ConnId: -100, // hadnshake packet
|
||||
if err := ws.stream.Send(&v1sync.TunnelMessage{
|
||||
ConnId: -100, // handshake packet
|
||||
PubkeyEcdhX25519: key.PublicKey().Bytes(),
|
||||
}); err != nil {
|
||||
return fmt.Errorf("send handshake packet: %w", err)
|
||||
}
|
||||
|
||||
go func() {
|
||||
timeoutTimer := time.NewTimer(5 * time.Second)
|
||||
defer timeoutTimer.Stop()
|
||||
select {
|
||||
case <-timeoutTimer.C:
|
||||
if ws.logger != nil {
|
||||
ws.logger.Warn("timeout waiting for handshake response")
|
||||
}
|
||||
ws.Shutdown()
|
||||
case <-ws.sharedSecretCh:
|
||||
if ws.logger != nil {
|
||||
ws.logger.Info("handshake response received, shared secret established")
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
// receive handshake packet
|
||||
handshake, err := ws.stream.Receive()
|
||||
if err != nil {
|
||||
@@ -170,18 +203,21 @@ func (ws *WrappedStream) HandlePackets(ctx context.Context) error {
|
||||
if err != nil {
|
||||
return fmt.Errorf("parse peer public key: %w", err)
|
||||
}
|
||||
|
||||
_, err = key.ECDH(peerKey)
|
||||
sharedSecret, err := key.ECDH(peerKey)
|
||||
if err != nil {
|
||||
return fmt.Errorf("compute shared key: %w", err)
|
||||
}
|
||||
// TODO: use the key for encryption and decryption of messages.
|
||||
ws.sharedSecret = sharedSecret
|
||||
close(ws.sharedSecretCh)
|
||||
|
||||
// cryptedStream := newCryptedStream(ws.stream, ws.sharedSecret)
|
||||
cryptedStream := ws.stream
|
||||
|
||||
newConn := func(connId int64) *connState {
|
||||
if ws.logger != nil {
|
||||
ws.logger.Info("new tunnel connection", zap.Int64("connId", connId))
|
||||
}
|
||||
new := newConnState(ws.stream, connId, ws.logger)
|
||||
new := newConnState(ws.stream, connId, ws.sharedSecret, ws.logger)
|
||||
ws.conns[connId] = new
|
||||
ws.provider.ProvideConn(new)
|
||||
return new
|
||||
@@ -191,10 +227,10 @@ func (ws *WrappedStream) HandlePackets(ctx context.Context) error {
|
||||
defer headOfLineBlockingTimer.Stop()
|
||||
|
||||
// send heartbeats in a separate goroutine if heartbeat interval is set
|
||||
go ws.sendHeartbeats(ctx)
|
||||
go ws.sendHeartbeats(ctx, cryptedStream)
|
||||
|
||||
for {
|
||||
msg, err := ws.stream.Receive()
|
||||
msg, err := cryptedStream.Receive()
|
||||
if err != nil {
|
||||
if ws.handlingPackets.Load() {
|
||||
return nil
|
||||
|
||||
Reference in New Issue
Block a user