212 lines
4.6 KiB
Go
212 lines
4.6 KiB
Go
package main
|
|
|
|
import (
|
|
"fmt"
|
|
"io"
|
|
"net"
|
|
"sync"
|
|
"time"
|
|
)
|
|
|
|
// sshCarrierConn absorbs the relatively small encrypted writes produced by the
|
|
// SSH packet layer and combines them before handing them to DragonTCP.
|
|
//
|
|
// This matters because a DragonTCP upload is a request/ack transaction. Without
|
|
// write combining, one SSH packet can become one full network round trip even
|
|
// when the discovered DragonTCP path supports much larger chunks. The wrapper
|
|
// deliberately behaves like a kernel socket send buffer: Write returns after
|
|
// the bytes have been copied into a bounded queue, while a single ordered
|
|
// writer drains that queue to the underlying DragonTCP stream.
|
|
type sshCarrierConn struct {
|
|
raw net.Conn
|
|
flushBytes int
|
|
maxBuffered int
|
|
flushDelay time.Duration
|
|
|
|
mu sync.Mutex
|
|
cond *sync.Cond
|
|
buf []byte
|
|
closing bool
|
|
writeErr error
|
|
coalesceLogged bool
|
|
done chan struct{}
|
|
closeOnce sync.Once
|
|
}
|
|
|
|
func newSSHCarrierConn(raw net.Conn, flushBytes, maxBuffered int, flushDelay time.Duration) net.Conn {
|
|
if raw == nil {
|
|
return nil
|
|
}
|
|
if flushBytes < 32*1024 {
|
|
flushBytes = 32 * 1024
|
|
}
|
|
if maxBuffered < flushBytes*2 {
|
|
maxBuffered = flushBytes * 2
|
|
}
|
|
if flushDelay <= 0 {
|
|
flushDelay = time.Millisecond
|
|
}
|
|
c := &sshCarrierConn{
|
|
raw: raw,
|
|
flushBytes: flushBytes,
|
|
maxBuffered: maxBuffered,
|
|
flushDelay: flushDelay,
|
|
done: make(chan struct{}),
|
|
}
|
|
c.cond = sync.NewCond(&c.mu)
|
|
go c.writeLoop()
|
|
return c
|
|
}
|
|
|
|
func (c *sshCarrierConn) writeLoop() {
|
|
defer close(c.done)
|
|
defer c.raw.Close()
|
|
|
|
for {
|
|
c.mu.Lock()
|
|
for len(c.buf) == 0 && !c.closing && c.writeErr == nil {
|
|
c.cond.Wait()
|
|
}
|
|
if c.writeErr != nil {
|
|
c.mu.Unlock()
|
|
return
|
|
}
|
|
if len(c.buf) == 0 && c.closing {
|
|
c.mu.Unlock()
|
|
return
|
|
}
|
|
shouldDelay := len(c.buf) < c.flushBytes && !c.closing
|
|
c.mu.Unlock()
|
|
|
|
// Give consecutive SSH packets a very small window to accumulate. The
|
|
// delay is tiny compared with a WAN RTT, but it lets 32 KiB SSH packets
|
|
// become a 256 KiB-1 MiB DragonTCP write on a busy stream.
|
|
if shouldDelay {
|
|
time.Sleep(c.flushDelay)
|
|
}
|
|
|
|
c.mu.Lock()
|
|
if c.writeErr != nil {
|
|
c.mu.Unlock()
|
|
return
|
|
}
|
|
n := len(c.buf)
|
|
if n > c.flushBytes {
|
|
n = c.flushBytes
|
|
}
|
|
batch := make([]byte, n)
|
|
copy(batch, c.buf[:n])
|
|
if n == len(c.buf) {
|
|
c.buf = c.buf[:0]
|
|
} else {
|
|
copy(c.buf, c.buf[n:])
|
|
c.buf = c.buf[:len(c.buf)-n]
|
|
}
|
|
c.cond.Broadcast()
|
|
c.mu.Unlock()
|
|
|
|
if !c.coalesceLogged && len(batch) >= 128*1024 {
|
|
c.coalesceLogged = true
|
|
fmt.Printf("ssh carrier: packet coalescing active batch=%d\n", len(batch))
|
|
}
|
|
|
|
if err := writeCarrierFull(c.raw, batch); err != nil {
|
|
c.mu.Lock()
|
|
if c.writeErr == nil {
|
|
c.writeErr = err
|
|
}
|
|
c.closing = true
|
|
c.cond.Broadcast()
|
|
c.mu.Unlock()
|
|
return
|
|
}
|
|
}
|
|
}
|
|
|
|
func writeCarrierFull(w io.Writer, p []byte) error {
|
|
for len(p) > 0 {
|
|
n, err := w.Write(p)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if n <= 0 {
|
|
return io.ErrShortWrite
|
|
}
|
|
p = p[n:]
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (c *sshCarrierConn) Read(p []byte) (int, error) {
|
|
n, err := c.raw.Read(p)
|
|
if n > 0 || err == nil {
|
|
return n, err
|
|
}
|
|
c.mu.Lock()
|
|
writeErr := c.writeErr
|
|
c.mu.Unlock()
|
|
if writeErr != nil {
|
|
return 0, writeErr
|
|
}
|
|
return n, err
|
|
}
|
|
|
|
func (c *sshCarrierConn) Write(p []byte) (int, error) {
|
|
if len(p) == 0 {
|
|
return 0, nil
|
|
}
|
|
written := 0
|
|
for written < len(p) {
|
|
c.mu.Lock()
|
|
for len(c.buf) >= c.maxBuffered && !c.closing && c.writeErr == nil {
|
|
c.cond.Wait()
|
|
}
|
|
if c.writeErr != nil {
|
|
err := c.writeErr
|
|
c.mu.Unlock()
|
|
return written, err
|
|
}
|
|
if c.closing {
|
|
c.mu.Unlock()
|
|
if written > 0 {
|
|
return written, net.ErrClosed
|
|
}
|
|
return 0, net.ErrClosed
|
|
}
|
|
room := c.maxBuffered - len(c.buf)
|
|
if room < 1 {
|
|
c.mu.Unlock()
|
|
continue
|
|
}
|
|
take := len(p) - written
|
|
if take > room {
|
|
take = room
|
|
}
|
|
c.buf = append(c.buf, p[written:written+take]...)
|
|
written += take
|
|
c.cond.Signal()
|
|
c.mu.Unlock()
|
|
}
|
|
return written, nil
|
|
}
|
|
|
|
func (c *sshCarrierConn) Close() error {
|
|
c.closeOnce.Do(func() {
|
|
c.mu.Lock()
|
|
c.closing = true
|
|
c.cond.Broadcast()
|
|
c.mu.Unlock()
|
|
<-c.done
|
|
})
|
|
c.mu.Lock()
|
|
err := c.writeErr
|
|
c.mu.Unlock()
|
|
return err
|
|
}
|
|
|
|
func (c *sshCarrierConn) LocalAddr() net.Addr { return c.raw.LocalAddr() }
|
|
func (c *sshCarrierConn) RemoteAddr() net.Addr { return c.raw.RemoteAddr() }
|
|
func (c *sshCarrierConn) SetDeadline(t time.Time) error { return c.raw.SetDeadline(t) }
|
|
func (c *sshCarrierConn) SetReadDeadline(t time.Time) error { return c.raw.SetReadDeadline(t) }
|
|
func (c *sshCarrierConn) SetWriteDeadline(t time.Time) error { return c.raw.SetWriteDeadline(t) }
|