Add doppel and tls packages

This commit is contained in:
9seconds
2026-03-12 19:07:10 +01:00
parent c886ffdd81
commit 1182b9ef6f
21 changed files with 1867 additions and 0 deletions
+88
View File
@@ -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
}
+160
View File
@@ -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{})
}
+30
View File
@@ -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
}
+48
View File
@@ -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
}
+125
View File
@@ -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{})
}