Change algorithm of TCP relaying

This commit is contained in:
9seconds
2021-08-27 17:22:36 +03:00
parent 4b7be8c565
commit 456ed5b051
18 changed files with 300 additions and 434 deletions
+25 -104
View File
@@ -3,127 +3,48 @@ package relay
import (
"context"
"io"
"net"
"sync"
"time"
)
type Relay struct {
ctx context.Context
ctxCancel context.CancelFunc
logger Logger
processMutex sync.Mutex
eastBuffer []byte
westBuffer []byte
tickChannel chan struct{}
errorChannel chan error
tickTimeout time.Duration
}
func Relay(ctx context.Context, log Logger, bufferSize int,
telegramConn net.Conn, clientConn io.ReadWriteCloser) {
defer telegramConn.Close()
defer clientConn.Close()
func (r *Relay) Reset() {
r.processMutex.Lock()
defer r.processMutex.Unlock()
ctx, cancel := context.WithTimeout(ctx, getConnectionTimeToLive())
defer cancel()
if r.ctxCancel != nil {
r.ctxCancel()
}
go func() {
<-ctx.Done()
telegramConn.Close()
clientConn.Close()
}()
r.ctx = nil
r.ctxCancel = nil
r.logger = nil
}
buffers := acquireEastWest(bufferSize)
defer releaseEastWest(buffers)
func (r *Relay) Process(eastConn, westConn io.ReadWriteCloser) error {
r.processMutex.Lock()
defer r.processMutex.Unlock()
eastConn = conn{
ReadWriteCloser: eastConn,
ctx: r.ctx,
tickChannel: r.tickChannel,
}
westConn = conn{
ReadWriteCloser: westConn,
ctx: r.ctx,
tickChannel: r.tickChannel,
telegramConn = conn{
Conn: telegramConn,
}
wg := &sync.WaitGroup{}
wg.Add(3) // nolint: gomnd
wg.Add(2) // nolint: gomnd
go r.runObserver(eastConn, westConn, wg)
go pump(log, telegramConn, clientConn, wg, buffers.east, "east -> west")
go r.transmit(eastConn, westConn, r.westBuffer, "west", wg)
r.transmit(westConn, eastConn, r.eastBuffer, "east", wg)
pump(log, clientConn, telegramConn, wg, buffers.west, "west -> east")
wg.Wait()
select {
case err := <-r.errorChannel:
return err
default:
return nil
}
}
func (r *Relay) transmit(src io.ReadCloser, dst io.WriteCloser,
buffer []byte, direction string, wg *sync.WaitGroup) {
func pump(log Logger, src io.ReadCloser, dst io.WriteCloser, wg *sync.WaitGroup,
buf []byte, direction string) {
defer wg.Done()
defer src.Close()
defer dst.Close()
defer func() {
r.ctxCancel()
src.Close()
dst.Close()
}()
if _, err := io.CopyBuffer(dst, src, buffer); err != nil {
r.logger.Printf("error '%v' happened on direction %s", err, direction)
select {
case <-r.ctx.Done():
err = r.ctx.Err()
default:
}
select {
case r.errorChannel <- err:
default:
}
}
}
func (r *Relay) runObserver(one, another io.Closer, wg *sync.WaitGroup) {
defer wg.Done()
ticker := time.NewTicker(time.Second)
defer func() {
one.Close()
another.Close()
ticker.Stop()
select {
case <-ticker.C:
default:
}
}()
lastTickAt := time.Now()
for {
select {
case <-r.ctx.Done():
return
case <-r.tickChannel:
lastTickAt = time.Now()
case <-ticker.C:
if time.Since(lastTickAt) > r.tickTimeout {
r.logger.Printf("exit due to a timeout")
r.ctxCancel()
return
}
}
if n, err := io.CopyBuffer(dst, src, buf); err != nil {
log.Printf("cannot pump %s (written %d bytes): %w", direction, n, err)
}
}