From 33c0fa9bf725acf57b7af51b5f7996a665251063 Mon Sep 17 00:00:00 2001 From: 9seconds Date: Fri, 13 Mar 2026 11:30:52 +0100 Subject: [PATCH] Add SyncWrite method to doppel.Conn --- mtglib/internal/doppel/conn.go | 55 +++++++++--- mtglib/internal/doppel/conn_test.go | 130 ++++++++++++++++++++++++++++ 2 files changed, 175 insertions(+), 10 deletions(-) diff --git a/mtglib/internal/doppel/conn.go b/mtglib/internal/doppel/conn.go index a2543cd..7e8ed30 100644 --- a/mtglib/internal/doppel/conn.go +++ b/mtglib/internal/doppel/conn.go @@ -16,18 +16,48 @@ type Conn struct { } type connPayload struct { - ctx context.Context - ctxCancel context.CancelCauseFunc - clock Clock - wg sync.WaitGroup - writeLock sync.Mutex - writeStream bytes.Buffer + ctx context.Context + ctxCancel context.CancelCauseFunc + clock Clock + wg sync.WaitGroup + syncWriteLock sync.RWMutex + writeStream bytes.Buffer + writeCond *sync.Cond } func (c Conn) Write(p []byte) (int, error) { - c.p.writeLock.Lock() + c.p.syncWriteLock.RLock() + defer c.p.syncWriteLock.RUnlock() + + c.p.writeCond.L.Lock() c.p.writeStream.Write(p) - c.p.writeLock.Unlock() + c.p.writeCond.L.Unlock() + + return len(p), context.Cause(c.p.ctx) +} + +func (c Conn) SyncWrite(p []byte) (int, error) { + c.p.syncWriteLock.Lock() + defer c.p.syncWriteLock.Unlock() + + c.p.writeCond.L.Lock() + // wait until buffer is exhausted + for c.p.writeStream.Len() != 0 && context.Cause(c.p.ctx) == nil { + c.p.writeCond.Wait() + } + c.p.writeStream.Write(p) + c.p.writeCond.L.Unlock() + + if err := context.Cause(c.p.ctx); err != nil { + return len(p), err + } + + c.p.writeCond.L.Lock() + // wait until data will be sent + for c.p.writeStream.Len() != 0 && context.Cause(c.p.ctx) == nil { + c.p.writeCond.Wait() + } + c.p.writeCond.L.Unlock() return len(p), context.Cause(c.p.ctx) } @@ -39,6 +69,8 @@ func (c Conn) Start() { } func (c Conn) start() { + defer c.p.writeCond.Broadcast() + buf := [tls.MaxRecordSize]byte{} for { @@ -48,9 +80,9 @@ func (c Conn) start() { case <-c.p.clock.tick: } - c.p.writeLock.Lock() + c.p.writeCond.L.Lock() n, err := c.p.writeStream.Read(buf[:c.p.clock.stats.Size()]) - c.p.writeLock.Unlock() + c.p.writeCond.L.Unlock() if n == 0 || err != nil { continue @@ -60,6 +92,8 @@ func (c Conn) start() { c.p.ctxCancel(err) return } + + c.p.writeCond.Signal() } } @@ -75,6 +109,7 @@ func NewConn(ctx context.Context, conn essentials.Conn, stats *Stats) Conn { p: &connPayload{ ctx: ctx, ctxCancel: cancel, + writeCond: sync.NewCond(&sync.Mutex{}), clock: Clock{ stats: stats, tick: make(chan struct{}), diff --git a/mtglib/internal/doppel/conn_test.go b/mtglib/internal/doppel/conn_test.go index e774bce..f8adb91 100644 --- a/mtglib/internal/doppel/conn_test.go +++ b/mtglib/internal/doppel/conn_test.go @@ -157,6 +157,136 @@ func (suite *ConnTestSuite) TestStopOnUnderlyingWriteError() { }, 2*time.Second, time.Millisecond) } +func (suite *ConnTestSuite) TestSyncWriteDataSent() { + suite.connMock. + On("Write", mock.AnythingOfType("[]uint8")). + Return(0, nil). + Maybe() + + c := suite.makeConn() + defer c.Stop() + + payload := []byte("sync hello") + n, err := c.SyncWrite(payload) + suite.NoError(err) + suite.Equal(len(payload), n) + + // SyncWrite returns only after data is flushed to the wire. + assembled := &bytes.Buffer{} + reader := bytes.NewReader(suite.connMock.Written()) + + for { + header := make([]byte, tls.SizeHeader) + if _, err := io.ReadFull(reader, header); err != nil { + break + } + + suite.Equal(byte(tls.TypeApplicationData), header[0]) + + length := binary.BigEndian.Uint16(header[tls.SizeRecordType+tls.SizeVersion:]) + rec := make([]byte, length) + _, err := io.ReadFull(reader, rec) + suite.NoError(err) + + assembled.Write(rec) + } + + suite.Equal(payload, assembled.Bytes()) +} + +func (suite *ConnTestSuite) TestSyncWriteDrainsBufferFirst() { + suite.connMock. + On("Write", mock.AnythingOfType("[]uint8")). + Return(0, nil). + Maybe() + + c := suite.makeConn() + defer c.Stop() + + // Buffer some data via async Write. + _, err := c.Write([]byte("first")) + suite.NoError(err) + + // SyncWrite must drain "first" before sending "second". + n, err := c.SyncWrite([]byte("second")) + suite.NoError(err) + suite.Equal(6, n) + + // All data should be on the wire now. + assembled := &bytes.Buffer{} + reader := bytes.NewReader(suite.connMock.Written()) + + for { + header := make([]byte, tls.SizeHeader) + if _, err := io.ReadFull(reader, header); err != nil { + break + } + + length := binary.BigEndian.Uint16(header[tls.SizeRecordType+tls.SizeVersion:]) + rec := make([]byte, length) + _, err := io.ReadFull(reader, rec) + suite.NoError(err) + + assembled.Write(rec) + } + + suite.Equal([]byte("firstsecond"), assembled.Bytes()) +} + +func (suite *ConnTestSuite) TestSyncWriteBlocksAsyncWrite() { + suite.connMock. + On("Write", mock.AnythingOfType("[]uint8")). + Return(0, nil). + Maybe() + + c := suite.makeConn() + defer c.Stop() + + // Start SyncWrite — it holds exclusive lock. + syncDone := make(chan struct{}) + + go func() { + defer close(syncDone) + c.SyncWrite([]byte("exclusive")) + }() + + // Give SyncWrite time to acquire the lock. + time.Sleep(10 * time.Millisecond) + + // Async Write should block until SyncWrite completes. + writeDone := make(chan struct{}) + + go func() { + defer close(writeDone) + c.Write([]byte("blocked")) + }() + + // SyncWrite should finish first. + <-syncDone + + select { + case <-writeDone: + // Write completed after SyncWrite — correct. + case <-time.After(2 * time.Second): + suite.Fail("async Write did not unblock after SyncWrite completed") + } +} + +func (suite *ConnTestSuite) TestSyncWriteReturnsErrorAfterStop() { + suite.connMock. + On("Write", mock.AnythingOfType("[]uint8")). + Return(0, nil). + Maybe() + + c := suite.makeConn() + c.Stop() + + time.Sleep(10 * time.Millisecond) + + _, err := c.SyncWrite([]byte("too late")) + suite.Error(err) +} + func TestConn(t *testing.T) { t.Parallel() suite.Run(t, &ConnTestSuite{})