Encrypt TCP tunnel payload with AES-GCM (#6)
CI / check-and-test (push) Failing after 10s

This commit was merged in pull request #6.
This commit is contained in:
s1d3sw1ped_bot
2026-08-31 22:30:40 -05:00
parent 1fef0b6a70
commit 283781fd47
6 changed files with 524 additions and 22 deletions
+42 -12
View File
@@ -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,
+120
View File
@@ -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")
}
}
+118
View File
@@ -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")
}
}
+40 -10
View File
@@ -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,