mirror of
https://github.com/ScuroNeko/mtg.git
synced 2026-08-31 12:04:03 +03:00
Add SyncWrite method to doppel.Conn
This commit is contained in:
@@ -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{}),
|
||||
|
||||
@@ -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{})
|
||||
|
||||
Reference in New Issue
Block a user