TCP forwardData copied plaintext after the handshake. Frame each chunk as uint32 length plus AES-GCM ciphertext in both client and server, both directions.
This commit is contained in:
+42
-12
@@ -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,
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user