mirror of
https://github.com/ScuroNeko/mtg.git
synced 2026-08-31 16:14:02 +03:00
Add doppel and tls packages
This commit is contained in:
@@ -0,0 +1,88 @@
|
||||
package tls
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
|
||||
"github.com/9seconds/mtg/v2/essentials"
|
||||
)
|
||||
|
||||
const (
|
||||
SizeRecordType = 1
|
||||
SizeVersion = 2
|
||||
SizeSize = 2
|
||||
SizeHeader = SizeRecordType + SizeVersion + SizeSize
|
||||
|
||||
MaxRecordSize = 16384
|
||||
MaxRecordPayloadSize = MaxRecordSize - SizeHeader
|
||||
DefaultBufferSize = 4096
|
||||
|
||||
TypeChangeCipherSpec = 0x14
|
||||
TypeHandshake = 0x16
|
||||
TypeApplicationData = 0x17
|
||||
)
|
||||
|
||||
var (
|
||||
// TLS 1.2 is used for both TLS 1.2 and 1.3
|
||||
TLSVersion = [SizeVersion]byte{3, 3}
|
||||
)
|
||||
|
||||
// Conn presents an established TLS 1.3 connection, after handshake
|
||||
type Conn struct {
|
||||
essentials.Conn
|
||||
|
||||
p *connPayload
|
||||
}
|
||||
|
||||
type connPayload struct {
|
||||
readBuf bytes.Buffer
|
||||
writeBuf bytes.Buffer
|
||||
connBuffered *bufio.Reader
|
||||
read bool
|
||||
write bool
|
||||
}
|
||||
|
||||
func (c Conn) Write(p []byte) (int, error) {
|
||||
if !c.p.write {
|
||||
return c.Conn.Write(p)
|
||||
}
|
||||
|
||||
return len(p), WriteRecord(c.Conn, p)
|
||||
}
|
||||
|
||||
func (c Conn) Read(p []byte) (int, error) {
|
||||
if !c.p.read {
|
||||
return c.Conn.Read(p)
|
||||
}
|
||||
|
||||
for {
|
||||
if n, err := c.p.readBuf.Read(p); err == nil {
|
||||
return n, nil
|
||||
}
|
||||
|
||||
recordType, _, err := ReadRecord(c.p.connBuffered, &c.p.readBuf)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
if recordType != TypeApplicationData {
|
||||
c.p.readBuf.Reset()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func New(conn essentials.Conn, read, write bool) Conn {
|
||||
newConn := Conn{
|
||||
Conn: conn,
|
||||
p: &connPayload{
|
||||
connBuffered: bufio.NewReaderSize(conn, DefaultBufferSize),
|
||||
read: read,
|
||||
write: write,
|
||||
},
|
||||
}
|
||||
|
||||
newConn.p.readBuf.Grow(DefaultBufferSize)
|
||||
newConn.p.writeBuf.Grow(DefaultBufferSize)
|
||||
|
||||
return newConn
|
||||
}
|
||||
@@ -0,0 +1,160 @@
|
||||
package tls
|
||||
|
||||
import (
|
||||
"io"
|
||||
"testing"
|
||||
|
||||
"github.com/9seconds/mtg/v2/internal/testlib"
|
||||
"github.com/stretchr/testify/mock"
|
||||
"github.com/stretchr/testify/suite"
|
||||
)
|
||||
|
||||
type ConnTestSuite struct {
|
||||
suite.Suite
|
||||
|
||||
connMock *testlib.EssentialsConnMock
|
||||
}
|
||||
|
||||
func (suite *ConnTestSuite) SetupTest() {
|
||||
suite.connMock = &testlib.EssentialsConnMock{}
|
||||
}
|
||||
|
||||
func (suite *ConnTestSuite) TearDownTest() {
|
||||
suite.connMock.AssertExpectations(suite.T())
|
||||
}
|
||||
|
||||
func (suite *ConnTestSuite) feedRead(raw []byte) {
|
||||
suite.connMock.
|
||||
On("Read", mock.AnythingOfType("[]uint8")).
|
||||
Run(func(args mock.Arguments) {
|
||||
copy(args.Get(0).([]byte), raw)
|
||||
}).
|
||||
Return(len(raw), nil).
|
||||
Once()
|
||||
suite.connMock.
|
||||
On("Read", mock.AnythingOfType("[]uint8")).
|
||||
Return(0, io.EOF).
|
||||
Maybe()
|
||||
}
|
||||
|
||||
func (suite *ConnTestSuite) TestReadTLSEnabled() {
|
||||
payload := []byte("hello world")
|
||||
suite.feedRead(MakeTLSRecord(0x17, payload))
|
||||
|
||||
conn := New(suite.connMock, true, false)
|
||||
|
||||
buf := make([]byte, 128)
|
||||
n, err := conn.Read(buf)
|
||||
|
||||
suite.NoError(err)
|
||||
suite.Equal(payload, buf[:n])
|
||||
}
|
||||
|
||||
func (suite *ConnTestSuite) TestReadTLSSkipsNonApplicationData() {
|
||||
raw := append(
|
||||
MakeTLSRecord(0x14, []byte{1}),
|
||||
MakeTLSRecord(0x17, []byte("real data"))...,
|
||||
)
|
||||
suite.feedRead(raw)
|
||||
|
||||
conn := New(suite.connMock, true, false)
|
||||
|
||||
buf := make([]byte, 128)
|
||||
n, err := conn.Read(buf)
|
||||
|
||||
suite.NoError(err)
|
||||
suite.Equal([]byte("real data"), buf[:n])
|
||||
}
|
||||
|
||||
func (suite *ConnTestSuite) TestReadTLSMultipleRecords() {
|
||||
raw := append(
|
||||
MakeTLSRecord(0x17, []byte("first")),
|
||||
MakeTLSRecord(0x17, []byte("second"))...,
|
||||
)
|
||||
suite.feedRead(raw)
|
||||
|
||||
conn := New(suite.connMock, true, false)
|
||||
buf := make([]byte, 128)
|
||||
|
||||
n, err := conn.Read(buf)
|
||||
suite.NoError(err)
|
||||
suite.Equal([]byte("first"), buf[:n])
|
||||
|
||||
n, err = conn.Read(buf)
|
||||
suite.NoError(err)
|
||||
suite.Equal([]byte("second"), buf[:n])
|
||||
}
|
||||
|
||||
func (suite *ConnTestSuite) TestReadTLSSmallBuffer() {
|
||||
payload := []byte("hello world, this is a longer payload")
|
||||
suite.feedRead(MakeTLSRecord(0x17, payload))
|
||||
|
||||
conn := New(suite.connMock, true, false)
|
||||
|
||||
small := make([]byte, 5)
|
||||
n, err := conn.Read(small)
|
||||
suite.NoError(err)
|
||||
suite.Equal(payload[:5], small[:n])
|
||||
|
||||
rest := make([]byte, 128)
|
||||
n, err = conn.Read(rest)
|
||||
suite.NoError(err)
|
||||
suite.Equal(payload[5:], rest[:n])
|
||||
}
|
||||
|
||||
func (suite *ConnTestSuite) TestReadPassthrough() {
|
||||
data := []byte("raw bytes")
|
||||
|
||||
suite.connMock.
|
||||
On("Read", mock.AnythingOfType("[]uint8")).
|
||||
Run(func(args mock.Arguments) {
|
||||
copy(args.Get(0).([]byte), data)
|
||||
}).
|
||||
Return(len(data), nil).
|
||||
Once()
|
||||
|
||||
conn := New(suite.connMock, false, false)
|
||||
|
||||
buf := make([]byte, 128)
|
||||
n, err := conn.Read(buf)
|
||||
|
||||
suite.NoError(err)
|
||||
suite.Equal(data, buf[:n])
|
||||
}
|
||||
|
||||
func (suite *ConnTestSuite) TestWritePassthrough() {
|
||||
data := []byte("outgoing data")
|
||||
|
||||
suite.connMock.
|
||||
On("Write", mock.AnythingOfType("[]uint8")).
|
||||
Return(len(data), nil).
|
||||
Once()
|
||||
|
||||
conn := New(suite.connMock, false, false)
|
||||
|
||||
n, err := conn.Write(data)
|
||||
|
||||
suite.NoError(err)
|
||||
suite.Equal(len(data), n)
|
||||
}
|
||||
|
||||
func (suite *ConnTestSuite) TestWriteTLSEnabled() {
|
||||
data := []byte("outgoing data")
|
||||
|
||||
suite.connMock.
|
||||
On("Write", mock.AnythingOfType("[]uint8")).
|
||||
Return(len(data), nil).
|
||||
Once()
|
||||
|
||||
conn := New(suite.connMock, false, true)
|
||||
|
||||
n, err := conn.Write(data)
|
||||
|
||||
suite.NoError(err)
|
||||
suite.Equal(len(data), n)
|
||||
}
|
||||
|
||||
func TestConn(t *testing.T) {
|
||||
t.Parallel()
|
||||
suite.Run(t, &ConnTestSuite{})
|
||||
}
|
||||
@@ -0,0 +1,30 @@
|
||||
package tls
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
|
||||
"github.com/stretchr/testify/mock"
|
||||
)
|
||||
|
||||
type WriterMock struct {
|
||||
mock.Mock
|
||||
}
|
||||
|
||||
func (m *WriterMock) Write(p []byte) (int, error) {
|
||||
args := m.Called(p)
|
||||
return args.Int(0), args.Error(1)
|
||||
}
|
||||
|
||||
// makeTLSRecord builds a raw TLS record from hardcoded offsets:
|
||||
// type(1) + version(2, {3,3}) + length(2, big-endian) + payload.
|
||||
func MakeTLSRecord(recordType byte, payload []byte) []byte {
|
||||
buf := make([]byte, 5+len(payload))
|
||||
|
||||
buf[0] = recordType
|
||||
buf[1] = 3
|
||||
buf[2] = 3
|
||||
binary.BigEndian.PutUint16(buf[3:5], uint16(len(payload)))
|
||||
copy(buf[5:], payload)
|
||||
|
||||
return buf
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
package tls
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"io"
|
||||
)
|
||||
|
||||
func ReadRecord(r io.Reader, w io.Writer) (byte, int64, error) {
|
||||
buf := [SizeHeader]byte{}
|
||||
|
||||
if _, err := io.ReadFull(r, buf[:]); err != nil {
|
||||
return 0, 0, err
|
||||
}
|
||||
|
||||
pVer := buf[SizeRecordType:]
|
||||
pLen := pVer[SizeVersion:]
|
||||
|
||||
if !bytes.Equal(TLSVersion[:], pVer[:SizeVersion]) {
|
||||
return 0, 0, fmt.Errorf("incorrect tls version %v", pVer)
|
||||
}
|
||||
|
||||
length := int64(binary.BigEndian.Uint16(pLen[:SizeSize]))
|
||||
_, err := io.CopyN(w, r, length)
|
||||
|
||||
return buf[0], length, err
|
||||
}
|
||||
|
||||
func WriteRecord(w io.Writer, payload []byte) error {
|
||||
buf := [MaxRecordSize]byte{}
|
||||
buf[0] = TypeApplicationData
|
||||
|
||||
bufV := buf[SizeRecordType:]
|
||||
copy(bufV[:SizeVersion], TLSVersion[:])
|
||||
|
||||
bufS := bufV[SizeVersion:]
|
||||
binary.BigEndian.PutUint16(bufS[:SizeSize], uint16(len(payload)))
|
||||
|
||||
bufP := buf[SizeHeader:]
|
||||
if n := copy(bufP, payload); n != len(payload) {
|
||||
return fmt.Errorf("copied %d bytes of payload instead of %d", n, len(payload))
|
||||
}
|
||||
|
||||
_, err := w.Write(buf[:SizeHeader+len(payload)])
|
||||
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,125 @@
|
||||
package tls
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/mock"
|
||||
"github.com/stretchr/testify/suite"
|
||||
)
|
||||
|
||||
type UtilsTestSuite struct {
|
||||
suite.Suite
|
||||
|
||||
dst *bytes.Buffer
|
||||
}
|
||||
|
||||
func (suite *UtilsTestSuite) SetupTest() {
|
||||
suite.dst = &bytes.Buffer{}
|
||||
}
|
||||
|
||||
func (suite *UtilsTestSuite) TestReadRecord() {
|
||||
payload := []byte("hello world")
|
||||
raw := MakeTLSRecord(0x17, payload)
|
||||
|
||||
recordType, length, err := ReadRecord(bytes.NewReader(raw), suite.dst)
|
||||
|
||||
suite.NoError(err)
|
||||
suite.Equal(byte(0x17), recordType)
|
||||
suite.Equal(int64(len(payload)), length)
|
||||
suite.Equal(payload, suite.dst.Bytes())
|
||||
}
|
||||
|
||||
func (suite *UtilsTestSuite) TestReadRecordChangeCipherSpec() {
|
||||
payload := []byte{1}
|
||||
raw := MakeTLSRecord(0x14, payload)
|
||||
|
||||
recordType, length, err := ReadRecord(bytes.NewReader(raw), suite.dst)
|
||||
|
||||
suite.NoError(err)
|
||||
suite.Equal(byte(0x14), recordType)
|
||||
suite.Equal(int64(1), length)
|
||||
}
|
||||
|
||||
func (suite *UtilsTestSuite) TestReadRecordRejectsWrongVersion() {
|
||||
record := []byte{0x17, 3, 1, 0, 5, 0, 0, 0, 0, 0}
|
||||
|
||||
_, _, err := ReadRecord(bytes.NewReader(record), suite.dst)
|
||||
suite.ErrorContains(err, "incorrect tls version")
|
||||
}
|
||||
|
||||
func (suite *UtilsTestSuite) TestReadRecordEmptyReader() {
|
||||
_, _, err := ReadRecord(bytes.NewReader(nil), suite.dst)
|
||||
suite.Error(err)
|
||||
}
|
||||
|
||||
func (suite *UtilsTestSuite) TestReadRecordTruncatedHeader() {
|
||||
_, _, err := ReadRecord(bytes.NewReader([]byte{0x17, 3}), suite.dst)
|
||||
suite.Error(err)
|
||||
}
|
||||
|
||||
func (suite *UtilsTestSuite) TestReadRecordTruncatedPayload() {
|
||||
raw := MakeTLSRecord(0x17, []byte("full payload"))
|
||||
truncated := raw[:5+3]
|
||||
|
||||
_, _, err := ReadRecord(bytes.NewReader(truncated), suite.dst)
|
||||
suite.Error(err)
|
||||
}
|
||||
|
||||
func (suite *UtilsTestSuite) TestWriteRecord() {
|
||||
payload := []byte("hello world")
|
||||
|
||||
err := WriteRecord(suite.dst, payload)
|
||||
suite.NoError(err)
|
||||
|
||||
written := suite.dst.Bytes()
|
||||
suite.Equal(byte(0x17), written[0])
|
||||
suite.Equal([]byte{3, 3}, written[1:3])
|
||||
|
||||
length := binary.BigEndian.Uint16(written[3:5])
|
||||
suite.Equal(uint16(len(payload)), length)
|
||||
suite.Equal(payload, written[5:])
|
||||
}
|
||||
|
||||
func (suite *UtilsTestSuite) TestWriteRecordRoundTrip() {
|
||||
payload := []byte("round trip test")
|
||||
|
||||
var wire bytes.Buffer
|
||||
|
||||
err := WriteRecord(&wire, payload)
|
||||
suite.NoError(err)
|
||||
|
||||
var recovered bytes.Buffer
|
||||
|
||||
recordType, length, err := ReadRecord(&wire, &recovered)
|
||||
|
||||
suite.NoError(err)
|
||||
suite.Equal(byte(0x17), recordType)
|
||||
suite.Equal(int64(len(payload)), length)
|
||||
suite.Equal(payload, recovered.Bytes())
|
||||
}
|
||||
|
||||
func (suite *UtilsTestSuite) TestWriteRecordPropagatesError() {
|
||||
m := &WriterMock{}
|
||||
m.
|
||||
On("Write", mock.AnythingOfType("[]uint8")).
|
||||
Once().
|
||||
Return(0, errors.New("dist full"))
|
||||
|
||||
err := WriteRecord(m, []byte("data"))
|
||||
suite.Error(err)
|
||||
|
||||
m.AssertExpectations(suite.T())
|
||||
}
|
||||
|
||||
func (suite *UtilsTestSuite) TestWriteRecordPayloadTooLarge() {
|
||||
err := WriteRecord(suite.dst, make([]byte, MaxRecordPayloadSize+1))
|
||||
suite.Error(err)
|
||||
}
|
||||
|
||||
func TestUtils(t *testing.T) {
|
||||
t.Parallel()
|
||||
suite.Run(t, &UtilsTestSuite{})
|
||||
}
|
||||
Reference in New Issue
Block a user