diff --git a/mtglib/internal/relay/init.go b/mtglib/internal/relay/init.go index 855d952..e54ac7a 100644 --- a/mtglib/internal/relay/init.go +++ b/mtglib/internal/relay/init.go @@ -1,7 +1,11 @@ package relay +import "time" + const ( - bufferSize = 32 * 1024 + copyBufferSize = 32 * 1024 + writerBufferSize = 2 * copyBufferSize + readTimeout = 10 * time.Millisecond ) type Logger interface { diff --git a/mtglib/internal/relay/pools.go b/mtglib/internal/relay/pools.go index 9a10527..2adff99 100644 --- a/mtglib/internal/relay/pools.go +++ b/mtglib/internal/relay/pools.go @@ -1,25 +1,31 @@ package relay -import "sync" +import ( + "bufio" + "io" + "net" + "sync" +) -type eastWest struct { - east []byte - west []byte -} - -var eastWestPool = sync.Pool{ +var syncPairPool = sync.Pool{ New: func() interface{} { - return &eastWest{ - east: make([]byte, bufferSize), - west: make([]byte, bufferSize), + return &syncPair{ + writer: bufio.NewWriterSize(nil, writerBufferSize), + copyBuf: make([]byte, copyBufferSize), } }, } -func acquireEastWest() *eastWest { - return eastWestPool.Get().(*eastWest) +func acquireSyncPair(reader net.Conn, writer io.Writer) *syncPair { + sp := syncPairPool.Get().(*syncPair) // nolint: forcetypeassert + sp.writer.Reset(writer) + sp.reader = reader + + return sp } -func releaseEastWest(ew *eastWest) { - eastWestPool.Put(ew) +func releaseSyncPair(sp *syncPair) { + sp.writer.Reset(nil) + sp.reader = nil + syncPairPool.Put(sp) } diff --git a/mtglib/internal/relay/relay.go b/mtglib/internal/relay/relay.go index 2dab380..294fa5d 100644 --- a/mtglib/internal/relay/relay.go +++ b/mtglib/internal/relay/relay.go @@ -2,14 +2,11 @@ package relay import ( "context" - "io" + "net" "sync" ) -func Relay(ctx context.Context, log Logger, telegramConn, clientConn io.ReadWriteCloser) { - defer telegramConn.Close() - defer clientConn.Close() - +func Relay(ctx context.Context, log Logger, telegramConn, clientConn net.Conn) { ctx, cancel := context.WithCancel(ctx) defer cancel() @@ -19,26 +16,25 @@ func Relay(ctx context.Context, log Logger, telegramConn, clientConn io.ReadWrit clientConn.Close() }() - buffers := acquireEastWest() - defer releaseEastWest(buffers) - wg := &sync.WaitGroup{} wg.Add(2) // nolint: gomnd - go pump(log, telegramConn, clientConn, wg, buffers.east, "east -> west") + go pump(log, telegramConn, clientConn, wg, "client -> telegram") - pump(log, clientConn, telegramConn, wg, buffers.west, "west -> east") + pump(log, clientConn, telegramConn, wg, "telegram -> client") wg.Wait() } -func pump(log Logger, src io.ReadCloser, dst io.WriteCloser, wg *sync.WaitGroup, - buf []byte, direction string) { +func pump(log Logger, src, dst net.Conn, wg *sync.WaitGroup, direction string) { defer wg.Done() defer src.Close() defer dst.Close() - if n, err := io.CopyBuffer(dst, src, buf); err != nil { + syncer := acquireSyncPair(src, dst) + defer releaseSyncPair(syncer) + + if n, err := syncer.Sync(); err != nil { log.Printf("cannot pump %s (written %d bytes): %w", direction, n, err) } } diff --git a/mtglib/internal/relay/sync_pair.go b/mtglib/internal/relay/sync_pair.go new file mode 100644 index 0000000..e83c58e --- /dev/null +++ b/mtglib/internal/relay/sync_pair.go @@ -0,0 +1,54 @@ +package relay + +import ( + "bufio" + "errors" + "fmt" + "io" + "net" + "os" + "time" +) + +type syncPair struct { + writer *bufio.Writer + copyBuf []byte + + 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) + + if errors.Is(err, os.ErrDeadlineExceeded) { + if err := s.writer.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) { + return s.writer.Write(p) // 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 +}