diff --git a/mtglib/internal/tls/fake/bytes_pool.go b/mtglib/internal/tls/fake/bytes_pool.go new file mode 100644 index 0000000..2eaca15 --- /dev/null +++ b/mtglib/internal/tls/fake/bytes_pool.go @@ -0,0 +1,23 @@ +package fake + +import ( + "bytes" + "sync" +) + +var ( + bytesPool = sync.Pool{ + New: func() any { + return &bytes.Buffer{} + }, + } +) + +func acquireBuffer() *bytes.Buffer { + return bytesPool.Get().(*bytes.Buffer) +} + +func releaseBuffer(b *bytes.Buffer) { + b.Reset() + bytesPool.Put(b) +} diff --git a/mtglib/internal/tls/fake/client_side.go b/mtglib/internal/tls/fake/client_side.go index dc41397..8d89b5a 100644 --- a/mtglib/internal/tls/fake/client_side.go +++ b/mtglib/internal/tls/fake/client_side.go @@ -20,9 +20,9 @@ const ( TypeHandshakeClient = 0x01 RandomLen = 32 - // record_type(1) + version(2) + size(2) + handshake_type(1) + uint24_length(3) + client_version(2) - clientRandomOffset = 1 + 2 + 2 + 1 + 3 + 2 + RandomOffset = 1 + 2 + 2 + 1 + 3 + 2 + sniDNSNamesListType = 0 ) @@ -81,7 +81,7 @@ func ReadClientHello(conn net.Conn, secret mtglib.Secret, tolerateTimeSkewness t digest := hmac.New(sha256.New, secret.Key[:]) // we write a copy of the handshake with client random all nullified. - digest.Write(handshakeCopyBuf.Next(clientRandomOffset)) + digest.Write(handshakeCopyBuf.Next(RandomOffset)) handshakeCopyBuf.Next(RandomLen) digest.Write(emptyRandom[:]) digest.Write(handshakeCopyBuf.Bytes()) diff --git a/mtglib/internal/tls/fake/server_side.go b/mtglib/internal/tls/fake/server_side.go new file mode 100644 index 0000000..8f920c5 --- /dev/null +++ b/mtglib/internal/tls/fake/server_side.go @@ -0,0 +1,141 @@ +package fake + +import ( + "bytes" + "crypto/hmac" + "crypto/rand" + "crypto/sha256" + "encoding/binary" + "io" + + "github.com/9seconds/mtg/v2/mtglib" + "github.com/9seconds/mtg/v2/mtglib/internal/tls" + "golang.org/x/crypto/curve25519" +) + +const ( + TypeHandshakeServer = 0x02 + + ChangeCipherValue = 0x01 + + EllipticCurveLen = 32 +) + +var ( + serverHelloSuffix = []byte{ + 0x00, // no compression + 0x00, 0x2e, // 46 bytes of data + 0x00, 0x2b, // Extension - Supported Versions + 0x00, 0x02, // 2 bytes are following + 0x03, 0x04, // TLS 1.3 + 0x00, 0x33, // Extension - Key Share + 0x00, 0x24, // 36 bytes + 0x00, 0x1d, // x25519 curve + 0x00, 0x20, // 32 bytes of key + } +) + +func SendServerHello(w io.Writer, secret mtglib.Secret, clientHello *ClientHello) ([]byte, error) { + buf := &bytes.Buffer{} + buf.Grow(tls.MaxRecordSize) + + generateServerHello(buf, clientHello) + generateChangeCipherValue(buf) + + noise := &bytes.Buffer{} + generateNoise(noise) + + packet := buf.Bytes() + digest := hmac.New(sha256.New, secret.Key[:]) + + digest.Write(clientHello.Random[:]) + digest.Write(packet) + digest.Write(noise.Bytes()) + copy(packet[RandomOffset:], digest.Sum(nil)) + + _, err := w.Write(packet) + + return noise.Bytes(), err +} + +func generateServerHello(buf *bytes.Buffer, hello *ClientHello) { + payload := acquireBuffer() + defer releaseBuffer(payload) + + generateServerHelloPayload(payload, hello) + + // 16 - type is 0x16 (handshake record) + // 03 03 - legacy protocol version of "3,3" (TLS 1.2) + // 00 7a - 0x7A (122) bytes of handshake message follows + + // 16 - type is 0x16 (handshake record) + buf.WriteByte(tls.TypeHandshake) + // 03 03 - legacy protocol version of "3,3" (TLS 1.2) + buf.Write(tls.TLSVersion[:]) + // 00 7a - 0x7A (122) bytes of handshake message follows + binary.Write(buf, binary.BigEndian, uint16(payload.Len())) //nolint: errcheck + + payload.WriteTo(buf) //nolint: errcheck +} + +func generateServerHelloPayload(buf *bytes.Buffer, hello *ClientHello) { + data := [4]byte{} + + payload := acquireBuffer() + defer releaseBuffer(payload) + + generateServerHelloHandshakePayload(payload, hello) + + // 02 - handshake message type 0x02 (server hello) + // 00 00 76 - 0x76 (118) bytes of server hello data follows + buf.WriteByte(TypeHandshakeServer) + // 00 00 76 - 0x76 (118) bytes of server hello data follows + binary.BigEndian.PutUint32(data[:], uint32(payload.Len())) + buf.Write(data[1:]) + + payload.WriteTo(buf) //nolint: errcheck +} + +func generateServerHelloHandshakePayload(buf *bytes.Buffer, hello *ClientHello) { + // The unusual version number ("3,3" representing TLS 1.2) is due to + // TLS 1.0 being a minor revision of the SSL 3.0 protocol. Therefore + // TLS 1.0 is represented by "3,1", TLS 1.1 is "3,2", and so on. + buf.Write(tls.TLSVersion[:]) + + buf.Write(emptyRandom[:]) + + // 20 - 0x20 (32) bytes of session ID follow + // e0 e1 ... fe ff - session ID copied from Client Hello + buf.WriteByte(byte(len(hello.SessionID))) + buf.Write(hello.SessionID) + + binary.Write(buf, binary.BigEndian, hello.CipherSuite) //nolint: errcheck + + buf.Write(serverHelloSuffix) + + scalar := [EllipticCurveLen]byte{} + + if _, err := rand.Read(scalar[:]); err != nil { + panic(err) + } + + curve, _ := curve25519.X25519(scalar[:], curve25519.Basepoint) + buf.Write(curve) +} + +func generateChangeCipherValue(buf *bytes.Buffer) { + buf.WriteByte(tls.TypeChangeCipherSpec) + buf.Write(tls.TLSVersion[:]) + binary.Write(buf, binary.BigEndian, uint16(1)) //nolint: errcheck + buf.WriteByte(ChangeCipherValue) +} + +func generateNoise(buf *bytes.Buffer) { + data := [1369]byte{} + + if _, err := rand.Read(data[:]); err != nil { + panic(err) + } + + tls.WriteRecord(buf, data[:]) //nolint: errcheck +} diff --git a/mtglib/internal/tls/fake/server_side_test.go b/mtglib/internal/tls/fake/server_side_test.go new file mode 100644 index 0000000..05508ba --- /dev/null +++ b/mtglib/internal/tls/fake/server_side_test.go @@ -0,0 +1,131 @@ +package fake_test + +import ( + "bytes" + "crypto/hmac" + "crypto/rand" + "crypto/sha256" + "testing" + + "github.com/9seconds/mtg/v2/mtglib" + "github.com/9seconds/mtg/v2/mtglib/internal/tls" + "github.com/9seconds/mtg/v2/mtglib/internal/tls/fake" + "github.com/stretchr/testify/suite" +) + +type SendServerHelloTestSuite struct { + suite.Suite + + hello *fake.ClientHello + buf *bytes.Buffer + secret mtglib.Secret +} + +func (suite *SendServerHelloTestSuite) SetupTest() { + suite.hello = &fake.ClientHello{ + CipherSuite: 4867, + SessionID: make([]byte, 32), + } + + _, err := rand.Read(suite.hello.SessionID) + suite.NoError(err) + + _, err = rand.Read(suite.hello.Random[:]) + suite.NoError(err) + + suite.buf = &bytes.Buffer{} + suite.secret = mtglib.GenerateSecret("google.com") +} + +func (suite *SendServerHelloTestSuite) TestRecordStructure() { + noise, err := fake.SendServerHello(suite.buf, suite.secret, suite.hello) + suite.NoError(err) + + var rec bytes.Buffer + + recordType, _, err := tls.ReadRecord(suite.buf, &rec) + suite.NoError(err) + suite.Equal(byte(tls.TypeHandshake), recordType) + + rec.Reset() + + recordType, _, err = tls.ReadRecord(suite.buf, &rec) + suite.NoError(err) + suite.Equal(byte(tls.TypeChangeCipherSpec), recordType) + + suite.Empty(suite.buf.Bytes()) + + noiseBuf := bytes.NewReader(noise) + rec.Reset() + + recordType, _, err = tls.ReadRecord(noiseBuf, &rec) + suite.NoError(err) + suite.Equal(byte(tls.TypeApplicationData), recordType) + suite.Zero(noiseBuf.Len()) +} + +func (suite *SendServerHelloTestSuite) TestHMAC() { + noise, err := fake.SendServerHello(suite.buf, suite.secret, suite.hello) + suite.NoError(err) + + packet := make([]byte, suite.buf.Len()) + copy(packet, suite.buf.Bytes()) + + random := make([]byte, fake.RandomLen) + copy(random, packet[fake.RandomOffset:]) + copy(packet[fake.RandomOffset:], make([]byte, fake.RandomLen)) + + mac := hmac.New(sha256.New, suite.secret.Key[:]) + mac.Write(suite.hello.Random[:]) + mac.Write(packet) + mac.Write(noise) + + suite.Equal(random, mac.Sum(nil)) +} + +func (suite *SendServerHelloTestSuite) TestHandshakePayload() { + _, err := fake.SendServerHello(suite.buf, suite.secret, suite.hello) + suite.NoError(err) + + packet := suite.buf.Bytes() + + // TLS record header: type(1) + version(2) + length(2) + suite.Equal(byte(tls.TypeHandshake), packet[0]) + suite.Equal([]byte{3, 3}, packet[1:3]) + + // Handshake header: type(1) + uint24_length(3) + suite.Equal(byte(fake.TypeHandshakeServer), packet[5]) + + // ServerHello version + suite.Equal([]byte{3, 3}, packet[9:11]) + + // Session ID + sessionIDOffset := fake.RandomOffset + fake.RandomLen + suite.Equal(byte(len(suite.hello.SessionID)), packet[sessionIDOffset]) + suite.Equal(suite.hello.SessionID, packet[sessionIDOffset+1:sessionIDOffset+1+len(suite.hello.SessionID)]) +} + +func (suite *SendServerHelloTestSuite) TestChangeCipherSpec() { + _, err := fake.SendServerHello(suite.buf, suite.secret, suite.hello) + suite.NoError(err) + + // Skip first record + var rec bytes.Buffer + + _, _, err = tls.ReadRecord(suite.buf, &rec) + suite.NoError(err) + + // Read ChangeCipherSpec record + rec.Reset() + + recordType, length, err := tls.ReadRecord(suite.buf, &rec) + suite.NoError(err) + suite.Equal(byte(tls.TypeChangeCipherSpec), recordType) + suite.Equal(int64(1), length) + suite.Equal([]byte{fake.ChangeCipherValue}, rec.Bytes()) +} + +func TestSendServerHello(t *testing.T) { + t.Parallel() + suite.Run(t, &SendServerHelloTestSuite{}) +}