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 minPipeline int maxPipeline 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 minPipeline 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 } if opts.maxPipeline < 1 { opts.maxPipeline = 1 } if opts.maxPipeline > 256 { opts.maxPipeline = 256 } if opts.minPipeline < 1 { opts.minPipeline = 1 } if opts.minPipeline > opts.maxPipeline { opts.minPipeline = opts.maxPipeline } 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 // Resolved silently: this runs once per proxied flow, so it must never log. if reconnect == 1 && profile.persistent { reconnect = 0 } sid, err := randomSessionID() if err != nil { return nil, err } // OPEN rides the upload lane instead of a throwaway connection. A dedicated // control connection cost one extra dial per proxied flow, which shows up on // the server as connection churn on top of the steady-state count. uploadLane := newRequestLane(serverAddr, opts.tcpBuffer, reconnect, opts.txnTimeout) payload, err := encodeOpen(token, targetHost, targetPort) if err != nil { uploadLane.Close() return nil, err } status, body, err := uploadLane.single(wire.ModeOpen, sid, 0, payload) if err != nil { uploadLane.Close() return nil, err } if status == wire.StatusError { uploadLane.Close() return nil, fmt.Errorf("%s", string(body)) } if status != wire.StatusOK || len(body) != 4 { uploadLane.Close() return nil, fmt.Errorf("bad OPEN response") } serverMax := int(binary.BigEndian.Uint32(body)) if serverMax < opts.minSize { uploadLane.Close() 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: uploadLane, downloadLane: newRequestLane(serverAddr, opts.tcpBuffer, reconnect, opts.txnTimeout), // Start at the configured ceiling. On transport failure the batch is // halved but never below minPipeline; successful data grows it back by // one. When min == max the depth is pinned and never adapts, which is // what paths that only work at one specific batch size need. pipeline: opts.maxPipeline, minPipeline: opts.minPipeline, maxPipeline: opts.maxPipeline, } 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 < c.minPipeline { count = c.minPipeline } if count > c.maxPipeline { count = c.maxPipeline } // Bound each batch to roughly 1 MiB of useful data, but never below the // configured floor: a pinned depth is a path requirement, not a hint. if maxCount := (1024 * 1024) / maxInt(chunk, 1); maxCount < count { count = maxInt(maxCount, c.minPipeline) } 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 > c.minPipeline { old := c.pipeline c.pipeline /= 2 if c.pipeline < c.minPipeline { c.pipeline = c.minPipeline } 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() { // Reuse the upload lane rather than dialling a connection just to say // goodbye; that was a second wasted dial per flow. _, _, _ = c.uploadLane.single(wire.ModeClose, c.sid, 0, nil) 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 }