V13
This commit is contained in:
@@ -0,0 +1,791 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"sort"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"dragontcp/internal/protocol"
|
||||
"dragontcp/internal/wire"
|
||||
)
|
||||
|
||||
type chunkClientOptions struct {
|
||||
startSize int
|
||||
minSize int
|
||||
maxSize int
|
||||
adaptive bool
|
||||
adaptSuccesses int
|
||||
adaptLog bool
|
||||
pollers int
|
||||
reconnectEvery int
|
||||
pollDelay time.Duration
|
||||
txnTimeout time.Duration
|
||||
tcpBuffer int
|
||||
}
|
||||
|
||||
type adaptiveSizer struct {
|
||||
mu sync.Mutex
|
||||
name string
|
||||
current int
|
||||
min int
|
||||
max int
|
||||
adaptive bool
|
||||
adaptSuccesses int
|
||||
successes int
|
||||
good int
|
||||
bad int
|
||||
logChanges bool
|
||||
}
|
||||
|
||||
func newAdaptiveSizer(name string, start int, opts chunkClientOptions) *adaptiveSizer {
|
||||
if start < opts.minSize {
|
||||
start = opts.minSize
|
||||
}
|
||||
if start > opts.maxSize {
|
||||
start = opts.maxSize
|
||||
}
|
||||
return &adaptiveSizer{
|
||||
name: name,
|
||||
current: start,
|
||||
min: opts.minSize,
|
||||
max: opts.maxSize,
|
||||
adaptive: opts.adaptive,
|
||||
adaptSuccesses: func() int {
|
||||
if opts.adaptSuccesses > 0 {
|
||||
return opts.adaptSuccesses
|
||||
}
|
||||
return 64
|
||||
}(),
|
||||
logChanges: opts.adaptLog,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *adaptiveSizer) Current() int {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return s.current
|
||||
}
|
||||
|
||||
func (s *adaptiveSizer) Success(attempted int) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if !s.adaptive || attempted != s.current || s.current >= s.max {
|
||||
return
|
||||
}
|
||||
if attempted > s.good {
|
||||
s.good = attempted
|
||||
}
|
||||
s.successes++
|
||||
growAfter := s.adaptSuccesses
|
||||
if s.bad > 0 && s.bad-s.good <= 64 {
|
||||
growAfter *= 8
|
||||
}
|
||||
if s.successes < growAfter {
|
||||
return
|
||||
}
|
||||
s.successes = 0
|
||||
|
||||
old := s.current
|
||||
next := 0
|
||||
if s.bad > old+1 {
|
||||
next = old + (s.bad-old)/2
|
||||
} else {
|
||||
if s.bad > 0 {
|
||||
s.bad = 0
|
||||
}
|
||||
step := old / 4
|
||||
if step < 32 {
|
||||
step = 32
|
||||
}
|
||||
next = old + step
|
||||
}
|
||||
if next > s.max {
|
||||
next = s.max
|
||||
}
|
||||
if next <= old {
|
||||
return
|
||||
}
|
||||
s.current = next
|
||||
if s.logChanges {
|
||||
fmt.Printf("adaptive %s chunk: %d -> %d after stable success\n", s.name, old, next)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *adaptiveSizer) Failure(attempted int) (int, int) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
old := s.current
|
||||
if !s.adaptive || attempted != s.current {
|
||||
return old, old
|
||||
}
|
||||
s.successes = 0
|
||||
if s.bad == 0 || attempted < s.bad {
|
||||
s.bad = attempted
|
||||
}
|
||||
next := attempted / 2
|
||||
if s.good > 0 && s.good < attempted {
|
||||
next = s.good
|
||||
} else {
|
||||
s.good = 0
|
||||
}
|
||||
if next < s.min {
|
||||
next = s.min
|
||||
}
|
||||
if next >= attempted && attempted > s.min {
|
||||
next = attempted - 1
|
||||
}
|
||||
if next < s.min {
|
||||
next = s.min
|
||||
}
|
||||
s.current = next
|
||||
if s.logChanges && old != next {
|
||||
fmt.Printf("adaptive %s chunk: %d -> %d after transport failure\n", s.name, old, next)
|
||||
}
|
||||
return old, next
|
||||
}
|
||||
|
||||
type physicalConn struct {
|
||||
conn net.Conn
|
||||
requests int
|
||||
}
|
||||
|
||||
type requestLane struct {
|
||||
mu sync.Mutex
|
||||
serverAddr string
|
||||
tcpBuffer int
|
||||
reconnectEvery int
|
||||
timeout time.Duration
|
||||
pc *physicalConn
|
||||
closed bool
|
||||
}
|
||||
|
||||
func newRequestLane(serverAddr string, tcpBuffer, reconnectEvery int, timeout time.Duration) *requestLane {
|
||||
return &requestLane{
|
||||
serverAddr: serverAddr,
|
||||
tcpBuffer: tcpBuffer,
|
||||
reconnectEvery: reconnectEvery,
|
||||
timeout: timeout,
|
||||
}
|
||||
}
|
||||
|
||||
func (l *requestLane) discardLocked() {
|
||||
if l.pc != nil {
|
||||
_ = l.pc.conn.Close()
|
||||
l.pc = nil
|
||||
}
|
||||
}
|
||||
|
||||
func (l *requestLane) closeAfterLocked() {
|
||||
if l.pc != nil && l.reconnectEvery > 0 && l.pc.requests >= l.reconnectEvery {
|
||||
l.discardLocked()
|
||||
}
|
||||
}
|
||||
|
||||
func (l *requestLane) ensureLocked() error {
|
||||
if l.closed {
|
||||
return net.ErrClosed
|
||||
}
|
||||
if l.pc != nil {
|
||||
if l.reconnectEvery <= 0 || l.pc.requests < l.reconnectEvery {
|
||||
return nil
|
||||
}
|
||||
l.discardLocked()
|
||||
}
|
||||
d := net.Dialer{Timeout: 10 * time.Second, KeepAlive: 30 * time.Second}
|
||||
conn, err := d.Dial("tcp", l.serverAddr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
protocol.TuneTCP(conn)
|
||||
protocol.TuneTCPBuffer(conn, l.tcpBuffer)
|
||||
l.pc = &physicalConn{conn: conn}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (l *requestLane) Close() {
|
||||
l.mu.Lock()
|
||||
l.closed = true
|
||||
l.discardLocked()
|
||||
l.mu.Unlock()
|
||||
}
|
||||
|
||||
func (l *requestLane) single(mode byte, sid wire.SessionID, seq uint64, payload []byte) (byte, []byte, error) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
if err := l.ensureLocked(); err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
timeout := l.timeout
|
||||
if timeout <= 0 {
|
||||
timeout = 5 * time.Second
|
||||
}
|
||||
_ = l.pc.conn.SetDeadline(time.Now().Add(timeout))
|
||||
if err := wire.WriteRequest(l.pc.conn, mode, sid, seq, payload); err != nil {
|
||||
l.discardLocked()
|
||||
return 0, nil, err
|
||||
}
|
||||
status, body, err := wire.ReadResponse(l.pc.conn)
|
||||
if err != nil {
|
||||
l.discardLocked()
|
||||
return 0, nil, err
|
||||
}
|
||||
l.pc.requests++
|
||||
_ = l.pc.conn.SetDeadline(time.Time{})
|
||||
l.closeAfterLocked()
|
||||
if status != wire.StatusError && len(body) > 0 {
|
||||
body = wire.DecodeMaskedResponse(status, body, sid, mode, seq)
|
||||
}
|
||||
return status, body, nil
|
||||
}
|
||||
|
||||
// download sends one compact request and consumes up to count response records.
|
||||
// startOffset is also the response keystream sequence. Each DATA response advances
|
||||
// it by exactly the returned byte count.
|
||||
func (l *requestLane) download(sid wire.SessionID, startOffset, ackOffset uint64, maxChunk, count int) ([][]byte, byte, error) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
if err := l.ensureLocked(); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
timeout := l.timeout
|
||||
if timeout <= 0 {
|
||||
timeout = 5 * time.Second
|
||||
}
|
||||
_ = l.pc.conn.SetDeadline(time.Now().Add(timeout))
|
||||
|
||||
payload := make([]byte, 14)
|
||||
binary.BigEndian.PutUint64(payload[0:8], ackOffset)
|
||||
binary.BigEndian.PutUint32(payload[8:12], uint32(maxChunk))
|
||||
binary.BigEndian.PutUint16(payload[12:14], uint16(count))
|
||||
if err := wire.WriteRequest(l.pc.conn, wire.ModeDownload, sid, startOffset, payload); err != nil {
|
||||
l.discardLocked()
|
||||
return nil, 0, err
|
||||
}
|
||||
|
||||
out := make([][]byte, 0, count)
|
||||
offset := startOffset
|
||||
lastStatus := wire.StatusOK
|
||||
for i := 0; i < count; i++ {
|
||||
status, body, err := wire.ReadResponse(l.pc.conn)
|
||||
if err != nil {
|
||||
l.discardLocked()
|
||||
return out, lastStatus, err
|
||||
}
|
||||
lastStatus = status
|
||||
switch status {
|
||||
case wire.StatusData:
|
||||
body = wire.DecodeMaskedResponse(status, body, sid, wire.ModeDownload, offset)
|
||||
if len(body) == 0 {
|
||||
l.discardLocked()
|
||||
return out, status, fmt.Errorf("empty DATA response")
|
||||
}
|
||||
out = append(out, body)
|
||||
offset += uint64(len(body))
|
||||
case wire.StatusWait, wire.StatusEOF:
|
||||
i = count // stop after this response
|
||||
case wire.StatusError:
|
||||
l.discardLocked()
|
||||
return out, status, fmt.Errorf("%s", string(body))
|
||||
default:
|
||||
l.discardLocked()
|
||||
return out, status, fmt.Errorf("unknown response status %d", status)
|
||||
}
|
||||
if status == wire.StatusWait || status == wire.StatusEOF {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
l.pc.requests++
|
||||
_ = l.pc.conn.SetDeadline(time.Time{})
|
||||
l.closeAfterLocked()
|
||||
return out, lastStatus, nil
|
||||
}
|
||||
|
||||
type pathProfile struct {
|
||||
upload int
|
||||
download int
|
||||
persistent bool
|
||||
at time.Time
|
||||
}
|
||||
|
||||
var profileState struct {
|
||||
sync.Mutex
|
||||
key string
|
||||
p pathProfile
|
||||
}
|
||||
|
||||
var probeSeq atomic.Uint64
|
||||
|
||||
func randomSessionID() (wire.SessionID, error) {
|
||||
var sid wire.SessionID
|
||||
_, err := rand.Read(sid[:])
|
||||
return sid, err
|
||||
}
|
||||
|
||||
func probePattern(n int) []byte {
|
||||
out := make([]byte, n)
|
||||
for i := range out {
|
||||
out[i] = byte((i*31 + 17) & 0xff)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func makeProbePayload(kind byte, value, total int, token string) []byte {
|
||||
base := 11 + len(token)
|
||||
if total < base {
|
||||
total = base
|
||||
}
|
||||
out := make([]byte, total)
|
||||
copy(out[:4], wire.ProbeMagic[:])
|
||||
out[4] = kind
|
||||
binary.BigEndian.PutUint16(out[5:7], uint16(len(token)))
|
||||
binary.BigEndian.PutUint32(out[7:11], uint32(value))
|
||||
copy(out[11:11+len(token)], token)
|
||||
for i := base; i < len(out); i++ {
|
||||
out[i] = byte((i*31 + 17) & 0xff)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func probeOne(serverAddr, token string, opts chunkClientOptions, kind byte, candidate int) bool {
|
||||
sid, err := randomSessionID()
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
timeout := opts.txnTimeout
|
||||
if timeout <= 0 || timeout > 2500*time.Millisecond {
|
||||
timeout = 2500 * time.Millisecond
|
||||
}
|
||||
lane := newRequestLane(serverAddr, opts.tcpBuffer, 1, timeout)
|
||||
defer lane.Close()
|
||||
seq := probeSeq.Add(1)
|
||||
|
||||
total := 0
|
||||
value := candidate
|
||||
if kind == wire.ProbeUpload {
|
||||
total = candidate
|
||||
}
|
||||
payload := makeProbePayload(kind, value, total, token)
|
||||
status, body, err := lane.single(wire.ModeProbe, sid, seq, payload)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
if kind == wire.ProbeUpload {
|
||||
return status == wire.StatusOK
|
||||
}
|
||||
if kind == wire.ProbeDownload {
|
||||
if status != wire.StatusData || len(body) != candidate {
|
||||
return false
|
||||
}
|
||||
want := probePattern(candidate)
|
||||
for i := range body {
|
||||
if body[i] != want[i] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
return status == wire.StatusOK
|
||||
}
|
||||
|
||||
func probePersistent(serverAddr, token string, opts chunkClientOptions) bool {
|
||||
sid, err := randomSessionID()
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
timeout := opts.txnTimeout
|
||||
if timeout <= 0 || timeout > 2500*time.Millisecond {
|
||||
timeout = 2500 * time.Millisecond
|
||||
}
|
||||
lane := newRequestLane(serverAddr, opts.tcpBuffer, 0, timeout)
|
||||
defer lane.Close()
|
||||
for i := 0; i < 8; i++ {
|
||||
seq := probeSeq.Add(1)
|
||||
payload := makeProbePayload(wire.ProbeKeepalive, i, 32+len(token), token)
|
||||
status, _, err := lane.single(wire.ModeProbe, sid, seq, payload)
|
||||
if err != nil || status != wire.StatusOK {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func probeCandidates(minSize, maxSize int) []int {
|
||||
base := []int{32, 64, 128, 256, 512, 1024, 1200, 1280, 1320, 1350, 1360, 1380, 1400, 1450, 1600, 2048, 3205, 4096, 8192, 16384, 32768, 65536, 98304, 131072, 262144, 524288, 786432, 1048576}
|
||||
seen := map[int]bool{}
|
||||
out := make([]int, 0, len(base)+2)
|
||||
for _, n := range base {
|
||||
if n >= minSize && n <= maxSize && !seen[n] {
|
||||
out = append(out, n)
|
||||
seen[n] = true
|
||||
}
|
||||
}
|
||||
if !seen[minSize] {
|
||||
out = append(out, minSize)
|
||||
}
|
||||
if !seen[maxSize] {
|
||||
out = append(out, maxSize)
|
||||
}
|
||||
sort.Ints(out)
|
||||
return out
|
||||
}
|
||||
|
||||
func probeMaximum(serverAddr, token string, opts chunkClientOptions, kind byte) int {
|
||||
candidates := probeCandidates(opts.minSize, opts.maxSize)
|
||||
lo, hi := 0, len(candidates)-1
|
||||
best := opts.minSize
|
||||
for lo <= hi {
|
||||
mid := lo + (hi-lo)/2
|
||||
candidate := candidates[mid]
|
||||
if probeOne(serverAddr, token, opts, kind, candidate) {
|
||||
best = candidate
|
||||
lo = mid + 1
|
||||
} else {
|
||||
hi = mid - 1
|
||||
}
|
||||
}
|
||||
if best < opts.minSize {
|
||||
best = opts.minSize
|
||||
}
|
||||
return best
|
||||
}
|
||||
|
||||
func getPathProfile(serverAddr, token string, opts chunkClientOptions) pathProfile {
|
||||
key := fmt.Sprintf("%s|%s|%d|%d", serverAddr, token, opts.minSize, opts.maxSize)
|
||||
profileState.Lock()
|
||||
if profileState.key == key && time.Since(profileState.p.at) < 30*time.Minute {
|
||||
p := profileState.p
|
||||
profileState.Unlock()
|
||||
return p
|
||||
}
|
||||
profileState.Unlock()
|
||||
|
||||
fallbackUp := minInt(opts.maxSize, maxInt(opts.minSize, 32768))
|
||||
fallbackDown := minInt(opts.maxSize, maxInt(opts.minSize, 1350))
|
||||
|
||||
upCh := make(chan int, 1)
|
||||
downCh := make(chan int, 1)
|
||||
go func() { upCh <- probeMaximum(serverAddr, token, opts, wire.ProbeUpload) }()
|
||||
go func() { downCh <- probeMaximum(serverAddr, token, opts, wire.ProbeDownload) }()
|
||||
|
||||
p := pathProfile{upload: fallbackUp, download: fallbackDown, persistent: false, at: time.Now()}
|
||||
select {
|
||||
case p.upload = <-upCh:
|
||||
case <-time.After(20 * time.Second):
|
||||
}
|
||||
select {
|
||||
case p.download = <-downCh:
|
||||
case <-time.After(20 * time.Second):
|
||||
}
|
||||
p.persistent = probePersistent(serverAddr, token, opts)
|
||||
|
||||
fmt.Printf("path probe: upload=%d download=%d persistent=%t\n", p.upload, p.download, p.persistent)
|
||||
|
||||
profileState.Lock()
|
||||
profileState.key = key
|
||||
profileState.p = p
|
||||
profileState.Unlock()
|
||||
return p
|
||||
}
|
||||
|
||||
func encodeOpen(token, host string, port int) ([]byte, error) {
|
||||
if len(token) > 65535 || len(host) > 65535 {
|
||||
return nil, fmt.Errorf("token or hostname too long")
|
||||
}
|
||||
out := make([]byte, 6+len(token)+len(host))
|
||||
binary.BigEndian.PutUint16(out[0:2], uint16(len(token)))
|
||||
binary.BigEndian.PutUint16(out[2:4], uint16(len(host)))
|
||||
binary.BigEndian.PutUint16(out[4:6], uint16(port))
|
||||
copy(out[6:6+len(token)], token)
|
||||
copy(out[6+len(token):], host)
|
||||
return out, nil
|
||||
}
|
||||
|
||||
type chunkConn struct {
|
||||
sid wire.SessionID
|
||||
opts chunkClientOptions
|
||||
uploadLane *requestLane
|
||||
downloadLane *requestLane
|
||||
serverMax int
|
||||
upSizer *adaptiveSizer
|
||||
downSizer *adaptiveSizer
|
||||
|
||||
writeMu sync.Mutex
|
||||
upOffset uint64
|
||||
|
||||
readMu sync.Mutex
|
||||
readBuf []byte
|
||||
downloadOffset uint64
|
||||
consumedOffset uint64
|
||||
eof bool
|
||||
pipeline int
|
||||
maxPipeline int
|
||||
|
||||
closeOnce sync.Once
|
||||
}
|
||||
|
||||
func openChunkTunnel(serverAddr, token, targetHost string, targetPort int, opts chunkClientOptions) (net.Conn, error) {
|
||||
if opts.minSize < 32 {
|
||||
opts.minSize = 32
|
||||
}
|
||||
if opts.maxSize < opts.minSize {
|
||||
opts.maxSize = opts.minSize
|
||||
}
|
||||
if opts.maxSize > 1024*1024 {
|
||||
opts.maxSize = 1024 * 1024
|
||||
}
|
||||
if opts.txnTimeout <= 0 {
|
||||
opts.txnTimeout = 5 * time.Second
|
||||
}
|
||||
if opts.adaptSuccesses < 1 {
|
||||
opts.adaptSuccesses = 64
|
||||
}
|
||||
if opts.reconnectEvery < 0 {
|
||||
opts.reconnectEvery = 0
|
||||
}
|
||||
|
||||
profile := getPathProfile(serverAddr, token, opts)
|
||||
reconnect := opts.reconnectEvery
|
||||
// Compatibility-friendly reconnect modes:
|
||||
// 0 = persistent (CLI explicit)
|
||||
// 1 = auto: persistent when the path probe succeeds, otherwise one request/connection
|
||||
// N>=2 = force connection rotation after N logical requests
|
||||
if reconnect == 1 {
|
||||
if profile.persistent {
|
||||
reconnect = 0
|
||||
fmt.Printf("path probe: reconnect mode auto -> persistent\n")
|
||||
} else {
|
||||
fmt.Printf("path probe: reconnect mode auto -> every request\n")
|
||||
}
|
||||
}
|
||||
|
||||
sid, err := randomSessionID()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
control := newRequestLane(serverAddr, opts.tcpBuffer, reconnect, opts.txnTimeout)
|
||||
payload, err := encodeOpen(token, targetHost, targetPort)
|
||||
if err != nil {
|
||||
control.Close()
|
||||
return nil, err
|
||||
}
|
||||
status, body, err := control.single(wire.ModeOpen, sid, 0, payload)
|
||||
if err != nil {
|
||||
control.Close()
|
||||
return nil, err
|
||||
}
|
||||
if status == wire.StatusError {
|
||||
control.Close()
|
||||
return nil, fmt.Errorf("%s", string(body))
|
||||
}
|
||||
if status != wire.StatusOK || len(body) != 4 {
|
||||
control.Close()
|
||||
return nil, fmt.Errorf("bad OPEN response")
|
||||
}
|
||||
serverMax := int(binary.BigEndian.Uint32(body))
|
||||
control.Close()
|
||||
if serverMax < opts.minSize {
|
||||
return nil, fmt.Errorf("server maximum chunk %d is below client minimum %d", serverMax, opts.minSize)
|
||||
}
|
||||
if opts.maxSize > serverMax {
|
||||
opts.maxSize = serverMax
|
||||
}
|
||||
upStart := minInt(profile.upload, opts.maxSize)
|
||||
downStart := minInt(profile.download, opts.maxSize)
|
||||
if upStart < opts.minSize {
|
||||
upStart = opts.minSize
|
||||
}
|
||||
if downStart < opts.minSize {
|
||||
downStart = opts.minSize
|
||||
}
|
||||
|
||||
c := &chunkConn{
|
||||
sid: sid,
|
||||
opts: opts,
|
||||
serverMax: serverMax,
|
||||
uploadLane: newRequestLane(serverAddr, opts.tcpBuffer, reconnect, opts.txnTimeout),
|
||||
downloadLane: newRequestLane(serverAddr, opts.tcpBuffer, reconnect, opts.txnTimeout),
|
||||
// BHTTP-style safe pipeline: request up to 256 records immediately.
|
||||
// Hybrid v1 already supported 256 on the wire/server; starting at 32
|
||||
// made tiny-path downloads spend many RTTs ramping up.
|
||||
pipeline: 256,
|
||||
maxPipeline: 256,
|
||||
}
|
||||
c.upSizer = newAdaptiveSizer("upload", upStart, opts)
|
||||
c.downSizer = newAdaptiveSizer("download", downStart, opts)
|
||||
return c, nil
|
||||
}
|
||||
|
||||
func (c *chunkConn) fillReadBuffer() error {
|
||||
if c.eof {
|
||||
return io.EOF
|
||||
}
|
||||
minFailures := 0
|
||||
for len(c.readBuf) == 0 && !c.eof {
|
||||
chunk := c.downSizer.Current()
|
||||
count := c.pipeline
|
||||
if count < 1 {
|
||||
count = 1
|
||||
}
|
||||
if count > c.maxPipeline {
|
||||
count = c.maxPipeline
|
||||
}
|
||||
// Bound each batch to roughly 1 MiB of useful data.
|
||||
if maxCount := (1024 * 1024) / maxInt(chunk, 1); maxCount < count {
|
||||
count = maxInt(maxCount, 1)
|
||||
}
|
||||
|
||||
data, status, err := c.downloadLane.download(c.sid, c.downloadOffset, c.consumedOffset, chunk, count)
|
||||
for _, part := range data {
|
||||
c.readBuf = append(c.readBuf, part...)
|
||||
c.downloadOffset += uint64(len(part))
|
||||
}
|
||||
if len(data) > 0 {
|
||||
c.downSizer.Success(chunk)
|
||||
if c.pipeline < c.maxPipeline {
|
||||
c.pipeline++
|
||||
}
|
||||
minFailures = 0
|
||||
}
|
||||
if err != nil {
|
||||
if c.pipeline > 1 {
|
||||
old := c.pipeline
|
||||
c.pipeline /= 2
|
||||
if c.pipeline < 1 {
|
||||
c.pipeline = 1
|
||||
}
|
||||
if c.opts.adaptLog && old != c.pipeline {
|
||||
fmt.Printf("adaptive download pipeline: %d -> %d after transport failure\n", old, c.pipeline)
|
||||
}
|
||||
} else {
|
||||
old, next := c.downSizer.Failure(chunk)
|
||||
if old == next && next == c.opts.minSize {
|
||||
minFailures++
|
||||
if minFailures >= 8 {
|
||||
return fmt.Errorf("download failed at minimum chunk %d: %w", next, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
time.Sleep(30 * time.Millisecond)
|
||||
if len(c.readBuf) > 0 {
|
||||
return nil
|
||||
}
|
||||
continue
|
||||
}
|
||||
switch status {
|
||||
case wire.StatusEOF:
|
||||
c.eof = true
|
||||
case wire.StatusWait:
|
||||
if c.opts.pollDelay > 0 {
|
||||
time.Sleep(c.opts.pollDelay)
|
||||
} else {
|
||||
time.Sleep(5 * time.Millisecond)
|
||||
}
|
||||
}
|
||||
if len(c.readBuf) > 0 {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
if c.eof && len(c.readBuf) == 0 {
|
||||
return io.EOF
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *chunkConn) Read(p []byte) (int, error) {
|
||||
c.readMu.Lock()
|
||||
defer c.readMu.Unlock()
|
||||
if len(p) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
if len(c.readBuf) == 0 {
|
||||
if err := c.fillReadBuffer(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
}
|
||||
n := copy(p, c.readBuf)
|
||||
c.readBuf = c.readBuf[n:]
|
||||
c.consumedOffset += uint64(n)
|
||||
return n, nil
|
||||
}
|
||||
|
||||
func (c *chunkConn) Write(p []byte) (int, error) {
|
||||
c.writeMu.Lock()
|
||||
defer c.writeMu.Unlock()
|
||||
if len(p) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
total := 0
|
||||
minFailures := 0
|
||||
for len(p) > 0 {
|
||||
size := c.upSizer.Current()
|
||||
n := minInt(size, len(p))
|
||||
status, body, err := c.uploadLane.single(wire.ModeUpload, c.sid, c.upOffset, p[:n])
|
||||
if err != nil {
|
||||
old, next := c.upSizer.Failure(size)
|
||||
if old == next && next == c.opts.minSize {
|
||||
minFailures++
|
||||
if minFailures >= 8 {
|
||||
return total, fmt.Errorf("upload failed at minimum chunk %d: %w", next, err)
|
||||
}
|
||||
} else {
|
||||
minFailures = 0
|
||||
}
|
||||
time.Sleep(30 * time.Millisecond)
|
||||
continue
|
||||
}
|
||||
if status == wire.StatusError {
|
||||
return total, fmt.Errorf("%s", string(body))
|
||||
}
|
||||
if status != wire.StatusOK {
|
||||
return total, fmt.Errorf("unexpected upload status %d", status)
|
||||
}
|
||||
c.upOffset += uint64(n)
|
||||
total += n
|
||||
p = p[n:]
|
||||
c.upSizer.Success(size)
|
||||
minFailures = 0
|
||||
}
|
||||
return total, nil
|
||||
}
|
||||
|
||||
func (c *chunkConn) Close() error {
|
||||
c.closeOnce.Do(func() {
|
||||
lane := newRequestLane(c.uploadLane.serverAddr, c.opts.tcpBuffer, 1, c.opts.txnTimeout)
|
||||
_, _, _ = lane.single(wire.ModeClose, c.sid, 0, nil)
|
||||
lane.Close()
|
||||
c.uploadLane.Close()
|
||||
c.downloadLane.Close()
|
||||
})
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *chunkConn) LocalAddr() net.Addr { return dummyAddr("dragontcp-binary-local") }
|
||||
func (c *chunkConn) RemoteAddr() net.Addr { return dummyAddr("dragontcp-binary-remote") }
|
||||
func (c *chunkConn) SetDeadline(time.Time) error { return nil }
|
||||
func (c *chunkConn) SetReadDeadline(time.Time) error { return nil }
|
||||
func (c *chunkConn) SetWriteDeadline(time.Time) error { return nil }
|
||||
|
||||
type dummyAddr string
|
||||
|
||||
func (d dummyAddr) Network() string { return "dragontcp-binary" }
|
||||
func (d dummyAddr) String() string { return string(d) }
|
||||
|
||||
func minInt(a, b int) int {
|
||||
if a < b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
func maxInt(a, b int) int {
|
||||
if a > b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
package main
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestAdaptiveSizerRecoversFromMinimum(t *testing.T) {
|
||||
opts := chunkClientOptions{
|
||||
startSize: 64,
|
||||
minSize: 32,
|
||||
maxSize: 1024,
|
||||
adaptive: true,
|
||||
adaptSuccesses: 2,
|
||||
}
|
||||
s := newAdaptiveSizer("test", 64, opts)
|
||||
_, next := s.Failure(64)
|
||||
if next != 32 {
|
||||
t.Fatalf("failure should reduce 64 -> 32, got %d", next)
|
||||
}
|
||||
for i := 0; i < 16; i++ {
|
||||
s.Success(32)
|
||||
}
|
||||
if got := s.Current(); got <= 32 {
|
||||
t.Fatalf("adaptive controller remained stuck at minimum: %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReconnectZeroMeansPersistent(t *testing.T) {
|
||||
lane := newRequestLane("127.0.0.1:1", 0, 0, 0)
|
||||
if lane.reconnectEvery != 0 {
|
||||
t.Fatalf("reconnectEvery=%d, want 0", lane.reconnectEvery)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,482 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"flag"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/url"
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"dragontcp/internal/protocol"
|
||||
)
|
||||
|
||||
const maxHeader = 128 * 1024
|
||||
|
||||
var requestCounter atomic.Uint32
|
||||
|
||||
func readHTTPHeaders(conn net.Conn) ([]byte, []byte, error) {
|
||||
buf := make([]byte, 0, 8192)
|
||||
tmp := make([]byte, 8192)
|
||||
|
||||
for {
|
||||
n, err := conn.Read(tmp)
|
||||
if n > 0 {
|
||||
buf = append(buf, tmp[:n]...)
|
||||
|
||||
if len(buf) > maxHeader {
|
||||
return nil, nil, fmt.Errorf("HTTP headers too large")
|
||||
}
|
||||
|
||||
if i := bytes.Index(buf, []byte("\r\n\r\n")); i >= 0 {
|
||||
end := i + 4
|
||||
return buf[:end], buf[end:], nil
|
||||
}
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func parseHostPort(authority string, defaultPort int) (string, int, error) {
|
||||
authority = strings.TrimSpace(authority)
|
||||
|
||||
if host, portText, err := net.SplitHostPort(authority); err == nil {
|
||||
port, err := strconv.Atoi(portText)
|
||||
return host, port, err
|
||||
}
|
||||
|
||||
// Host without port.
|
||||
if strings.HasPrefix(authority, "[") && strings.HasSuffix(authority, "]") {
|
||||
return strings.Trim(authority, "[]"), defaultPort, nil
|
||||
}
|
||||
|
||||
if strings.Count(authority, ":") == 0 {
|
||||
return authority, defaultPort, nil
|
||||
}
|
||||
|
||||
// Bare IPv6.
|
||||
if ip := net.ParseIP(authority); ip != nil {
|
||||
return authority, defaultPort, nil
|
||||
}
|
||||
|
||||
return "", 0, fmt.Errorf("invalid authority: %s", authority)
|
||||
}
|
||||
|
||||
func rewritePlainHTTPRequest(header []byte) (string, int, []byte, error) {
|
||||
text := string(header)
|
||||
lines := strings.Split(text, "\r\n")
|
||||
if len(lines) == 0 {
|
||||
return "", 0, nil, fmt.Errorf("empty request")
|
||||
}
|
||||
|
||||
parts := strings.SplitN(lines[0], " ", 3)
|
||||
if len(parts) != 3 {
|
||||
return "", 0, nil, fmt.Errorf("invalid request line")
|
||||
}
|
||||
|
||||
method, target, version := parts[0], parts[1], parts[2]
|
||||
|
||||
var (
|
||||
hostHeader string
|
||||
headers []string
|
||||
)
|
||||
|
||||
for _, line := range lines[1:] {
|
||||
if line == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
k, v, ok := strings.Cut(line, ":")
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
|
||||
lk := strings.ToLower(strings.TrimSpace(k))
|
||||
|
||||
if lk == "host" {
|
||||
hostHeader = strings.TrimSpace(v)
|
||||
}
|
||||
|
||||
if lk == "connection" ||
|
||||
lk == "proxy-connection" ||
|
||||
lk == "proxy-authorization" {
|
||||
continue
|
||||
}
|
||||
|
||||
headers = append(headers, k+": "+strings.TrimSpace(v))
|
||||
}
|
||||
|
||||
u, err := url.Parse(target)
|
||||
if err != nil {
|
||||
return "", 0, nil, err
|
||||
}
|
||||
|
||||
var host string
|
||||
var port int
|
||||
path := target
|
||||
|
||||
if u.Hostname() != "" {
|
||||
if strings.ToLower(u.Scheme) != "http" {
|
||||
return "", 0, nil, fmt.Errorf("unsupported plain HTTP scheme: %s", u.Scheme)
|
||||
}
|
||||
|
||||
host = u.Hostname()
|
||||
port = 80
|
||||
|
||||
if u.Port() != "" {
|
||||
port, err = strconv.Atoi(u.Port())
|
||||
if err != nil {
|
||||
return "", 0, nil, err
|
||||
}
|
||||
}
|
||||
|
||||
path = u.EscapedPath()
|
||||
if path == "" {
|
||||
path = "/"
|
||||
}
|
||||
if u.RawQuery != "" {
|
||||
path += "?" + u.RawQuery
|
||||
}
|
||||
} else {
|
||||
if hostHeader == "" {
|
||||
return "", 0, nil, fmt.Errorf("missing Host header")
|
||||
}
|
||||
|
||||
host, port, err = parseHostPort(hostHeader, 80)
|
||||
if err != nil {
|
||||
return "", 0, nil, err
|
||||
}
|
||||
if path == "" {
|
||||
path = "/"
|
||||
}
|
||||
}
|
||||
|
||||
var out strings.Builder
|
||||
fmt.Fprintf(&out, "%s %s %s\r\n", method, path, version)
|
||||
|
||||
sawHost := false
|
||||
for _, h := range headers {
|
||||
if strings.HasPrefix(strings.ToLower(h), "host:") {
|
||||
sawHost = true
|
||||
}
|
||||
out.WriteString(h)
|
||||
out.WriteString("\r\n")
|
||||
}
|
||||
|
||||
if !sawHost {
|
||||
if port == 80 {
|
||||
fmt.Fprintf(&out, "Host: %s\r\n", host)
|
||||
} else {
|
||||
fmt.Fprintf(&out, "Host: %s\r\n", net.JoinHostPort(host, strconv.Itoa(port)))
|
||||
}
|
||||
}
|
||||
|
||||
out.WriteString("Connection: close\r\n\r\n")
|
||||
|
||||
return host, port, []byte(out.String()), nil
|
||||
}
|
||||
|
||||
func openDragonTCPTunnel(serverAddr, token, targetHost string, targetPort int, transport string, tcpBuffer int) (net.Conn, error) {
|
||||
d := net.Dialer{
|
||||
Timeout: 10 * time.Second,
|
||||
KeepAlive: 30 * time.Second,
|
||||
}
|
||||
|
||||
conn, err := d.Dial("tcp", serverAddr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
protocol.TuneTCP(conn)
|
||||
protocol.TuneTCPBuffer(conn, tcpBuffer)
|
||||
_ = conn.SetDeadline(time.Now().Add(15 * time.Second))
|
||||
|
||||
// Correlation only; cryptographic randomness is unnecessary here.
|
||||
requestID := requestCounter.Add(1)
|
||||
|
||||
var command []byte
|
||||
if transport == "raw" {
|
||||
command = []byte(fmt.Sprintf("TUNNEL2 %s %s %d RAW", token, targetHost, targetPort))
|
||||
} else {
|
||||
// Legacy XOR command remains compatible with the older server.
|
||||
command = []byte(fmt.Sprintf("TUNNEL %s %s %d", token, targetHost, targetPort))
|
||||
}
|
||||
|
||||
if err := protocol.WriteRequestFrame(conn, requestID, command); err != nil {
|
||||
conn.Close()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
responseID, response, err := protocol.ReadResponseFrame(conn)
|
||||
if err != nil {
|
||||
conn.Close()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if responseID != requestID {
|
||||
conn.Close()
|
||||
return nil, fmt.Errorf("request ID mismatch")
|
||||
}
|
||||
|
||||
if string(response) != "CONNECTED" {
|
||||
conn.Close()
|
||||
return nil, fmt.Errorf("%s", response)
|
||||
}
|
||||
|
||||
_ = conn.SetDeadline(time.Time{})
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
func writeHTTPError(conn net.Conn, code int, reason, detail string) {
|
||||
if detail == "" {
|
||||
detail = reason
|
||||
}
|
||||
|
||||
body := []byte(detail)
|
||||
|
||||
fmt.Fprintf(
|
||||
conn,
|
||||
"HTTP/1.1 %d %s\r\nContent-Type: text/plain; charset=utf-8\r\nContent-Length: %d\r\nConnection: close\r\n\r\n",
|
||||
code,
|
||||
reason,
|
||||
len(body),
|
||||
)
|
||||
_, _ = conn.Write(body)
|
||||
}
|
||||
|
||||
func handleLocal(conn net.Conn, serverAddr, token, transport string, tcpBuffer int, chunkOpts chunkClientOptions, slots chan struct{}) {
|
||||
defer func() {
|
||||
<-slots
|
||||
_ = conn.Close()
|
||||
}()
|
||||
|
||||
protocol.TuneTCP(conn)
|
||||
protocol.TuneTCPBuffer(conn, tcpBuffer)
|
||||
_ = conn.SetDeadline(time.Now().Add(15 * time.Second))
|
||||
|
||||
header, extra, err := readHTTPHeaders(conn)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
firstLine := strings.SplitN(string(header), "\r\n", 2)[0]
|
||||
parts := strings.SplitN(firstLine, " ", 3)
|
||||
|
||||
if len(parts) != 3 {
|
||||
writeHTTPError(conn, 400, "Bad Request", "invalid HTTP request line")
|
||||
return
|
||||
}
|
||||
|
||||
method, target := parts[0], parts[1]
|
||||
|
||||
if strings.EqualFold(method, "CONNECT") {
|
||||
host, port, err := parseHostPort(target, 443)
|
||||
if err != nil {
|
||||
writeHTTPError(conn, 400, "Bad Request", err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
var remote net.Conn
|
||||
if transport == "chunk" {
|
||||
remote, err = openChunkTunnel(serverAddr, token, host, port, chunkOpts)
|
||||
} else {
|
||||
remote, err = openDragonTCPTunnel(serverAddr, token, host, port, transport, tcpBuffer)
|
||||
}
|
||||
if err != nil {
|
||||
writeHTTPError(conn, 502, "Bad Gateway", err.Error())
|
||||
return
|
||||
}
|
||||
defer remote.Close()
|
||||
|
||||
_, _ = conn.Write([]byte(
|
||||
"HTTP/1.1 200 Connection Established\r\n" +
|
||||
"Proxy-Agent: dragontcp-proxy/2.0\r\n\r\n",
|
||||
))
|
||||
|
||||
if len(extra) > 0 {
|
||||
if transport == "xor" {
|
||||
protocol.XorInPlace(extra)
|
||||
}
|
||||
if _, err := remote.Write(extra); err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
_ = conn.SetDeadline(time.Time{})
|
||||
if transport == "xor" {
|
||||
protocol.RelayXOR(conn, remote)
|
||||
} else {
|
||||
// raw and chunk connections expose a normal plaintext net.Conn.
|
||||
protocol.RelayRaw(conn, remote)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
host, port, rewritten, err := rewritePlainHTTPRequest(header)
|
||||
if err != nil {
|
||||
writeHTTPError(conn, 400, "Bad Request", err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
var remote net.Conn
|
||||
if transport == "chunk" {
|
||||
remote, err = openChunkTunnel(serverAddr, token, host, port, chunkOpts)
|
||||
} else {
|
||||
remote, err = openDragonTCPTunnel(serverAddr, token, host, port, transport, tcpBuffer)
|
||||
}
|
||||
if err != nil {
|
||||
writeHTTPError(conn, 502, "Bad Gateway", err.Error())
|
||||
return
|
||||
}
|
||||
defer remote.Close()
|
||||
|
||||
initial := make([]byte, 0, len(rewritten)+len(extra))
|
||||
initial = append(initial, rewritten...)
|
||||
initial = append(initial, extra...)
|
||||
if transport == "xor" {
|
||||
protocol.XorInPlace(initial)
|
||||
}
|
||||
|
||||
if _, err := remote.Write(initial); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
_ = conn.SetDeadline(time.Time{})
|
||||
if transport == "xor" {
|
||||
protocol.RelayXOR(conn, remote)
|
||||
} else {
|
||||
protocol.RelayRaw(conn, remote)
|
||||
}
|
||||
}
|
||||
|
||||
func main() {
|
||||
var (
|
||||
listenHost = flag.String("listen-host", "127.0.0.1", "local proxy listen host")
|
||||
listenPort = flag.Int("listen-port", 8080, "local proxy listen port")
|
||||
serverHost = flag.String("server-host", "", "remote DragonTCP server host")
|
||||
serverPort = flag.Int("server-port", 53, "remote DragonTCP server port")
|
||||
token = flag.String("token", "", "optional shared token")
|
||||
maxConnections = flag.Int("max-connections", 20000, "max simultaneous proxy connections")
|
||||
transport = flag.String("transport", "chunk", "transport: chunk (DragonTCP binary adaptive transport)")
|
||||
tcpBuffer = flag.Int("tcp-buffer", 0, "optional TCP read/write buffer bytes; 0 keeps OS autotuning")
|
||||
chunkStart = flag.Int("chunk-start", 1048576, "initial adaptive chunk payload bytes")
|
||||
chunkMin = flag.Int("chunk-min", 32, "minimum adaptive chunk payload bytes")
|
||||
chunkMax = flag.Int("chunk-max", 1048576, "maximum adaptive chunk payload bytes (up to 1 MiB)")
|
||||
chunkAdaptive = flag.Bool("chunk-adaptive", true, "automatically shrink on failures and grow after stable success")
|
||||
chunkSuccesses = flag.Int("chunk-grow-after", 16, "successful data records required before increasing chunk size")
|
||||
chunkAdaptLog = flag.Bool("chunk-adapt-log", true, "print adaptive chunk size changes")
|
||||
chunkSizeLegacy = flag.Int("chunk-size", 0, "legacy fixed chunk size; nonzero disables adaptation")
|
||||
chunkPollers = flag.Int("chunk-pollers", 1, "reserved compatibility setting; binary transport uses one download worker")
|
||||
chunkReconnect = flag.Int("chunk-reconnect-every", 0, "force reconnect after N logical requests; 0 = persistent/automatic")
|
||||
chunkPollDelay = flag.Duration("chunk-poll-delay", 2*time.Millisecond, "delay after an empty chunk poll")
|
||||
chunkTimeout = flag.Duration("chunk-timeout", 2*time.Second, "per-record transaction timeout before adaptive shrink")
|
||||
)
|
||||
flag.Parse()
|
||||
|
||||
if *serverHost == "" {
|
||||
fmt.Fprintln(os.Stderr, "--server-host is required")
|
||||
os.Exit(2)
|
||||
}
|
||||
|
||||
*transport = strings.ToLower(*transport)
|
||||
if *transport != "chunk" {
|
||||
fmt.Fprintln(os.Stderr, "DragonTCP requires --transport chunk (binary adaptive TCP/53 transport)")
|
||||
os.Exit(2)
|
||||
}
|
||||
if *chunkSizeLegacy != 0 {
|
||||
if *chunkSizeLegacy < 32 || *chunkSizeLegacy > protocol.MaxChunkPayload {
|
||||
fmt.Fprintf(os.Stderr, "--chunk-size must be between 32 and %d\n", protocol.MaxChunkPayload)
|
||||
os.Exit(2)
|
||||
}
|
||||
*chunkStart = *chunkSizeLegacy
|
||||
*chunkMin = *chunkSizeLegacy
|
||||
*chunkMax = *chunkSizeLegacy
|
||||
*chunkAdaptive = false
|
||||
}
|
||||
if *chunkMin < 32 || *chunkMax > protocol.MaxChunkPayload || *chunkMin > *chunkStart || *chunkStart > *chunkMax {
|
||||
fmt.Fprintf(os.Stderr, "require 32 <= --chunk-min <= --chunk-start <= --chunk-max <= %d\n", protocol.MaxChunkPayload)
|
||||
os.Exit(2)
|
||||
}
|
||||
if *chunkSuccesses < 1 {
|
||||
fmt.Fprintln(os.Stderr, "--chunk-grow-after must be at least 1")
|
||||
os.Exit(2)
|
||||
}
|
||||
if *chunkPollers < 1 || *chunkPollers > 128 {
|
||||
fmt.Fprintln(os.Stderr, "--chunk-pollers must be between 1 and 128")
|
||||
os.Exit(2)
|
||||
}
|
||||
if *chunkReconnect < 0 {
|
||||
fmt.Fprintln(os.Stderr, "--chunk-reconnect-every must be 0 or greater")
|
||||
os.Exit(2)
|
||||
}
|
||||
chunkOpts := chunkClientOptions{
|
||||
startSize: *chunkStart,
|
||||
minSize: *chunkMin,
|
||||
maxSize: *chunkMax,
|
||||
adaptive: *chunkAdaptive,
|
||||
adaptSuccesses: *chunkSuccesses,
|
||||
adaptLog: *chunkAdaptLog,
|
||||
pollers: *chunkPollers,
|
||||
reconnectEvery: *chunkReconnect,
|
||||
pollDelay: *chunkPollDelay,
|
||||
txnTimeout: *chunkTimeout,
|
||||
tcpBuffer: *tcpBuffer,
|
||||
}
|
||||
|
||||
listenAddr := net.JoinHostPort(*listenHost, strconv.Itoa(*listenPort))
|
||||
serverAddr := net.JoinHostPort(*serverHost, strconv.Itoa(*serverPort))
|
||||
|
||||
ln, err := net.Listen("tcp", listenAddr)
|
||||
if err != nil {
|
||||
fmt.Fprintln(os.Stderr, err)
|
||||
os.Exit(1)
|
||||
}
|
||||
defer ln.Close()
|
||||
|
||||
fmt.Printf("local Go HTTP proxy listening on %s\n", listenAddr)
|
||||
fmt.Printf("remote DragonTCP endpoint=%s\n", serverAddr)
|
||||
fmt.Printf("max_connections=%d transport=%s tcp_buffer=%d\n", *maxConnections, *transport, *tcpBuffer)
|
||||
if *transport == "chunk" {
|
||||
fmt.Printf(
|
||||
"adaptive_chunk=%v start=%d min=%d max=%d grow_after=%d pollers=%d reconnect_every=%d timeout=%s\n",
|
||||
*chunkAdaptive,
|
||||
*chunkStart,
|
||||
*chunkMin,
|
||||
*chunkMax,
|
||||
*chunkSuccesses,
|
||||
*chunkPollers,
|
||||
*chunkReconnect,
|
||||
chunkTimeout.String(),
|
||||
)
|
||||
}
|
||||
|
||||
slots := make(chan struct{}, *maxConnections)
|
||||
|
||||
for {
|
||||
conn, err := ln.Accept()
|
||||
if err != nil {
|
||||
fmt.Fprintln(os.Stderr, "accept:", err)
|
||||
continue
|
||||
}
|
||||
|
||||
select {
|
||||
case slots <- struct{}{}:
|
||||
go handleLocal(conn, serverAddr, *token, *transport, *tcpBuffer, chunkOpts, slots)
|
||||
default:
|
||||
writeHTTPError(
|
||||
conn,
|
||||
503,
|
||||
"Service Unavailable",
|
||||
"proxy connection limit reached",
|
||||
)
|
||||
_ = conn.Close()
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,507 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"net"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"dragontcp/internal/wire"
|
||||
)
|
||||
|
||||
type streamSession struct {
|
||||
sid wire.SessionID
|
||||
target net.Conn
|
||||
targetName string
|
||||
maxChunk int
|
||||
maxBuffer int
|
||||
debug *serverDebug
|
||||
|
||||
mu sync.Mutex
|
||||
notify chan struct{}
|
||||
buf []byte
|
||||
base uint64
|
||||
eof bool
|
||||
closed bool
|
||||
lastSeen time.Time
|
||||
|
||||
upMu sync.Mutex
|
||||
expectedUp uint64
|
||||
}
|
||||
|
||||
func newStreamSession(sid wire.SessionID, target net.Conn, targetName string, maxChunk, maxBuffer int, debug *serverDebug) *streamSession {
|
||||
s := &streamSession{
|
||||
sid: sid,
|
||||
target: target,
|
||||
targetName: targetName,
|
||||
maxChunk: maxChunk,
|
||||
maxBuffer: maxBuffer,
|
||||
debug: debug,
|
||||
notify: make(chan struct{}),
|
||||
lastSeen: time.Now(),
|
||||
}
|
||||
go s.readTarget()
|
||||
return s
|
||||
}
|
||||
|
||||
func (s *streamSession) signalLocked() {
|
||||
close(s.notify)
|
||||
s.notify = make(chan struct{})
|
||||
}
|
||||
|
||||
func (s *streamSession) touchLocked() { s.lastSeen = time.Now() }
|
||||
|
||||
func (s *streamSession) readTarget() {
|
||||
tmp := make([]byte, 64*1024)
|
||||
for {
|
||||
n, err := s.target.Read(tmp)
|
||||
if n > 0 {
|
||||
data := append([]byte(nil), tmp[:n]...)
|
||||
for len(data) > 0 {
|
||||
s.mu.Lock()
|
||||
for !s.closed && len(s.buf) >= s.maxBuffer {
|
||||
ch := s.notify
|
||||
s.mu.Unlock()
|
||||
<-ch
|
||||
s.mu.Lock()
|
||||
}
|
||||
if s.closed {
|
||||
s.mu.Unlock()
|
||||
return
|
||||
}
|
||||
room := s.maxBuffer - len(s.buf)
|
||||
take := len(data)
|
||||
if take > room {
|
||||
take = room
|
||||
}
|
||||
s.buf = append(s.buf, data[:take]...)
|
||||
data = data[take:]
|
||||
s.touchLocked()
|
||||
s.signalLocked()
|
||||
s.mu.Unlock()
|
||||
if s.debug != nil && s.debug.enabled {
|
||||
s.debug.bytesDown.Add(uint64(take))
|
||||
}
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
s.mu.Lock()
|
||||
if !s.closed {
|
||||
s.eof = true
|
||||
s.touchLocked()
|
||||
s.signalLocked()
|
||||
}
|
||||
s.mu.Unlock()
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *streamSession) ackLocked(offset uint64) {
|
||||
if offset <= s.base {
|
||||
return
|
||||
}
|
||||
end := s.base + uint64(len(s.buf))
|
||||
if offset > end {
|
||||
offset = end
|
||||
}
|
||||
drop := int(offset - s.base)
|
||||
if drop <= 0 {
|
||||
return
|
||||
}
|
||||
s.buf = s.buf[drop:]
|
||||
s.base = offset
|
||||
if len(s.buf) == 0 {
|
||||
s.buf = nil
|
||||
} else if cap(s.buf) > 4*len(s.buf) && cap(s.buf) > 1024*1024 {
|
||||
compact := append([]byte(nil), s.buf...)
|
||||
s.buf = compact
|
||||
}
|
||||
s.signalLocked()
|
||||
}
|
||||
|
||||
func (s *streamSession) ack(offset uint64) {
|
||||
s.mu.Lock()
|
||||
s.ackLocked(offset)
|
||||
s.touchLocked()
|
||||
s.mu.Unlock()
|
||||
}
|
||||
|
||||
func (s *streamSession) readAt(offset uint64, limit int, wait time.Duration) ([]byte, byte, error) {
|
||||
if limit < 1 || limit > s.maxChunk {
|
||||
return nil, wire.StatusError, fmt.Errorf("invalid download limit %d", limit)
|
||||
}
|
||||
deadline := time.Now().Add(wait)
|
||||
firstDataAt := time.Time{}
|
||||
|
||||
for {
|
||||
s.mu.Lock()
|
||||
s.touchLocked()
|
||||
if offset < s.base {
|
||||
s.mu.Unlock()
|
||||
return nil, wire.StatusError, fmt.Errorf("download offset %d was already acknowledged (base=%d)", offset, s.base)
|
||||
}
|
||||
rel64 := offset - s.base
|
||||
if rel64 <= uint64(len(s.buf)) {
|
||||
rel := int(rel64)
|
||||
available := len(s.buf) - rel
|
||||
if available > 0 {
|
||||
if firstDataAt.IsZero() {
|
||||
firstDataAt = time.Now()
|
||||
}
|
||||
// Coalesce tiny target reads briefly. This prevents a 1-2 byte
|
||||
// producer read from becoming a permanent tiny tunnel record.
|
||||
if available < limit && !s.eof && wait > 0 && time.Since(firstDataAt) < 2*time.Millisecond {
|
||||
ch := s.notify
|
||||
s.mu.Unlock()
|
||||
select {
|
||||
case <-ch:
|
||||
case <-time.After(2 * time.Millisecond):
|
||||
}
|
||||
continue
|
||||
}
|
||||
n := available
|
||||
if n > limit {
|
||||
n = limit
|
||||
}
|
||||
out := append([]byte(nil), s.buf[rel:rel+n]...)
|
||||
s.mu.Unlock()
|
||||
return out, wire.StatusData, nil
|
||||
}
|
||||
if s.eof || s.closed {
|
||||
s.mu.Unlock()
|
||||
return nil, wire.StatusEOF, nil
|
||||
}
|
||||
} else {
|
||||
s.mu.Unlock()
|
||||
return nil, wire.StatusError, fmt.Errorf("download offset %d is beyond buffered stream end %d", offset, s.base+uint64(len(s.buf)))
|
||||
}
|
||||
|
||||
if wait <= 0 || time.Now().After(deadline) {
|
||||
s.mu.Unlock()
|
||||
return nil, wire.StatusWait, nil
|
||||
}
|
||||
ch := s.notify
|
||||
remaining := time.Until(deadline)
|
||||
s.mu.Unlock()
|
||||
select {
|
||||
case <-ch:
|
||||
case <-time.After(remaining):
|
||||
return nil, wire.StatusWait, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *streamSession) upload(offset uint64, data []byte) error {
|
||||
if len(data) == 0 || len(data) > s.maxChunk {
|
||||
return fmt.Errorf("invalid upload size %d", len(data))
|
||||
}
|
||||
s.upMu.Lock()
|
||||
defer s.upMu.Unlock()
|
||||
|
||||
if offset < s.expectedUp {
|
||||
// Idempotent retry after a lost ACK.
|
||||
if offset+uint64(len(data)) <= s.expectedUp {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("overlapping upload retry at %d", offset)
|
||||
}
|
||||
if offset != s.expectedUp {
|
||||
return fmt.Errorf("upload gap: got %d expected %d", offset, s.expectedUp)
|
||||
}
|
||||
if _, err := s.target.Write(data); err != nil {
|
||||
return err
|
||||
}
|
||||
s.expectedUp += uint64(len(data))
|
||||
s.mu.Lock()
|
||||
s.touchLocked()
|
||||
s.mu.Unlock()
|
||||
if s.debug != nil && s.debug.enabled {
|
||||
s.debug.bytesUp.Add(uint64(len(data)))
|
||||
s.debug.pushRecords.Add(1)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *streamSession) close() {
|
||||
s.mu.Lock()
|
||||
if s.closed {
|
||||
s.mu.Unlock()
|
||||
return
|
||||
}
|
||||
s.closed = true
|
||||
s.signalLocked()
|
||||
s.mu.Unlock()
|
||||
_ = s.target.Close()
|
||||
}
|
||||
|
||||
type streamManager struct {
|
||||
mu sync.RWMutex
|
||||
sessions map[string]*streamSession
|
||||
timeout time.Duration
|
||||
debug *serverDebug
|
||||
}
|
||||
|
||||
func sidKey(sid wire.SessionID) string { return string(sid[:]) }
|
||||
|
||||
func newStreamManager(timeout time.Duration, debug *serverDebug) *streamManager {
|
||||
m := &streamManager{sessions: make(map[string]*streamSession), timeout: timeout, debug: debug}
|
||||
go m.cleanupLoop()
|
||||
return m
|
||||
}
|
||||
|
||||
func (m *streamManager) get(sid wire.SessionID) *streamSession {
|
||||
m.mu.RLock()
|
||||
s := m.sessions[sidKey(sid)]
|
||||
m.mu.RUnlock()
|
||||
return s
|
||||
}
|
||||
|
||||
func (m *streamManager) addOrGet(sid wire.SessionID, s *streamSession) (*streamSession, bool) {
|
||||
key := sidKey(sid)
|
||||
m.mu.Lock()
|
||||
if old := m.sessions[key]; old != nil {
|
||||
m.mu.Unlock()
|
||||
s.close()
|
||||
return old, false
|
||||
}
|
||||
m.sessions[key] = s
|
||||
m.mu.Unlock()
|
||||
return s, true
|
||||
}
|
||||
|
||||
func (m *streamManager) remove(sid wire.SessionID) {
|
||||
key := sidKey(sid)
|
||||
m.mu.Lock()
|
||||
s := m.sessions[key]
|
||||
delete(m.sessions, key)
|
||||
m.mu.Unlock()
|
||||
if s != nil {
|
||||
s.close()
|
||||
if m.debug != nil && m.debug.enabled {
|
||||
m.debug.sessionsClosed.Add(1)
|
||||
m.debug.activeSessions.Add(-1)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (m *streamManager) count() int { m.mu.RLock(); n := len(m.sessions); m.mu.RUnlock(); return n }
|
||||
|
||||
func (m *streamManager) cleanupLoop() {
|
||||
ticker := time.NewTicker(30 * time.Second)
|
||||
defer ticker.Stop()
|
||||
for range ticker.C {
|
||||
cutoff := time.Now().Add(-m.timeout)
|
||||
var stale []wire.SessionID
|
||||
m.mu.RLock()
|
||||
for _, s := range m.sessions {
|
||||
s.mu.Lock()
|
||||
last := s.lastSeen
|
||||
closed := s.closed
|
||||
sid := s.sid
|
||||
s.mu.Unlock()
|
||||
if closed || last.Before(cutoff) {
|
||||
stale = append(stale, sid)
|
||||
}
|
||||
}
|
||||
m.mu.RUnlock()
|
||||
for _, sid := range stale {
|
||||
m.remove(sid)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func parseProbe(payload []byte) (kind byte, value int, token string, err error) {
|
||||
if len(payload) < 11 || !bytes.Equal(payload[:4], wire.ProbeMagic[:]) {
|
||||
return 0, 0, "", fmt.Errorf("bad probe payload")
|
||||
}
|
||||
kind = payload[4]
|
||||
tl := int(binary.BigEndian.Uint16(payload[5:7]))
|
||||
value = int(binary.BigEndian.Uint32(payload[7:11]))
|
||||
if 11+tl > len(payload) {
|
||||
return 0, 0, "", fmt.Errorf("bad probe token length")
|
||||
}
|
||||
token = string(payload[11 : 11+tl])
|
||||
return
|
||||
}
|
||||
|
||||
func probePattern(n int) []byte {
|
||||
out := make([]byte, n)
|
||||
for i := range out {
|
||||
out[i] = byte((i*31 + 17) & 0xff)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func parseOpen(payload []byte) (token, host string, port int, err error) {
|
||||
if len(payload) < 6 {
|
||||
return "", "", 0, fmt.Errorf("bad OPEN payload")
|
||||
}
|
||||
tl := int(binary.BigEndian.Uint16(payload[0:2]))
|
||||
hl := int(binary.BigEndian.Uint16(payload[2:4]))
|
||||
port = int(binary.BigEndian.Uint16(payload[4:6]))
|
||||
if port < 1 || 6+tl+hl != len(payload) {
|
||||
return "", "", 0, fmt.Errorf("bad OPEN lengths")
|
||||
}
|
||||
token = string(payload[6 : 6+tl])
|
||||
host = string(payload[6+tl:])
|
||||
if host == "" {
|
||||
return "", "", 0, fmt.Errorf("empty target host")
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
func processWireRequest(conn net.Conn, req wire.Request, token string, allowPrivate bool, cache *dnsCache, tcpBuffer int, manager *streamManager, maxChunk, maxBuffer int, pollWait time.Duration, debug *serverDebug) error {
|
||||
switch req.Mode {
|
||||
case wire.ModeProbe:
|
||||
kind, value, supplied, err := parseProbe(req.Payload)
|
||||
if err != nil {
|
||||
return wire.WriteResponse(conn, wire.StatusError, []byte(err.Error()))
|
||||
}
|
||||
if !tokenEqual(supplied, token) {
|
||||
return wire.WriteResponse(conn, wire.StatusError, []byte("authentication failed"))
|
||||
}
|
||||
switch kind {
|
||||
case wire.ProbeUpload:
|
||||
if len(req.Payload) > maxChunk {
|
||||
return wire.WriteResponse(conn, wire.StatusError, []byte("probe too large"))
|
||||
}
|
||||
return wire.WriteResponse(conn, wire.StatusOK, nil)
|
||||
case wire.ProbeDownload:
|
||||
if value < 1 || value > maxChunk {
|
||||
return wire.WriteResponse(conn, wire.StatusError, []byte("probe too large"))
|
||||
}
|
||||
return wire.WriteMaskedResponse(conn, wire.StatusData, probePattern(value), req.Session, wire.ModeProbe, req.Seq)
|
||||
case wire.ProbeKeepalive:
|
||||
return wire.WriteResponse(conn, wire.StatusOK, nil)
|
||||
case wire.ProbeBatch:
|
||||
count := value
|
||||
if count < 1 {
|
||||
count = 1
|
||||
}
|
||||
if count > 16 {
|
||||
count = 16
|
||||
}
|
||||
for i := 0; i < count; i++ {
|
||||
data := probePattern(32)
|
||||
if err := wire.WriteMaskedResponse(conn, wire.StatusData, data, req.Session, wire.ModeProbe, req.Seq+uint64(i)); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
default:
|
||||
return wire.WriteResponse(conn, wire.StatusError, []byte("unknown probe kind"))
|
||||
}
|
||||
|
||||
case wire.ModeOpen:
|
||||
supplied, host, port, err := parseOpen(req.Payload)
|
||||
if err != nil {
|
||||
return wire.WriteResponse(conn, wire.StatusError, []byte(err.Error()))
|
||||
}
|
||||
if !tokenEqual(supplied, token) {
|
||||
return wire.WriteResponse(conn, wire.StatusError, []byte("authentication failed"))
|
||||
}
|
||||
if old := manager.get(req.Session); old != nil {
|
||||
body := make([]byte, 4)
|
||||
binary.BigEndian.PutUint32(body, uint32(maxChunk))
|
||||
return wire.WriteMaskedResponse(conn, wire.StatusOK, body, req.Session, wire.ModeOpen, req.Seq)
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
target, err := dialTarget(ctx, host, port, allowPrivate, cache, tcpBuffer)
|
||||
cancel()
|
||||
if err != nil {
|
||||
return wire.WriteResponse(conn, wire.StatusError, []byte(err.Error()))
|
||||
}
|
||||
session := newStreamSession(req.Session, target, fmt.Sprintf("%s:%d", host, port), maxChunk, maxBuffer, debug)
|
||||
_, created := manager.addOrGet(req.Session, session)
|
||||
if created && debug != nil && debug.enabled {
|
||||
debug.sessionsOpened.Add(1)
|
||||
debug.activeSessions.Add(1)
|
||||
debug.logf("SESSION OPEN sid=%x target=%s:%d active_sessions=%d", req.Session[:4], host, port, manager.count())
|
||||
}
|
||||
body := make([]byte, 4)
|
||||
binary.BigEndian.PutUint32(body, uint32(maxChunk))
|
||||
return wire.WriteMaskedResponse(conn, wire.StatusOK, body, req.Session, wire.ModeOpen, req.Seq)
|
||||
|
||||
case wire.ModeUpload:
|
||||
s := manager.get(req.Session)
|
||||
if s == nil {
|
||||
return wire.WriteResponse(conn, wire.StatusError, []byte("unknown session"))
|
||||
}
|
||||
if len(req.Payload) > maxChunk {
|
||||
return wire.WriteResponse(conn, wire.StatusError, []byte("upload too large"))
|
||||
}
|
||||
if err := s.upload(req.Seq, req.Payload); err != nil {
|
||||
return wire.WriteResponse(conn, wire.StatusError, []byte(err.Error()))
|
||||
}
|
||||
return wire.WriteResponse(conn, wire.StatusOK, nil)
|
||||
|
||||
case wire.ModeDownload:
|
||||
s := manager.get(req.Session)
|
||||
if s == nil {
|
||||
return wire.WriteResponse(conn, wire.StatusError, []byte("unknown session"))
|
||||
}
|
||||
if len(req.Payload) != 14 {
|
||||
return wire.WriteResponse(conn, wire.StatusError, []byte("bad download request"))
|
||||
}
|
||||
ack := binary.BigEndian.Uint64(req.Payload[0:8])
|
||||
limit := int(binary.BigEndian.Uint32(req.Payload[8:12]))
|
||||
count := int(binary.BigEndian.Uint16(req.Payload[12:14]))
|
||||
if limit < 1 {
|
||||
limit = 1
|
||||
}
|
||||
if limit > maxChunk {
|
||||
limit = maxChunk
|
||||
}
|
||||
if count < 1 {
|
||||
count = 1
|
||||
}
|
||||
if count > 256 {
|
||||
count = 256
|
||||
}
|
||||
s.ack(ack)
|
||||
offset := req.Seq
|
||||
if debug != nil && debug.enabled {
|
||||
debug.pullRequests.Add(1)
|
||||
}
|
||||
for i := 0; i < count; i++ {
|
||||
wait := time.Duration(0)
|
||||
if i == 0 {
|
||||
wait = pollWait
|
||||
}
|
||||
data, status, err := s.readAt(offset, limit, wait)
|
||||
if err != nil {
|
||||
return wire.WriteResponse(conn, wire.StatusError, []byte(err.Error()))
|
||||
}
|
||||
switch status {
|
||||
case wire.StatusData:
|
||||
if debug != nil && debug.enabled {
|
||||
debug.dataRecords.Add(1)
|
||||
}
|
||||
if err := wire.WriteMaskedResponse(conn, wire.StatusData, data, req.Session, wire.ModeDownload, offset); err != nil {
|
||||
return err
|
||||
}
|
||||
offset += uint64(len(data))
|
||||
case wire.StatusWait:
|
||||
if debug != nil && debug.enabled {
|
||||
debug.waitRecords.Add(1)
|
||||
}
|
||||
return wire.WriteResponse(conn, wire.StatusWait, nil)
|
||||
case wire.StatusEOF:
|
||||
return wire.WriteResponse(conn, wire.StatusEOF, nil)
|
||||
default:
|
||||
return wire.WriteResponse(conn, wire.StatusError, []byte("invalid session read status"))
|
||||
}
|
||||
}
|
||||
return nil
|
||||
|
||||
case wire.ModeClose:
|
||||
manager.remove(req.Session)
|
||||
return wire.WriteResponse(conn, wire.StatusOK, nil)
|
||||
default:
|
||||
return wire.WriteResponse(conn, wire.StatusError, []byte("unknown mode"))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestParseOpenAllowsEmptyToken(t *testing.T) {
|
||||
host := "example.com"
|
||||
p := make([]byte, 6+len(host))
|
||||
binary.BigEndian.PutUint16(p[0:2], 0)
|
||||
binary.BigEndian.PutUint16(p[2:4], uint16(len(host)))
|
||||
binary.BigEndian.PutUint16(p[4:6], 443)
|
||||
copy(p[6:], host)
|
||||
token, gotHost, port, err := parseOpen(p)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if token != "" || gotHost != host || port != 443 {
|
||||
t.Fatalf("got token=%q host=%q port=%d", token, gotHost, port)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,83 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
)
|
||||
|
||||
type serverDebug struct {
|
||||
enabled bool
|
||||
chunks bool
|
||||
statsEvery time.Duration
|
||||
started time.Time
|
||||
|
||||
sessionsOpened atomic.Uint64
|
||||
sessionsClosed atomic.Uint64
|
||||
activeSessions atomic.Int64
|
||||
bytesUp atomic.Uint64
|
||||
bytesDown atomic.Uint64
|
||||
pushRecords atomic.Uint64
|
||||
pullRequests atomic.Uint64
|
||||
dataRecords atomic.Uint64
|
||||
waitRecords atomic.Uint64
|
||||
errors atomic.Uint64
|
||||
}
|
||||
|
||||
func newServerDebug(enabled, chunks bool, statsEvery time.Duration) *serverDebug {
|
||||
d := &serverDebug{
|
||||
enabled: enabled || chunks,
|
||||
chunks: chunks,
|
||||
statsEvery: statsEvery,
|
||||
started: time.Now(),
|
||||
}
|
||||
if d.enabled && d.statsEvery > 0 {
|
||||
go d.statsLoop()
|
||||
}
|
||||
return d
|
||||
}
|
||||
|
||||
func (d *serverDebug) logf(format string, args ...any) {
|
||||
if d == nil || !d.enabled {
|
||||
return
|
||||
}
|
||||
fmt.Fprintf(os.Stderr, "%s [DEBUG] "+format+"\n", append([]any{time.Now().Format("2006-01-02 15:04:05.000")}, args...)...)
|
||||
}
|
||||
|
||||
func (d *serverDebug) chunkf(format string, args ...any) {
|
||||
if d == nil || !d.chunks {
|
||||
return
|
||||
}
|
||||
fmt.Fprintf(os.Stderr, "%s [CHUNK] "+format+"\n", append([]any{time.Now().Format("2006-01-02 15:04:05.000")}, args...)...)
|
||||
}
|
||||
|
||||
func (d *serverDebug) errorf(format string, args ...any) {
|
||||
if d == nil || !d.enabled {
|
||||
return
|
||||
}
|
||||
d.errors.Add(1)
|
||||
fmt.Fprintf(os.Stderr, "%s [ERROR] "+format+"\n", append([]any{time.Now().Format("2006-01-02 15:04:05.000")}, args...)...)
|
||||
}
|
||||
|
||||
func (d *serverDebug) statsLoop() {
|
||||
ticker := time.NewTicker(d.statsEvery)
|
||||
defer ticker.Stop()
|
||||
for range ticker.C {
|
||||
d.logf(
|
||||
"STATS uptime=%s active_connections=%d active_sessions=%d sessions_opened=%d sessions_closed=%d bytes_up=%d bytes_down=%d push_records=%d pull_requests=%d data_records=%d waits=%d errors=%d",
|
||||
time.Since(d.started).Round(time.Second),
|
||||
atomic.LoadInt64(&active),
|
||||
d.activeSessions.Load(),
|
||||
d.sessionsOpened.Load(),
|
||||
d.sessionsClosed.Load(),
|
||||
d.bytesUp.Load(),
|
||||
d.bytesDown.Load(),
|
||||
d.pushRecords.Load(),
|
||||
d.pullRequests.Load(),
|
||||
d.dataRecords.Load(),
|
||||
d.waitRecords.Load(),
|
||||
d.errors.Load(),
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,290 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/subtle"
|
||||
"flag"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/netip"
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"dragontcp/internal/protocol"
|
||||
"dragontcp/internal/wire"
|
||||
)
|
||||
|
||||
var active int64
|
||||
|
||||
type dnsEntry struct {
|
||||
ips []netip.Addr
|
||||
expires time.Time
|
||||
}
|
||||
|
||||
type dnsCache struct {
|
||||
mu sync.RWMutex
|
||||
entries map[string]dnsEntry
|
||||
ttl time.Duration
|
||||
max int
|
||||
}
|
||||
|
||||
func newDNSCache(ttl time.Duration, max int) *dnsCache {
|
||||
return &dnsCache{
|
||||
entries: make(map[string]dnsEntry),
|
||||
ttl: ttl,
|
||||
max: max,
|
||||
}
|
||||
}
|
||||
|
||||
func (c *dnsCache) resolve(ctx context.Context, host string) ([]netip.Addr, error) {
|
||||
if ip, err := netip.ParseAddr(host); err == nil {
|
||||
return []netip.Addr{ip}, nil
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
c.mu.RLock()
|
||||
entry, ok := c.entries[host]
|
||||
c.mu.RUnlock()
|
||||
if ok && now.Before(entry.expires) {
|
||||
return entry.ips, nil
|
||||
}
|
||||
|
||||
ips, err := net.DefaultResolver.LookupNetIP(ctx, "ip", host)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
c.mu.Lock()
|
||||
if len(c.entries) >= c.max {
|
||||
// Simple bounded reset keeps the hot cache cheap and prevents growth.
|
||||
c.entries = make(map[string]dnsEntry, c.max)
|
||||
}
|
||||
c.entries[host] = dnsEntry{ips: ips, expires: now.Add(c.ttl)}
|
||||
c.mu.Unlock()
|
||||
|
||||
return ips, nil
|
||||
}
|
||||
|
||||
func tokenEqual(a, b string) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
}
|
||||
return subtle.ConstantTimeCompare([]byte(a), []byte(b)) == 1
|
||||
}
|
||||
|
||||
var blockedSpecial = []netip.Prefix{
|
||||
netip.MustParsePrefix("0.0.0.0/8"),
|
||||
netip.MustParsePrefix("100.64.0.0/10"),
|
||||
netip.MustParsePrefix("192.0.0.0/24"),
|
||||
netip.MustParsePrefix("192.0.2.0/24"),
|
||||
netip.MustParsePrefix("198.18.0.0/15"),
|
||||
netip.MustParsePrefix("198.51.100.0/24"),
|
||||
netip.MustParsePrefix("203.0.113.0/24"),
|
||||
netip.MustParsePrefix("240.0.0.0/4"),
|
||||
netip.MustParsePrefix("2001:db8::/32"),
|
||||
}
|
||||
|
||||
func addressAllowed(addr netip.Addr, allowPrivate bool) bool {
|
||||
if addr.IsUnspecified() || addr.IsMulticast() {
|
||||
return false
|
||||
}
|
||||
|
||||
if allowPrivate {
|
||||
return true
|
||||
}
|
||||
|
||||
if !addr.IsGlobalUnicast() ||
|
||||
addr.IsPrivate() ||
|
||||
addr.IsLoopback() ||
|
||||
addr.IsLinkLocalUnicast() {
|
||||
return false
|
||||
}
|
||||
|
||||
for _, prefix := range blockedSpecial {
|
||||
if prefix.Contains(addr) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
func dialTarget(ctx context.Context, host string, port int, allowPrivate bool, cache *dnsCache, tcpBuffer int) (net.Conn, error) {
|
||||
ips, err := cache.resolve(ctx, host)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var lastErr error
|
||||
var blocked []string
|
||||
|
||||
d := net.Dialer{
|
||||
Timeout: 10 * time.Second,
|
||||
KeepAlive: 30 * time.Second,
|
||||
}
|
||||
|
||||
for _, ip := range ips {
|
||||
if !addressAllowed(ip, allowPrivate) {
|
||||
blocked = append(blocked, ip.String())
|
||||
continue
|
||||
}
|
||||
|
||||
addr := net.JoinHostPort(ip.String(), strconv.Itoa(port))
|
||||
conn, err := d.DialContext(ctx, "tcp", addr)
|
||||
if err == nil {
|
||||
protocol.TuneTCP(conn)
|
||||
protocol.TuneTCPBuffer(conn, tcpBuffer)
|
||||
return conn, nil
|
||||
}
|
||||
lastErr = err
|
||||
}
|
||||
|
||||
if lastErr != nil {
|
||||
return nil, lastErr
|
||||
}
|
||||
if len(blocked) > 0 {
|
||||
return nil, fmt.Errorf("target resolves only to blocked addresses: %s", strings.Join(blocked, ","))
|
||||
}
|
||||
return nil, fmt.Errorf("no usable target address")
|
||||
}
|
||||
|
||||
func handle(
|
||||
conn net.Conn,
|
||||
token string,
|
||||
allowPrivate bool,
|
||||
cache *dnsCache,
|
||||
tcpBuffer int,
|
||||
slots chan struct{},
|
||||
manager *streamManager,
|
||||
chunkMax int,
|
||||
bufferBytes int,
|
||||
chunkPollWait time.Duration,
|
||||
debug *serverDebug,
|
||||
) {
|
||||
defer func() {
|
||||
<-slots
|
||||
atomic.AddInt64(&active, -1)
|
||||
_ = conn.Close()
|
||||
}()
|
||||
|
||||
protocol.TuneTCP(conn)
|
||||
protocol.TuneTCPBuffer(conn, tcpBuffer)
|
||||
|
||||
for {
|
||||
_ = conn.SetDeadline(time.Now().Add(30 * time.Second))
|
||||
req, err := wire.ReadRequest(conn)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
if err := processWireRequest(
|
||||
conn,
|
||||
req,
|
||||
token,
|
||||
allowPrivate,
|
||||
cache,
|
||||
tcpBuffer,
|
||||
manager,
|
||||
chunkMax,
|
||||
bufferBytes,
|
||||
chunkPollWait,
|
||||
debug,
|
||||
); err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func main() {
|
||||
var (
|
||||
host = flag.String("host", "0.0.0.0", "listen host")
|
||||
port = flag.Int("port", 53, "listen port")
|
||||
token = flag.String("token", "", "optional shared token")
|
||||
maxConnections = flag.Int("max-connections", 20000, "max simultaneous tunnels")
|
||||
allowPrivate = flag.Bool("allow-private", false, "allow private/loopback targets")
|
||||
dnsCacheTTL = flag.Duration("dns-cache-ttl", 30*time.Second, "server DNS cache TTL")
|
||||
dnsCacheSize = flag.Int("dns-cache-size", 4096, "maximum cached DNS hostnames")
|
||||
tcpBuffer = flag.Int("tcp-buffer", 0, "optional TCP read/write buffer bytes; 0 keeps OS autotuning")
|
||||
chunkMax = flag.Int("chunk-max", 1048576, "maximum adaptive chunk payload bytes (32 bytes to 1 MiB)")
|
||||
chunkBuffered = flag.Int("chunk-buffered", 256, "compatibility buffer units; 256 = about 16 MiB per active session")
|
||||
chunkPollWait = flag.Duration("chunk-poll-wait", 200*time.Millisecond, "server long-poll wait for chunk data")
|
||||
sessionTimeout = flag.Duration("chunk-session-timeout", 2*time.Minute, "idle chunk session timeout")
|
||||
debugEnabled = flag.Bool("debug", false, "log session/connect/errors and periodic statistics")
|
||||
debugChunks = flag.Bool("debug-chunks", false, "log every chunk protocol record; very verbose")
|
||||
debugStats = flag.Duration("debug-stats-interval", 5*time.Second, "periodic debug statistics interval; 0 disables")
|
||||
)
|
||||
flag.Parse()
|
||||
|
||||
if *chunkMax < 32 || *chunkMax > protocol.MaxChunkPayload {
|
||||
fmt.Fprintf(os.Stderr, "--chunk-max must be between 32 and %d\n", protocol.MaxChunkPayload)
|
||||
os.Exit(2)
|
||||
}
|
||||
if *chunkBuffered < 8 {
|
||||
fmt.Fprintln(os.Stderr, "--chunk-buffered must be at least 8")
|
||||
os.Exit(2)
|
||||
}
|
||||
|
||||
listenAddr := net.JoinHostPort(*host, strconv.Itoa(*port))
|
||||
ln, err := net.Listen("tcp", listenAddr)
|
||||
if err != nil {
|
||||
fmt.Fprintln(os.Stderr, err)
|
||||
os.Exit(1)
|
||||
}
|
||||
defer ln.Close()
|
||||
|
||||
fmt.Printf("DragonTCP Go server listening on %s\n", listenAddr)
|
||||
fmt.Printf("max_connections=%d tcp_buffer=%d\n", *maxConnections, *tcpBuffer)
|
||||
|
||||
slots := make(chan struct{}, *maxConnections)
|
||||
cache := newDNSCache(*dnsCacheTTL, *dnsCacheSize)
|
||||
debug := newServerDebug(*debugEnabled, *debugChunks, *debugStats)
|
||||
bufferBytes := *chunkBuffered * 65536
|
||||
if bufferBytes < 1024*1024 {
|
||||
bufferBytes = 1024 * 1024
|
||||
}
|
||||
if bufferBytes > 64*1024*1024 {
|
||||
bufferBytes = 64 * 1024 * 1024
|
||||
}
|
||||
manager := newStreamManager(*sessionTimeout, debug)
|
||||
fmt.Printf("binary_transport=true chunk_max=%d buffer_bytes=%d poll_wait=%s\n", *chunkMax, bufferBytes, chunkPollWait.String())
|
||||
if debug.enabled {
|
||||
fmt.Printf("debug=true debug_chunks=%t stats_interval=%s\n", debug.chunks, debug.statsEvery)
|
||||
}
|
||||
|
||||
for {
|
||||
conn, err := ln.Accept()
|
||||
if err != nil {
|
||||
fmt.Fprintln(os.Stderr, "accept:", err)
|
||||
continue
|
||||
}
|
||||
|
||||
select {
|
||||
case slots <- struct{}{}:
|
||||
atomic.AddInt64(&active, 1)
|
||||
if debug.enabled {
|
||||
debug.logf("ACCEPT peer=%v active_connections=%d", conn.RemoteAddr(), atomic.LoadInt64(&active))
|
||||
}
|
||||
go handle(
|
||||
conn,
|
||||
*token,
|
||||
*allowPrivate,
|
||||
cache,
|
||||
*tcpBuffer,
|
||||
slots,
|
||||
manager,
|
||||
*chunkMax,
|
||||
bufferBytes,
|
||||
*chunkPollWait,
|
||||
debug,
|
||||
)
|
||||
default:
|
||||
if debug.enabled {
|
||||
debug.errorf("REJECT peer=%v reason=max-connections", conn.RemoteAddr())
|
||||
}
|
||||
_ = conn.Close()
|
||||
}
|
||||
}
|
||||
}
|
||||
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,3 @@
|
||||
module dragontcp
|
||||
|
||||
go 1.22
|
||||
@@ -0,0 +1,211 @@
|
||||
package protocol
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
XORKey byte = 0xAD
|
||||
|
||||
// MaxChunkPayload is the hard application-record payload ceiling.
|
||||
// The adaptive chunk protocol may use any size from 32 bytes through 1 MiB.
|
||||
MaxChunkPayload = 1024 * 1024
|
||||
|
||||
// Framed CPUSH/DATA messages include text metadata in addition to chunk
|
||||
// bytes, so keep the frame ceiling comfortably above MaxChunkPayload.
|
||||
MaxHandshake = 2 * 1024 * 1024
|
||||
)
|
||||
|
||||
// 64 KiB balances throughput with memory use at high connection counts.
|
||||
var BufferPool = sync.Pool{
|
||||
New: func() any {
|
||||
b := make([]byte, 64*1024)
|
||||
return &b
|
||||
},
|
||||
}
|
||||
|
||||
func ReadRequestFrame(r io.Reader) (uint32, uint32, []byte, error) {
|
||||
var header [14]byte
|
||||
if _, err := io.ReadFull(r, header[:]); err != nil {
|
||||
return 0, 0, nil, err
|
||||
}
|
||||
|
||||
if header[0] != 'U' || header[1] != 'P' {
|
||||
return 0, 0, nil, errors.New("bad request magic")
|
||||
}
|
||||
|
||||
requestID := binary.BigEndian.Uint32(header[2:6])
|
||||
reserved := binary.BigEndian.Uint32(header[6:10])
|
||||
length := binary.BigEndian.Uint32(header[10:14])
|
||||
|
||||
if length > MaxHandshake {
|
||||
return 0, 0, nil, errors.New("handshake payload too large")
|
||||
}
|
||||
|
||||
payload := make([]byte, int(length))
|
||||
if _, err := io.ReadFull(r, payload); err != nil {
|
||||
return 0, 0, nil, err
|
||||
}
|
||||
XorInPlace(payload)
|
||||
|
||||
return requestID, reserved, payload, nil
|
||||
}
|
||||
|
||||
func WriteRequestFrame(w io.Writer, requestID uint32, payload []byte) error {
|
||||
if len(payload) > MaxHandshake {
|
||||
return errors.New("request frame payload too large")
|
||||
}
|
||||
packet := make([]byte, 14+len(payload))
|
||||
packet[0], packet[1] = 'U', 'P'
|
||||
binary.BigEndian.PutUint32(packet[2:6], requestID)
|
||||
binary.BigEndian.PutUint32(packet[6:10], 0)
|
||||
binary.BigEndian.PutUint32(packet[10:14], uint32(len(payload)))
|
||||
copy(packet[14:], payload)
|
||||
XorInPlace(packet[14:])
|
||||
return writeAll(w, packet)
|
||||
}
|
||||
|
||||
func ReadResponseFrame(r io.Reader) (uint32, []byte, error) {
|
||||
var header [10]byte
|
||||
if _, err := io.ReadFull(r, header[:]); err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
|
||||
if header[0] != 'O' || header[1] != 'K' {
|
||||
return 0, nil, fmt.Errorf("bad response magic: %q", header[:2])
|
||||
}
|
||||
|
||||
requestID := binary.BigEndian.Uint32(header[2:6])
|
||||
length := binary.BigEndian.Uint32(header[6:10])
|
||||
|
||||
if length > MaxHandshake {
|
||||
return 0, nil, errors.New("handshake response too large")
|
||||
}
|
||||
|
||||
payload := make([]byte, int(length))
|
||||
if _, err := io.ReadFull(r, payload); err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
XorInPlace(payload)
|
||||
|
||||
return requestID, payload, nil
|
||||
}
|
||||
|
||||
func WriteResponseFrame(w io.Writer, requestID uint32, payload []byte) error {
|
||||
if len(payload) > MaxHandshake {
|
||||
return errors.New("response frame payload too large")
|
||||
}
|
||||
packet := make([]byte, 10+len(payload))
|
||||
packet[0], packet[1] = 'O', 'K'
|
||||
binary.BigEndian.PutUint32(packet[2:6], requestID)
|
||||
binary.BigEndian.PutUint32(packet[6:10], uint32(len(payload)))
|
||||
copy(packet[10:], payload)
|
||||
XorInPlace(packet[10:])
|
||||
return writeAll(w, packet)
|
||||
}
|
||||
|
||||
func writeAll(w io.Writer, b []byte) error {
|
||||
for len(b) > 0 {
|
||||
n, err := w.Write(b)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
b = b[n:]
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func CopyXOR(dst net.Conn, src net.Conn) error {
|
||||
ptr := BufferPool.Get().(*[]byte)
|
||||
buf := *ptr
|
||||
defer BufferPool.Put(ptr)
|
||||
|
||||
for {
|
||||
n, err := src.Read(buf)
|
||||
if n > 0 {
|
||||
chunk := buf[:n]
|
||||
XorInPlace(chunk)
|
||||
|
||||
if err2 := writeAll(dst, chunk); err2 != nil {
|
||||
return err2
|
||||
}
|
||||
|
||||
// No restore pass is needed. The next Read overwrites these bytes.
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
if errors.Is(err, io.EOF) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func relayPair(a, b net.Conn, copier func(net.Conn, net.Conn) error) {
|
||||
done := make(chan struct{}, 2)
|
||||
|
||||
go func() {
|
||||
_ = copier(b, a)
|
||||
if cw, ok := b.(interface{ CloseWrite() error }); ok {
|
||||
_ = cw.CloseWrite()
|
||||
}
|
||||
done <- struct{}{}
|
||||
}()
|
||||
|
||||
go func() {
|
||||
_ = copier(a, b)
|
||||
if cw, ok := a.(interface{ CloseWrite() error }); ok {
|
||||
_ = cw.CloseWrite()
|
||||
}
|
||||
done <- struct{}{}
|
||||
}()
|
||||
|
||||
// Preserve normal TCP half-close semantics. The old implementation set a
|
||||
// 2-second deadline on both connections after the first copy direction
|
||||
// ended, which truncated slow or large responses. Wait for the remaining
|
||||
// direction to drain naturally instead.
|
||||
<-done
|
||||
<-done
|
||||
}
|
||||
|
||||
func RelayXOR(a, b net.Conn) {
|
||||
relayPair(a, b, CopyXOR)
|
||||
}
|
||||
|
||||
// RelayRaw allows Go/Linux to use the optimized TCP io.Copy path. On Linux,
|
||||
// TCP-to-TCP copies can use splice, eliminating the userspace XOR/copy loop.
|
||||
func RelayRaw(a, b net.Conn) {
|
||||
relayPair(a, b, func(dst, src net.Conn) error {
|
||||
_, err := io.Copy(dst, src)
|
||||
return err
|
||||
})
|
||||
}
|
||||
|
||||
func TuneTCP(conn net.Conn) {
|
||||
if tcp, ok := conn.(*net.TCPConn); ok {
|
||||
_ = tcp.SetNoDelay(true)
|
||||
_ = tcp.SetKeepAlive(true)
|
||||
_ = tcp.SetKeepAlivePeriod(30 * time.Second)
|
||||
}
|
||||
}
|
||||
|
||||
// TuneTCPBuffer optionally requests larger kernel socket buffers. A value <= 0
|
||||
// leaves Linux/Android autotuning untouched, which is the recommended default
|
||||
// for large connection counts. For a small number of high-BDP mobile links,
|
||||
// values such as 1048576 or 4194304 can improve throughput.
|
||||
func TuneTCPBuffer(conn net.Conn, size int) {
|
||||
if size <= 0 {
|
||||
return
|
||||
}
|
||||
if tcp, ok := conn.(*net.TCPConn); ok {
|
||||
_ = tcp.SetReadBuffer(size)
|
||||
_ = tcp.SetWriteBuffer(size)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,44 @@
|
||||
//go:build arm || 386
|
||||
|
||||
package protocol
|
||||
|
||||
import "unsafe"
|
||||
|
||||
const xorWordMask32 uint32 = 0xADADADAD
|
||||
|
||||
// XorInPlace is the 32-bit optimized path used by ARMv7/386 builds.
|
||||
// It aligns once, then processes 32 bytes per iteration with native uint32
|
||||
// operations instead of a byte-at-a-time loop.
|
||||
func XorInPlace(b []byte) {
|
||||
n := len(b)
|
||||
if n == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
i := 0
|
||||
for i < n && (uintptr(unsafe.Pointer(&b[i]))&3) != 0 {
|
||||
b[i] ^= XORKey
|
||||
i++
|
||||
}
|
||||
|
||||
for ; i+32 <= n; i += 32 {
|
||||
p := unsafe.Pointer(&b[i])
|
||||
*(*uint32)(unsafe.Add(p, 0)) ^= xorWordMask32
|
||||
*(*uint32)(unsafe.Add(p, 4)) ^= xorWordMask32
|
||||
*(*uint32)(unsafe.Add(p, 8)) ^= xorWordMask32
|
||||
*(*uint32)(unsafe.Add(p, 12)) ^= xorWordMask32
|
||||
*(*uint32)(unsafe.Add(p, 16)) ^= xorWordMask32
|
||||
*(*uint32)(unsafe.Add(p, 20)) ^= xorWordMask32
|
||||
*(*uint32)(unsafe.Add(p, 24)) ^= xorWordMask32
|
||||
*(*uint32)(unsafe.Add(p, 28)) ^= xorWordMask32
|
||||
}
|
||||
|
||||
for ; i+4 <= n; i += 4 {
|
||||
p := (*uint32)(unsafe.Pointer(&b[i]))
|
||||
*p ^= xorWordMask32
|
||||
}
|
||||
|
||||
for ; i < n; i++ {
|
||||
b[i] ^= XORKey
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
//go:build amd64 || arm64
|
||||
|
||||
package protocol
|
||||
|
||||
import "unsafe"
|
||||
|
||||
const xorWordMask uint64 = 0xADADADADADADADAD
|
||||
|
||||
// XorInPlace is optimized for 64-bit targets (amd64/arm64).
|
||||
//
|
||||
// It aligns the input once, then XORs 64 bytes per loop iteration using
|
||||
// eight native 64-bit operations. This removes the encoding/binary call
|
||||
// overhead from the hot relay path and lets the compiler generate a tight
|
||||
// load/xor/store loop.
|
||||
func XorInPlace(b []byte) {
|
||||
n := len(b)
|
||||
if n == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
i := 0
|
||||
|
||||
// Align the pointer for native uint64 accesses. This is normally already
|
||||
// aligned for pooled relay buffers, but also makes this safe for subslices.
|
||||
for i < n && (uintptr(unsafe.Pointer(&b[i]))&7) != 0 {
|
||||
b[i] ^= XORKey
|
||||
i++
|
||||
}
|
||||
|
||||
for ; i+64 <= n; i += 64 {
|
||||
p := unsafe.Pointer(&b[i])
|
||||
*(*uint64)(unsafe.Add(p, 0)) ^= xorWordMask
|
||||
*(*uint64)(unsafe.Add(p, 8)) ^= xorWordMask
|
||||
*(*uint64)(unsafe.Add(p, 16)) ^= xorWordMask
|
||||
*(*uint64)(unsafe.Add(p, 24)) ^= xorWordMask
|
||||
*(*uint64)(unsafe.Add(p, 32)) ^= xorWordMask
|
||||
*(*uint64)(unsafe.Add(p, 40)) ^= xorWordMask
|
||||
*(*uint64)(unsafe.Add(p, 48)) ^= xorWordMask
|
||||
*(*uint64)(unsafe.Add(p, 56)) ^= xorWordMask
|
||||
}
|
||||
|
||||
for ; i+8 <= n; i += 8 {
|
||||
p := (*uint64)(unsafe.Pointer(&b[i]))
|
||||
*p ^= xorWordMask
|
||||
}
|
||||
|
||||
for ; i < n; i++ {
|
||||
b[i] ^= XORKey
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
//go:build !amd64 && !arm64 && !arm && !386
|
||||
|
||||
package protocol
|
||||
|
||||
// Generic fallback for 32-bit and uncommon architectures.
|
||||
func XorInPlace(b []byte) {
|
||||
for i := range b {
|
||||
b[i] ^= XORKey
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,177 @@
|
||||
package wire
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
)
|
||||
|
||||
const (
|
||||
RequestHeaderSize = 29
|
||||
ResponseHeaderSize = 5
|
||||
MaxPayload = 2 * 1024 * 1024
|
||||
|
||||
ModeProbe byte = 0
|
||||
ModeOpen byte = 1
|
||||
ModeUpload byte = 2
|
||||
ModeDownload byte = 3
|
||||
ModeClose byte = 4
|
||||
|
||||
StatusOK byte = 0
|
||||
StatusError byte = 1
|
||||
StatusData byte = 2
|
||||
StatusWait byte = 3
|
||||
StatusEOF byte = 4
|
||||
|
||||
ProbeUpload byte = 1
|
||||
ProbeDownload byte = 2
|
||||
ProbeKeepalive byte = 3
|
||||
ProbeBatch byte = 4
|
||||
)
|
||||
|
||||
var ProbeMagic = [4]byte{'D', 'T', 'P', '2'}
|
||||
|
||||
type SessionID [16]byte
|
||||
|
||||
type Request struct {
|
||||
Mode byte
|
||||
Session SessionID
|
||||
Seq uint64
|
||||
Payload []byte
|
||||
}
|
||||
|
||||
func MaskInPlace(data []byte, sid SessionID, mode byte, seq uint64, response bool) {
|
||||
if len(data) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
var seed [30]byte
|
||||
copy(seed[:16], sid[:])
|
||||
seed[16] = mode
|
||||
binary.BigEndian.PutUint64(seed[17:25], seq)
|
||||
if response {
|
||||
seed[25] = 1
|
||||
}
|
||||
|
||||
var counter uint32
|
||||
for off := 0; off < len(data); {
|
||||
binary.BigEndian.PutUint32(seed[26:30], counter)
|
||||
block := sha256.Sum256(seed[:])
|
||||
n := len(data) - off
|
||||
if n > len(block) {
|
||||
n = len(block)
|
||||
}
|
||||
for i := 0; i < n; i++ {
|
||||
data[off+i] ^= block[i]
|
||||
}
|
||||
off += n
|
||||
counter++
|
||||
}
|
||||
}
|
||||
|
||||
func WriteRequest(w io.Writer, mode byte, sid SessionID, seq uint64, plaintext []byte) error {
|
||||
if len(plaintext) > MaxPayload {
|
||||
return fmt.Errorf("request payload too large: %d", len(plaintext))
|
||||
}
|
||||
|
||||
packet := make([]byte, RequestHeaderSize+len(plaintext))
|
||||
packet[0] = mode
|
||||
copy(packet[1:17], sid[:])
|
||||
binary.BigEndian.PutUint64(packet[17:25], seq)
|
||||
binary.BigEndian.PutUint32(packet[25:29], uint32(len(plaintext)))
|
||||
copy(packet[29:], plaintext)
|
||||
MaskInPlace(packet[29:], sid, mode, seq, false)
|
||||
return writeAll(w, packet)
|
||||
}
|
||||
|
||||
func ReadRequest(r io.Reader) (Request, error) {
|
||||
var req Request
|
||||
var header [RequestHeaderSize]byte
|
||||
if _, err := io.ReadFull(r, header[:]); err != nil {
|
||||
return req, err
|
||||
}
|
||||
|
||||
req.Mode = header[0]
|
||||
copy(req.Session[:], header[1:17])
|
||||
req.Seq = binary.BigEndian.Uint64(header[17:25])
|
||||
n := binary.BigEndian.Uint32(header[25:29])
|
||||
if n > MaxPayload {
|
||||
return req, errors.New("request payload too large")
|
||||
}
|
||||
|
||||
if n > 0 {
|
||||
req.Payload = make([]byte, int(n))
|
||||
if _, err := io.ReadFull(r, req.Payload); err != nil {
|
||||
return req, err
|
||||
}
|
||||
MaskInPlace(req.Payload, req.Session, req.Mode, req.Seq, false)
|
||||
}
|
||||
return req, nil
|
||||
}
|
||||
|
||||
func WriteResponse(w io.Writer, status byte, body []byte) error {
|
||||
if len(body) > MaxPayload {
|
||||
return fmt.Errorf("response body too large: %d", len(body))
|
||||
}
|
||||
packet := make([]byte, ResponseHeaderSize+len(body))
|
||||
packet[0] = status
|
||||
binary.BigEndian.PutUint32(packet[1:5], uint32(len(body)))
|
||||
copy(packet[5:], body)
|
||||
return writeAll(w, packet)
|
||||
}
|
||||
|
||||
func WriteMaskedResponse(w io.Writer, status byte, body []byte, sid SessionID, mode byte, seq uint64) error {
|
||||
if len(body) > MaxPayload {
|
||||
return fmt.Errorf("response body too large: %d", len(body))
|
||||
}
|
||||
packet := make([]byte, ResponseHeaderSize+len(body))
|
||||
packet[0] = status
|
||||
binary.BigEndian.PutUint32(packet[1:5], uint32(len(body)))
|
||||
copy(packet[5:], body)
|
||||
MaskInPlace(packet[5:], sid, mode, seq, true)
|
||||
return writeAll(w, packet)
|
||||
}
|
||||
|
||||
func ReadResponse(r io.Reader) (byte, []byte, error) {
|
||||
var header [ResponseHeaderSize]byte
|
||||
if _, err := io.ReadFull(r, header[:]); err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
n := binary.BigEndian.Uint32(header[1:5])
|
||||
if n > MaxPayload {
|
||||
return 0, nil, errors.New("response body too large")
|
||||
}
|
||||
var body []byte
|
||||
if n > 0 {
|
||||
body = make([]byte, int(n))
|
||||
if _, err := io.ReadFull(r, body); err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
}
|
||||
return header[0], body, nil
|
||||
}
|
||||
|
||||
func DecodeMaskedResponse(status byte, body []byte, sid SessionID, mode byte, seq uint64) []byte {
|
||||
if len(body) == 0 || status == StatusError {
|
||||
return body
|
||||
}
|
||||
out := append([]byte(nil), body...)
|
||||
MaskInPlace(out, sid, mode, seq, true)
|
||||
return out
|
||||
}
|
||||
|
||||
func writeAll(w io.Writer, b []byte) error {
|
||||
for len(b) > 0 {
|
||||
n, err := w.Write(b)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if n <= 0 {
|
||||
return io.ErrShortWrite
|
||||
}
|
||||
b = b[n:]
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
package wire
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestMaskChangesWithSequenceAndRoundTrips(t *testing.T) {
|
||||
var sid SessionID
|
||||
for i := range sid { sid[i] = byte(i+1) }
|
||||
plain := bytes.Repeat([]byte("DragonTCP"), 100)
|
||||
a := append([]byte(nil), plain...)
|
||||
b := append([]byte(nil), plain...)
|
||||
MaskInPlace(a, sid, ModeUpload, 1, false)
|
||||
MaskInPlace(b, sid, ModeUpload, 2, false)
|
||||
if bytes.Equal(a, b) { t.Fatal("different sequences produced identical wire bytes") }
|
||||
MaskInPlace(a, sid, ModeUpload, 1, false)
|
||||
if !bytes.Equal(a, plain) { t.Fatal("mask did not round-trip") }
|
||||
}
|
||||
|
||||
func BenchmarkMask1MiB(b *testing.B) {
|
||||
var sid SessionID
|
||||
data := make([]byte, 1024*1024)
|
||||
b.SetBytes(int64(len(data)))
|
||||
b.ResetTimer()
|
||||
for i:=0;i<b.N;i++ {
|
||||
MaskInPlace(data,sid,ModeUpload,uint64(i),false)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user