diff --git a/internal/api/syncapi/tunnel/aes.go b/internal/api/syncapi/tunnel/aes.go new file mode 100644 index 00000000..7c944483 --- /dev/null +++ b/internal/api/syncapi/tunnel/aes.go @@ -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) +} diff --git a/internal/api/syncapi/tunnel/aes_test.go b/internal/api/syncapi/tunnel/aes_test.go new file mode 100644 index 00000000..96a97e9d --- /dev/null +++ b/internal/api/syncapi/tunnel/aes_test.go @@ -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) + } + } +} diff --git a/internal/api/syncapi/tunnel/connstate.go b/internal/api/syncapi/tunnel/connstate.go index 4f9c1df9..c52f4e9b 100644 --- a/internal/api/syncapi/tunnel/connstate.go +++ b/internal/api/syncapi/tunnel/connstate.go @@ -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 { diff --git a/internal/api/syncapi/tunnel/streamutil.go b/internal/api/syncapi/tunnel/streamutil.go index dda632b8..e7d4394f 100644 --- a/internal/api/syncapi/tunnel/streamutil.go +++ b/internal/api/syncapi/tunnel/streamutil.go @@ -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() { diff --git a/internal/api/syncapi/tunnel/tunnel_test.go b/internal/api/syncapi/tunnel/tunnel_test.go index 793b097c..b690ed35 100644 --- a/internal/api/syncapi/tunnel/tunnel_test.go +++ b/internal/api/syncapi/tunnel/tunnel_test.go @@ -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() { diff --git a/internal/api/syncapi/tunnel/wrappedstream.go b/internal/api/syncapi/tunnel/wrappedstream.go index 5901b998..b8886275 100644 --- a/internal/api/syncapi/tunnel/wrappedstream.go +++ b/internal/api/syncapi/tunnel/wrappedstream.go @@ -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