diff --git a/internal/client/client.go b/internal/client/client.go index ea9a983..c60faa1 100644 --- a/internal/client/client.go +++ b/internal/client/client.go @@ -28,6 +28,7 @@ type TeleportClient struct { cancel context.CancelFunc connectionPool chan net.Conn maxPoolSize int + derivedKey []byte } // NewTeleportClient creates a new teleport client @@ -49,6 +50,7 @@ func NewTeleportClient(config *config.Config) *TeleportClient { cancel: cancel, connectionPool: make(chan net.Conn, maxPoolSize), maxPoolSize: maxPoolSize, + derivedKey: encryption.DeriveKey(config.EncryptionKey), } } @@ -650,23 +652,34 @@ func (tc *TeleportClient) handleTCPConnection(clientConn net.Conn, rule config.P var wg sync.WaitGroup wg.Add(2) - // Forward data from client to server + // Encrypt local plaintext onto the tunnel go func() { defer wg.Done() - tc.forwardData(clientConn, serverConn) + tc.forwardData(clientConn, serverConn, true) }() - // Forward data from server to client + // Decrypt tunnel frames to the local connection go func() { defer wg.Done() - tc.forwardData(serverConn, clientConn) + tc.forwardData(serverConn, clientConn, false) }() wg.Wait() } -// forwardData forwards data from src to dst -func (tc *TeleportClient) forwardData(src, dst net.Conn) { +// forwardData copies between a local/plaintext conn and the tunnel. +// When toTunnel is true, plaintext is AES-GCM encrypted and written as +// uint32 length + ciphertext. When false, framed ciphertext is decrypted +// and written as plaintext. +func (tc *TeleportClient) forwardData(src, dst net.Conn, toTunnel bool) { + if toTunnel { + tc.forwardPlainToTunnel(src, dst) + return + } + tc.forwardTunnelToPlain(src, dst) +} + +func (tc *TeleportClient) forwardPlainToTunnel(src, dst net.Conn) { buffer := make([]byte, 4096) for { select { @@ -675,14 +688,32 @@ func (tc *TeleportClient) forwardData(src, dst net.Conn) { default: n, err := src.Read(buffer) if err != nil { - // Close the destination connection when source closes dst.Close() return } + if n == 0 { + continue + } + if err := encryption.WriteEncryptedFrame(dst, buffer[:n], tc.derivedKey); err != nil { + src.Close() + return + } + } + } +} - _, err = dst.Write(buffer[:n]) +func (tc *TeleportClient) forwardTunnelToPlain(src, dst net.Conn) { + for { + select { + case <-tc.ctx.Done(): + return + default: + plain, err := encryption.ReadEncryptedFrame(src, tc.derivedKey) if err != nil { - // Close the source connection when destination closes + dst.Close() + return + } + if _, err := dst.Write(plain); err != nil { src.Close() return } @@ -703,9 +734,8 @@ func (tc *TeleportClient) sendRequestToConnection(conn net.Conn, request types.P return err } - // Encrypt the data - key := encryption.DeriveKey(tc.config.EncryptionKey) - encryptedData, err := encryption.EncryptData(data, key) + // Encrypt the data (key derived once at process start) + encryptedData, err := encryption.EncryptData(data, tc.derivedKey) if err != nil { logger.WithFields(map[string]interface{}{ "error": err, diff --git a/internal/client/forward_test.go b/internal/client/forward_test.go new file mode 100644 index 0000000..b52efdf --- /dev/null +++ b/internal/client/forward_test.go @@ -0,0 +1,120 @@ +package client + +import ( + "bytes" + "encoding/binary" + "io" + "net" + "testing" + "time" + + "teleport/pkg/config" + "teleport/pkg/encryption" +) + +func testClient(t *testing.T) *TeleportClient { + t.Helper() + return NewTeleportClient(&config.Config{ + EncryptionKey: "a0e3dd20a761b118ca234160dd8b87230a001e332a97c9cfe3b8b9c99efaae03", + }) +} + +func TestForwardDataHidesHTTPPlaintextOnTunnel(t *testing.T) { + tc := testClient(t) + defer tc.Stop() + + localA, localB := net.Pipe() + tunnelA, tunnelB := net.Pipe() + defer localA.Close() + defer localB.Close() + defer tunnelA.Close() + defer tunnelB.Close() + + go tc.forwardData(localA, tunnelA, true) + + req := []byte("GET /secret HTTP/1.1\r\nHost: example.com\r\n\r\n") + writeErr := make(chan error, 1) + go func() { + _, err := localB.Write(req) + writeErr <- err + }() + + tunnelB.SetReadDeadline(time.Now().Add(5 * time.Second)) + var length uint32 + if err := binary.Read(tunnelB, binary.BigEndian, &length); err != nil { + t.Fatalf("failed to read frame length from tunnel: %v", err) + } + if length == 0 || length > encryption.MaxFrameSize { + t.Fatalf("unexpected frame length %d", length) + } + ciphertext := make([]byte, length) + if _, err := io.ReadFull(tunnelB, ciphertext); err != nil { + t.Fatalf("failed to read ciphertext from tunnel: %v", err) + } + + if bytes.Contains(ciphertext, []byte("GET /")) || + bytes.Contains(ciphertext, []byte("Host:")) || + bytes.Contains(ciphertext, []byte("example.com")) { + t.Fatal("HTTP plaintext visible on client->server tunnel") + } + + var frame bytes.Buffer + if err := binary.Write(&frame, binary.BigEndian, length); err != nil { + t.Fatal(err) + } + frame.Write(ciphertext) + plain, err := encryption.ReadEncryptedFrame(&frame, tc.derivedKey) + if err != nil { + t.Fatalf("decrypt of tunneled frame failed: %v", err) + } + if !bytes.Equal(plain, req) { + t.Errorf("decrypted payload mismatch: got %q want %q", plain, req) + } + + select { + case err := <-writeErr: + if err != nil { + t.Fatalf("write to local conn failed: %v", err) + } + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for local write") + } +} + +func TestForwardDataDecryptsTunnelToLocal(t *testing.T) { + tc := testClient(t) + defer tc.Stop() + + tunnelA, tunnelB := net.Pipe() + localA, localB := net.Pipe() + defer tunnelA.Close() + defer tunnelB.Close() + defer localA.Close() + defer localB.Close() + + go tc.forwardData(tunnelA, localA, false) + + payload := []byte("HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\nhello") + writeErr := make(chan error, 1) + go func() { + writeErr <- encryption.WriteEncryptedFrame(tunnelB, payload, tc.derivedKey) + }() + + localB.SetReadDeadline(time.Now().Add(5 * time.Second)) + got := make([]byte, len(payload)) + if _, err := io.ReadFull(localB, got); err != nil { + t.Fatalf("failed to read decrypted payload: %v", err) + } + if !bytes.Equal(got, payload) { + t.Errorf("got %q want %q", got, payload) + } + + select { + case err := <-writeErr: + if err != nil { + t.Fatalf("WriteEncryptedFrame failed: %v", err) + } + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for tunnel write") + } +} diff --git a/internal/server/forward_test.go b/internal/server/forward_test.go new file mode 100644 index 0000000..93fc255 --- /dev/null +++ b/internal/server/forward_test.go @@ -0,0 +1,118 @@ +package server + +import ( + "bytes" + "encoding/binary" + "io" + "net" + "testing" + "time" + + "teleport/pkg/config" + "teleport/pkg/encryption" +) + +func testServer(t *testing.T) *TeleportServer { + t.Helper() + return NewTeleportServer(&config.Config{ + EncryptionKey: "a0e3dd20a761b118ca234160dd8b87230a001e332a97c9cfe3b8b9c99efaae03", + }) +} + +func TestForwardDataHidesHTTPPlaintextOnTunnel(t *testing.T) { + ts := testServer(t) + defer ts.Stop() + + targetA, targetB := net.Pipe() + tunnelA, tunnelB := net.Pipe() + defer targetA.Close() + defer targetB.Close() + defer tunnelA.Close() + defer tunnelB.Close() + + // target -> tunnel (encrypt) + go ts.forwardData(targetA, tunnelA, true) + + resp := []byte("HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\n\r\nhello") + writeErr := make(chan error, 1) + go func() { + _, err := targetB.Write(resp) + writeErr <- err + }() + + tunnelB.SetReadDeadline(time.Now().Add(5 * time.Second)) + var length uint32 + if err := binary.Read(tunnelB, binary.BigEndian, &length); err != nil { + t.Fatalf("failed to read frame length from tunnel: %v", err) + } + ciphertext := make([]byte, length) + if _, err := io.ReadFull(tunnelB, ciphertext); err != nil { + t.Fatalf("failed to read ciphertext from tunnel: %v", err) + } + + if bytes.Contains(ciphertext, []byte("HTTP/1.1")) || + bytes.Contains(ciphertext, []byte("Content-Type")) || + bytes.Contains(ciphertext, []byte("hello")) { + t.Fatal("HTTP plaintext visible on server->client tunnel") + } + + var frame bytes.Buffer + if err := binary.Write(&frame, binary.BigEndian, length); err != nil { + t.Fatal(err) + } + frame.Write(ciphertext) + plain, err := encryption.ReadEncryptedFrame(&frame, ts.derivedKey) + if err != nil { + t.Fatalf("decrypt of tunneled frame failed: %v", err) + } + if !bytes.Equal(plain, resp) { + t.Errorf("decrypted payload mismatch: got %q want %q", plain, resp) + } + + select { + case err := <-writeErr: + if err != nil { + t.Fatalf("write to target conn failed: %v", err) + } + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for target write") + } +} + +func TestForwardDataDecryptsTunnelToTarget(t *testing.T) { + ts := testServer(t) + defer ts.Stop() + + tunnelA, tunnelB := net.Pipe() + targetA, targetB := net.Pipe() + defer tunnelA.Close() + defer tunnelB.Close() + defer targetA.Close() + defer targetB.Close() + + go ts.forwardData(tunnelA, targetA, false) + + req := []byte("GET / HTTP/1.1\r\nHost: example.com\r\n\r\n") + writeErr := make(chan error, 1) + go func() { + writeErr <- encryption.WriteEncryptedFrame(tunnelB, req, ts.derivedKey) + }() + + targetB.SetReadDeadline(time.Now().Add(5 * time.Second)) + got := make([]byte, len(req)) + if _, err := io.ReadFull(targetB, got); err != nil { + t.Fatalf("failed to read decrypted payload: %v", err) + } + if !bytes.Equal(got, req) { + t.Errorf("got %q want %q", got, req) + } + + select { + case err := <-writeErr: + if err != nil { + t.Fatalf("WriteEncryptedFrame failed: %v", err) + } + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for tunnel write") + } +} diff --git a/internal/server/server.go b/internal/server/server.go index 6a385db..a5ad42f 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -35,6 +35,7 @@ type TeleportServer struct { goroutineSem chan struct{} // Semaphore for limiting concurrent goroutines maxGoroutines int // Maximum concurrent goroutines metrics *metrics.Metrics + derivedKey []byte } // NewTeleportServer creates a new teleport server @@ -76,6 +77,7 @@ func NewTeleportServer(config *config.Config) *TeleportServer { goroutineSem: make(chan struct{}, maxGoroutines), maxGoroutines: maxGoroutines, metrics: metricsInstance, + derivedKey: encryption.DeriveKey(config.EncryptionKey), } } @@ -403,12 +405,12 @@ func (ts *TeleportServer) handleTCPForward(clientConn net.Conn, rule *config.Por go func() { defer wg.Done() - ts.forwardData(clientConn, targetConn) + ts.forwardData(clientConn, targetConn, false) // decrypt from tunnel }() go func() { defer wg.Done() - ts.forwardData(targetConn, clientConn) + ts.forwardData(targetConn, clientConn, true) // encrypt onto tunnel }() wg.Wait() @@ -704,8 +706,19 @@ func (ts *TeleportServer) deserializeTaggedUDPPacket(data []byte) (types.TaggedU }, nil } -// forwardData forwards data between two connections -func (ts *TeleportServer) forwardData(src, dst net.Conn) { +// forwardData copies between a local/target conn and the tunnel. +// When toTunnel is true, plaintext is AES-GCM encrypted and written as +// uint32 length + ciphertext. When false, framed ciphertext is decrypted +// and written as plaintext. +func (ts *TeleportServer) forwardData(src, dst net.Conn, toTunnel bool) { + if toTunnel { + ts.forwardPlainToTunnel(src, dst) + return + } + ts.forwardTunnelToPlain(src, dst) +} + +func (ts *TeleportServer) forwardPlainToTunnel(src, dst net.Conn) { buffer := make([]byte, 4096) for { select { @@ -714,14 +727,32 @@ func (ts *TeleportServer) forwardData(src, dst net.Conn) { default: n, err := src.Read(buffer) if err != nil { - // Close the destination connection when source closes dst.Close() return } + if n == 0 { + continue + } + if err := encryption.WriteEncryptedFrame(dst, buffer[:n], ts.derivedKey); err != nil { + src.Close() + return + } + } + } +} - _, err = dst.Write(buffer[:n]) +func (ts *TeleportServer) forwardTunnelToPlain(src, dst net.Conn) { + for { + select { + case <-ts.ctx.Done(): + return + default: + plain, err := encryption.ReadEncryptedFrame(src, ts.derivedKey) if err != nil { - // Close the source connection when destination closes + dst.Close() + return + } + if _, err := dst.Write(plain); err != nil { src.Close() return } @@ -757,9 +788,8 @@ func (ts *TeleportServer) readRequest(conn net.Conn, request *types.PortForwardR bytesRead += n } - // Decrypt the data - key := encryption.DeriveKey(ts.config.EncryptionKey) - decryptedData, err := encryption.DecryptData(encryptedData, key) + // Decrypt the data (key derived once at process start) + decryptedData, err := encryption.DecryptData(encryptedData, ts.derivedKey) if err != nil { logger.WithFields(map[string]interface{}{ "error": err, diff --git a/pkg/encryption/frame.go b/pkg/encryption/frame.go new file mode 100644 index 0000000..5603a7a --- /dev/null +++ b/pkg/encryption/frame.go @@ -0,0 +1,68 @@ +package encryption + +import ( + "encoding/binary" + "fmt" + "io" +) + +// MaxFrameSize is the maximum AES-GCM ciphertext length accepted on a framed +// TCP tunnel (uint32 length prefix + ciphertext). +const MaxFrameSize = 1024 * 1024 // 1 MiB + +// WriteEncryptedFrame encrypts plaintext with AES-GCM (EncryptData) and writes +// a uint32 big-endian length prefix followed by the ciphertext. +func WriteEncryptedFrame(w io.Writer, plaintext, key []byte) error { + ciphertext, err := EncryptData(plaintext, key) + if err != nil { + return err + } + if len(ciphertext) == 0 { + return fmt.Errorf("encrypted data cannot be empty") + } + if len(ciphertext) > MaxFrameSize { + return fmt.Errorf("frame too large: %d bytes (max %d)", len(ciphertext), MaxFrameSize) + } + + var length [4]byte + binary.BigEndian.PutUint32(length[:], uint32(len(ciphertext))) + if err := writeFull(w, length[:]); err != nil { + return err + } + return writeFull(w, ciphertext) +} + +// ReadEncryptedFrame reads a uint32 big-endian length, caps it, reads the +// ciphertext, and decrypts it with AES-GCM (DecryptData). +func ReadEncryptedFrame(r io.Reader, key []byte) ([]byte, error) { + var length uint32 + if err := binary.Read(r, binary.BigEndian, &length); err != nil { + return nil, err + } + if length == 0 { + return nil, fmt.Errorf("frame length cannot be zero") + } + if length > MaxFrameSize { + return nil, fmt.Errorf("frame too large: %d bytes (max %d)", length, MaxFrameSize) + } + + ciphertext := make([]byte, length) + if _, err := io.ReadFull(r, ciphertext); err != nil { + return nil, fmt.Errorf("failed to read frame data: %v", err) + } + return DecryptData(ciphertext, key) +} + +func writeFull(w io.Writer, data []byte) error { + for len(data) > 0 { + n, err := w.Write(data) + if err != nil { + return err + } + if n == 0 { + return io.ErrShortWrite + } + data = data[n:] + } + return nil +} diff --git a/pkg/encryption/frame_test.go b/pkg/encryption/frame_test.go new file mode 100644 index 0000000..eb4564a --- /dev/null +++ b/pkg/encryption/frame_test.go @@ -0,0 +1,136 @@ +package encryption + +import ( + "bytes" + "encoding/binary" + "io" + "strings" + "testing" +) + +func testKey(t *testing.T) []byte { + t.Helper() + return DeriveKey("test-encryption-key-for-tcp-frames") +} + +func TestEncryptedFrameRoundTrip(t *testing.T) { + key := testKey(t) + original := []byte("Hello, framed TCP tunnel!") + + var buf bytes.Buffer + if err := WriteEncryptedFrame(&buf, original, key); err != nil { + t.Fatalf("WriteEncryptedFrame failed: %v", err) + } + + got, err := ReadEncryptedFrame(&buf, key) + if err != nil { + t.Fatalf("ReadEncryptedFrame failed: %v", err) + } + if !bytes.Equal(got, original) { + t.Errorf("round trip mismatch: got %q want %q", got, original) + } +} + +func TestEncryptedFrameMultipleChunks(t *testing.T) { + key := testKey(t) + chunks := [][]byte{ + []byte("chunk-one"), + []byte("chunk-two-is-longer"), + []byte{0x00, 0x01, 0xff}, + } + + var buf bytes.Buffer + for _, c := range chunks { + if err := WriteEncryptedFrame(&buf, c, key); err != nil { + t.Fatalf("WriteEncryptedFrame failed: %v", err) + } + } + + for i, want := range chunks { + got, err := ReadEncryptedFrame(&buf, key) + if err != nil { + t.Fatalf("ReadEncryptedFrame chunk %d failed: %v", i, err) + } + if !bytes.Equal(got, want) { + t.Errorf("chunk %d mismatch: got %q want %q", i, got, want) + } + } +} + +func TestEncryptedFrameRejectsZeroLength(t *testing.T) { + key := testKey(t) + var buf bytes.Buffer + if err := binary.Write(&buf, binary.BigEndian, uint32(0)); err != nil { + t.Fatal(err) + } + if _, err := ReadEncryptedFrame(&buf, key); err == nil { + t.Fatal("expected error for zero-length frame") + } +} + +func TestEncryptedFrameRejectsOversizeLength(t *testing.T) { + key := testKey(t) + var buf bytes.Buffer + if err := binary.Write(&buf, binary.BigEndian, uint32(MaxFrameSize+1)); err != nil { + t.Fatal(err) + } + if _, err := ReadEncryptedFrame(&buf, key); err == nil { + t.Fatal("expected error for oversize frame length") + } else if !strings.Contains(err.Error(), "too large") { + t.Errorf("expected too-large error, got %v", err) + } +} + +func TestEncryptedFrameWrongKey(t *testing.T) { + key1 := DeriveKey("tcp-frame-key-1") + key2 := DeriveKey("tcp-frame-key-2") + + var buf bytes.Buffer + if err := WriteEncryptedFrame(&buf, []byte("secret"), key1); err != nil { + t.Fatal(err) + } + if _, err := ReadEncryptedFrame(&buf, key2); err == nil { + t.Fatal("decrypting with the wrong key should fail") + } +} + +func TestHTTPPlaintextNotVisibleOnFramedTunnel(t *testing.T) { + key := testKey(t) + req := []byte("GET / HTTP/1.1\r\nHost: example.com\r\n\r\n") + + var wire bytes.Buffer + if err := WriteEncryptedFrame(&wire, req, key); err != nil { + t.Fatalf("WriteEncryptedFrame failed: %v", err) + } + sniffed := wire.Bytes() + + if bytes.Contains(sniffed, []byte("GET /")) { + t.Fatal("sniffer saw HTTP request line on the tunnel") + } + if bytes.Contains(sniffed, []byte("Host:")) { + t.Fatal("sniffer saw HTTP Host header on the tunnel") + } + if bytes.Contains(sniffed, []byte("example.com")) { + t.Fatal("sniffer saw HTTP hostname on the tunnel") + } + + // Length prefix is 4 bytes; ciphertext must be longer than plaintext. + if len(sniffed) < 4+len(req) { + t.Fatalf("framed ciphertext too short: %d", len(sniffed)) + } + + got, err := ReadEncryptedFrame(bytes.NewReader(sniffed), key) + if err != nil { + t.Fatalf("decrypt failed: %v", err) + } + if !bytes.Equal(got, req) { + t.Errorf("decrypted HTTP mismatch: got %q", got) + } +} + +func TestReadEncryptedFrameEOF(t *testing.T) { + key := testKey(t) + if _, err := ReadEncryptedFrame(bytes.NewReader(nil), key); err != io.EOF { + t.Errorf("expected io.EOF, got %v", err) + } +}