mirror of
https://github.com/ScuroNeko/mtg.git
synced 2026-08-31 12:54:01 +03:00
Merge pull request #234 from 9seconds/cant-send-attachment
This commit is contained in:
@@ -8,14 +8,12 @@ import (
|
|||||||
// CloseableReader is a reader interface that can close its reading end.
|
// CloseableReader is a reader interface that can close its reading end.
|
||||||
type CloseableReader interface {
|
type CloseableReader interface {
|
||||||
io.Reader
|
io.Reader
|
||||||
|
|
||||||
CloseRead() error
|
CloseRead() error
|
||||||
}
|
}
|
||||||
|
|
||||||
// CloseableWriter is a writer that can close its writing end.
|
// CloseableWriter is a writer that can close its writing end.
|
||||||
type CloseableWriter interface {
|
type CloseableWriter interface {
|
||||||
io.Writer
|
io.Writer
|
||||||
|
|
||||||
CloseWrite() error
|
CloseWrite() error
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,11 +1,7 @@
|
|||||||
package relay
|
package relay
|
||||||
|
|
||||||
import "time"
|
|
||||||
|
|
||||||
const (
|
const (
|
||||||
copyBufferSize = 64 * 1024
|
copyBufferSize = 64 * 1024
|
||||||
writerBufferSize = 128 * 1024
|
|
||||||
readTimeout = 10 * time.Millisecond
|
|
||||||
)
|
)
|
||||||
|
|
||||||
type Logger interface {
|
type Logger interface {
|
||||||
|
|||||||
@@ -1,31 +1,19 @@
|
|||||||
package relay
|
package relay
|
||||||
|
|
||||||
import (
|
import "sync"
|
||||||
"bufio"
|
|
||||||
"io"
|
|
||||||
"net"
|
|
||||||
"sync"
|
|
||||||
)
|
|
||||||
|
|
||||||
var syncPairPool = sync.Pool{
|
var copyBufferPool = sync.Pool{
|
||||||
New: func() interface{} {
|
New: func() interface{} {
|
||||||
return &syncPair{
|
rv := make([]byte, copyBufferSize)
|
||||||
writer: bufio.NewWriterSize(nil, writerBufferSize),
|
|
||||||
copyBuf: make([]byte, copyBufferSize),
|
return &rv
|
||||||
}
|
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
func acquireSyncPair(reader net.Conn, writer io.Writer) *syncPair {
|
func acquireCopyBuffer() *[]byte {
|
||||||
sp := syncPairPool.Get().(*syncPair) // nolint: forcetypeassert
|
return copyBufferPool.Get().(*[]byte)
|
||||||
sp.writer.Reset(writer)
|
|
||||||
sp.reader = reader
|
|
||||||
|
|
||||||
return sp
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func releaseSyncPair(sp *syncPair) {
|
func releaseCopyBuffer(buf *[]byte) {
|
||||||
sp.writer.Reset(nil)
|
copyBufferPool.Put(buf)
|
||||||
sp.reader = nil
|
|
||||||
syncPairPool.Put(sp)
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,7 +4,6 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"io"
|
"io"
|
||||||
"sync"
|
|
||||||
|
|
||||||
"github.com/9seconds/mtg/v2/essentials"
|
"github.com/9seconds/mtg/v2/essentials"
|
||||||
)
|
)
|
||||||
@@ -22,28 +21,27 @@ func Relay(ctx context.Context, log Logger, telegramConn, clientConn essentials.
|
|||||||
clientConn.Close()
|
clientConn.Close()
|
||||||
}()
|
}()
|
||||||
|
|
||||||
wg := &sync.WaitGroup{}
|
closeChan := make(chan struct{})
|
||||||
wg.Add(2) // nolint: gomnd
|
|
||||||
|
|
||||||
go pump(log, telegramConn, clientConn, wg, "client -> telegram")
|
go func() {
|
||||||
|
defer close(closeChan)
|
||||||
|
|
||||||
pump(log, clientConn, telegramConn, wg, "telegram -> client")
|
pump(log, telegramConn, clientConn, "client -> telegram")
|
||||||
|
|
||||||
wg.Wait()
|
|
||||||
}
|
|
||||||
|
|
||||||
func pump(log Logger, src, dst essentials.Conn, wg *sync.WaitGroup, direction string) {
|
|
||||||
syncer := acquireSyncPair(src, dst)
|
|
||||||
|
|
||||||
defer func() {
|
|
||||||
syncer.Flush()
|
|
||||||
releaseSyncPair(syncer)
|
|
||||||
src.CloseRead() // nolint: errcheck
|
|
||||||
dst.CloseWrite() // nolint: errcheck
|
|
||||||
wg.Done()
|
|
||||||
}()
|
}()
|
||||||
|
|
||||||
n, err := syncer.Sync()
|
pump(log, clientConn, telegramConn, "telegram -> client")
|
||||||
|
|
||||||
|
<-closeChan
|
||||||
|
}
|
||||||
|
|
||||||
|
func pump(log Logger, src, dst essentials.Conn, direction string) {
|
||||||
|
defer src.CloseRead() // nolint: errcheck
|
||||||
|
defer dst.CloseWrite() // nolint: errcheck
|
||||||
|
|
||||||
|
copyBuffer := acquireCopyBuffer()
|
||||||
|
defer releaseCopyBuffer(copyBuffer)
|
||||||
|
|
||||||
|
n, err := io.CopyBuffer(src, dst, *copyBuffer)
|
||||||
|
|
||||||
switch {
|
switch {
|
||||||
case err == nil:
|
case err == nil:
|
||||||
|
|||||||
@@ -42,14 +42,12 @@ func (suite *RelayTestSuite) TestExit() {
|
|||||||
suite.telegramConnMock.On("CloseWrite").Return(nil).Once()
|
suite.telegramConnMock.On("CloseWrite").Return(nil).Once()
|
||||||
suite.telegramConnMock.On("Read", mock.Anything).Return(10, io.EOF).Once()
|
suite.telegramConnMock.On("Read", mock.Anything).Return(10, io.EOF).Once()
|
||||||
suite.telegramConnMock.On("Write", mock.Anything).Return(10, io.EOF).Maybe()
|
suite.telegramConnMock.On("Write", mock.Anything).Return(10, io.EOF).Maybe()
|
||||||
suite.telegramConnMock.On("SetReadDeadline", mock.Anything).Return(nil).Maybe()
|
|
||||||
|
|
||||||
suite.clientConnMock.On("Read", mock.Anything).Return(0, io.EOF).Once()
|
suite.clientConnMock.On("Read", mock.Anything).Return(0, io.EOF).Once()
|
||||||
suite.clientConnMock.On("Write", mock.Anything).Return(10, io.EOF).Maybe()
|
suite.clientConnMock.On("Write", mock.Anything).Return(10, io.EOF).Maybe()
|
||||||
suite.clientConnMock.On("Close").Return(nil)
|
suite.clientConnMock.On("Close").Return(nil)
|
||||||
suite.clientConnMock.On("CloseRead").Return(nil).Once()
|
suite.clientConnMock.On("CloseRead").Return(nil).Once()
|
||||||
suite.clientConnMock.On("CloseWrite").Return(nil).Once()
|
suite.clientConnMock.On("CloseWrite").Return(nil).Once()
|
||||||
suite.clientConnMock.On("SetReadDeadline", mock.Anything).Return(nil).Maybe()
|
|
||||||
|
|
||||||
relay.Relay(suite.ctx, suite.loggerMock, suite.telegramConnMock, suite.clientConnMock)
|
relay.Relay(suite.ctx, suite.loggerMock, suite.telegramConnMock, suite.clientConnMock)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,77 +0,0 @@
|
|||||||
package relay
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bufio"
|
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"net"
|
|
||||||
"os"
|
|
||||||
"sync"
|
|
||||||
"time"
|
|
||||||
)
|
|
||||||
|
|
||||||
type syncPair struct {
|
|
||||||
writer *bufio.Writer
|
|
||||||
copyBuf []byte
|
|
||||||
|
|
||||||
mutex sync.Mutex
|
|
||||||
reader net.Conn
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *syncPair) Sync() (int64, error) {
|
|
||||||
return io.CopyBuffer(s, s, s.copyBuf) // nolint: wrapcheck
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *syncPair) Read(p []byte) (int, error) {
|
|
||||||
n, err := s.readBlocking(p, false)
|
|
||||||
|
|
||||||
// nothing has been delivered for readTimeout time. Let's flush.
|
|
||||||
if errors.Is(err, os.ErrDeadlineExceeded) {
|
|
||||||
if err := s.Flush(); err != nil {
|
|
||||||
return 0, fmt.Errorf("cannot flush writer hand-side: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return s.readBlocking(p, true)
|
|
||||||
}
|
|
||||||
|
|
||||||
return n, err
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *syncPair) Write(p []byte) (int, error) {
|
|
||||||
s.mutex.Lock()
|
|
||||||
defer s.mutex.Unlock()
|
|
||||||
|
|
||||||
n, err := s.writer.Write(p)
|
|
||||||
|
|
||||||
// optimization for a case when we have a small package and want to avoid a
|
|
||||||
// delay in readTimeout. In that case, we assume that peer has finished to
|
|
||||||
// sent a data it wants to send so we can flush without waiting for anything
|
|
||||||
// else.
|
|
||||||
if err == nil && n < copyBufferSize {
|
|
||||||
err = s.writer.Flush()
|
|
||||||
}
|
|
||||||
|
|
||||||
return n, err // nolint: wrapcheck
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *syncPair) Flush() error {
|
|
||||||
s.mutex.Lock()
|
|
||||||
defer s.mutex.Unlock()
|
|
||||||
|
|
||||||
return s.writer.Flush() // nolint: wrapcheck
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *syncPair) readBlocking(p []byte, blocking bool) (int, error) {
|
|
||||||
var deadline time.Time
|
|
||||||
|
|
||||||
if !blocking {
|
|
||||||
deadline = time.Now().Add(readTimeout)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := s.reader.SetReadDeadline(deadline); err != nil {
|
|
||||||
return 0, fmt.Errorf("cannot set read deadline: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return s.reader.Read(p) // nolint: wrapcheck
|
|
||||||
}
|
|
||||||
Reference in New Issue
Block a user