package main import ( "crypto/rand" "errors" "flag" "fmt" "io" "net" "net/netip" "os" "os/signal" "strconv" "sync" "sync/atomic" "syscall" "time" "dragontcpvpn/internal/protocol" ) var requestCounter atomic.Uint32 type txnLane struct { mu sync.Mutex serverAddr string timeout time.Duration reconnectEvery int conn net.Conn count int closed bool } func newTxnLane(addr string, timeout time.Duration, reconnectEvery int) *txnLane { return &txnLane{serverAddr: addr, timeout: timeout, reconnectEvery: reconnectEvery} } func (l *txnLane) closeLocked() { if l.conn != nil { _ = l.conn.Close() l.conn = nil } l.count = 0 } func (l *txnLane) Close() { l.mu.Lock(); l.closed = true; l.closeLocked(); l.mu.Unlock() } func (l *txnLane) ensureConn() error { if l.closed { return net.ErrClosed } if l.conn != nil && (l.reconnectEvery <= 0 || l.count < l.reconnectEvery) { return nil } l.closeLocked() d := net.Dialer{Timeout: 10 * time.Second, KeepAlive: 30 * time.Second} c, err := d.Dial("tcp", l.serverAddr) if err != nil { return err } protocol.TuneTCP(c) l.conn = c return nil } func (l *txnLane) Do(payload []byte) ([]byte, error) { l.mu.Lock() defer l.mu.Unlock() if err := l.ensureConn(); err != nil { return nil, err } timeout := l.timeout if timeout <= 0 { timeout = 3 * time.Second } _ = l.conn.SetDeadline(time.Now().Add(timeout)) id := requestCounter.Add(1) if err := protocol.WriteRequestFrame(l.conn, id, payload); err != nil { l.closeLocked() return nil, err } rid, resp, err := protocol.ReadResponseFrame(l.conn) if err != nil { l.closeLocked() return nil, err } if rid != id { l.closeLocked() return nil, errors.New("request ID mismatch") } l.count++ _ = l.conn.SetDeadline(time.Time{}) return resp, nil } func doControl(l *txnLane, payload []byte) ([]byte, error) { var last error for i := 0; i < 6; i++ { r, e := l.Do(payload) if e == nil { return r, nil } last = e time.Sleep(time.Duration(i+1) * 50 * time.Millisecond) } return nil, last } type adaptiveSizer struct { mu sync.Mutex name string current, min, max int successes int growAfter int log bool } func newSizer(name string, start, min, max, growAfter int, log bool) *adaptiveSizer { if min < 32 { min = 32 } if max > protocol.VPNMaxFragment { max = protocol.VPNMaxFragment } if max < min { max = min } if start < min { start = min } if start > max { start = max } if growAfter < 1 { growAfter = 32 } return &adaptiveSizer{name: name, current: start, min: min, max: max, growAfter: growAfter, log: log} } func (s *adaptiveSizer) Current() int { s.mu.Lock(); v := s.current; s.mu.Unlock(); return v } func (s *adaptiveSizer) Failure(actual int) { s.mu.Lock() defer s.mu.Unlock() old := s.current s.successes = 0 basis := actual if basis <= 0 || basis > old { basis = old } next := basis / 2 if next < s.min { next = s.min } if next >= old && old > s.min { next = old / 2 if next < s.min { next = s.min } } if next < old { s.current = next if s.log { fmt.Printf("adaptive %s chunk: %d -> %d after transport failure (record=%d)\n", s.name, old, next, actual) } } } func (s *adaptiveSizer) Success(actual int, full bool) { s.mu.Lock() defer s.mu.Unlock() if s.current >= s.max || !full { return } s.successes++ if s.successes < s.growAfter { return } s.successes = 0 old := s.current step := old / 4 if step < 32 { step = 32 } next := old + step if next > s.max { next = s.max } if next > old { s.current = next if s.log { fmt.Printf("adaptive %s chunk: %d -> %d after stable success\n", s.name, old, next) } } } func receiveTunFD(path string, timeout time.Duration) (*os.File, error) { _ = os.Remove(path) addr := &net.UnixAddr{Name: path, Net: "unix"} ln, err := net.ListenUnix("unix", addr) if err != nil { return nil, err } defer func() { ln.Close(); os.Remove(path) }() _ = os.Chmod(path, 0600) fmt.Printf("TUNFD READY %s\n", path) _ = ln.SetDeadline(time.Now().Add(timeout)) c, err := ln.AcceptUnix() if err != nil { return nil, err } defer c.Close() buf := make([]byte, 1) oob := make([]byte, 128) n, oobn, _, _, err := c.ReadMsgUnix(buf, oob) if err != nil { return nil, err } if n < 1 { return nil, errors.New("missing TUN fd marker") } msgs, err := syscall.ParseSocketControlMessage(oob[:oobn]) if err != nil { return nil, err } for _, m := range msgs { fds, e := syscall.ParseUnixRights(&m) if e == nil && len(fds) > 0 { return os.NewFile(uintptr(fds[0]), "android-tun"), nil } } return nil, errors.New("TUN file descriptor was not received") } func randomSID() (protocol.VPNSessionID, error) { var sid protocol.VPNSessionID _, err := io.ReadFull(rand.Reader, sid[:]) return sid, err } type vpnClient struct { tun *os.File sid protocol.VPNSessionID serverAddr string token string ipv4, ipv6 netip.Addr mtu int timeout time.Duration reconnectEvery int upSizer, downSizer *adaptiveSizer control, upload, download *txnLane upPackets, downPackets, upBytes, downBytes atomic.Uint64 stopped chan struct{} stopOnce sync.Once } func newVPNClient(tun *os.File, addr, token string, v4, v6 netip.Addr, mtu, start, min, max, growAfter, reconnectEvery int, timeout time.Duration, adaptLog bool) (*vpnClient, error) { sid, err := randomSID() if err != nil { return nil, err } return &vpnClient{tun: tun, sid: sid, serverAddr: addr, token: token, ipv4: v4, ipv6: v6, mtu: mtu, timeout: timeout, reconnectEvery: reconnectEvery, upSizer: newSizer("upload", start, min, max, growAfter, adaptLog), downSizer: newSizer("download", start, min, max, growAfter, adaptLog), control: newTxnLane(addr, timeout, reconnectEvery), upload: newTxnLane(addr, timeout, reconnectEvery), download: newTxnLane(addr, timeout, reconnectEvery), stopped: make(chan struct{})}, nil } func (v *vpnClient) open() error { req, err := protocol.BuildVPNOpen(v.sid, v.token, v.ipv4, v.ipv6, v.mtu) if err != nil { return err } resp, err := doControl(v.control, req) if err != nil { return err } max, err := protocol.ParseVPNOpened(resp) if err != nil { return err } if max < v.upSizer.max { v.upSizer.max = max if v.upSizer.current > max { v.upSizer.current = max } } if max < v.downSizer.max { v.downSizer.max = max if v.downSizer.current > max { v.downSizer.current = max } } fmt.Printf("VPN SESSION OPEN ipv4=%s ipv6=%s mtu=%d server_chunk_max=%d\n", v.ipv4, v.ipv6, v.mtu, max) return nil } func (v *vpnClient) close() { v.stopOnce.Do(func() { close(v.stopped) if p, err := protocol.BuildVPNClose(v.sid), error(nil); err == nil { _, _ = v.control.Do(p) } v.control.Close() v.upload.Close() v.download.Close() _ = v.tun.Close() }) } func (v *vpnClient) uploadLoop(errs chan<- error) { buf := make([]byte, 65535) var seq uint32 for { n, err := v.tun.Read(buf) if err != nil { errs <- err return } if n < 1 { continue } packet := append([]byte(nil), buf[:n]...) if n > 65535 { continue } offset := 0 for offset < n { limit := v.upSizer.Current() size := n - offset if size > limit { size = limit } req, e := protocol.BuildVPNPush(v.sid, seq, offset, n, packet[offset:offset+size]) if e != nil { errs <- e return } resp, e := v.upload.Do(req) if e != nil { v.upSizer.Failure(size) continue } rseq, accepted, e := protocol.ParseVPNAck(resp) if e != nil { errs <- e return } if rseq != seq || accepted < offset || accepted > n { errs <- errors.New("bad server upload ACK") return } fullRecord := size == limit v.upSizer.Success(size, fullRecord) offset = accepted } v.upPackets.Add(1) v.upBytes.Add(uint64(n)) seq++ } } func (v *vpnClient) downloadLoop(errs chan<- error) { var want uint32 ack := protocol.VPNNoAck offset := 0 var packet []byte total := 0 for { limit := v.downSizer.Current() req, e := protocol.BuildVPNPull(v.sid, ack, want, offset, limit) if e != nil { errs <- e return } resp, e := v.download.Do(req) if e != nil { v.downSizer.Failure(limit) continue } seq, roff, rtotal, data, wait, e := protocol.ParseVPNData(resp) if e != nil { errs <- e return } if wait { continue } if seq != want || roff != offset || rtotal < 1 || rtotal > 65535 { errs <- errors.New("bad server download sequence") return } if offset == 0 { total = rtotal packet = make([]byte, 0, total) } else if rtotal != total { errs <- errors.New("download packet size changed") return } packet = append(packet, data...) offset += len(data) v.downSizer.Success(len(data), len(data) == limit) if offset < total { continue } if offset != total { errs <- errors.New("download packet overflow") return } n, e := v.tun.Write(packet) if e != nil { errs <- e return } if n != len(packet) { errs <- io.ErrShortWrite return } v.downPackets.Add(1) v.downBytes.Add(uint64(n)) ack = want want++ offset = 0 packet = nil total = 0 } } func (v *vpnClient) run() error { if err := v.open(); err != nil { return err } fmt.Println("VPN READY") errs := make(chan error, 2) go v.uploadLoop(errs) go v.downloadLoop(errs) ticker := time.NewTicker(5 * time.Second) defer ticker.Stop() for { select { case err := <-errs: return err case <-ticker.C: fmt.Printf("STATS up_packets=%d down_packets=%d up_bytes=%d down_bytes=%d upload_chunk=%d download_chunk=%d pollers=1\n", v.upPackets.Load(), v.downPackets.Load(), v.upBytes.Load(), v.downBytes.Load(), v.upSizer.Current(), v.downSizer.Current()) case <-v.stopped: return nil } } } func main() { serverHost := flag.String("server-host", "", "DragonTCP VPN server host/IP") serverPort := flag.Int("server-port", 53, "DragonTCP VPN server TCP port") token := flag.String("token", "change-this-token", "shared token") tunFDSocket := flag.String("tun-fd-socket", "", "Unix socket path used by Android to pass the VpnService TUN fd") tunFD := flag.Int("tun-fd", -1, "existing TUN fd for testing/non-Android use") ipv4Text := flag.String("vpn-ipv4", "10.123.0.2", "client VPN IPv4 address") ipv6Text := flag.String("vpn-ipv6", "fd7a:4472:6167:6f6e::2", "client VPN IPv6 address") mtu := flag.Int("vpn-mtu", 1280, "VPN interface MTU") chunkMax := flag.Int("chunk-max", 65535, "maximum adaptive record bytes") chunkMin := flag.Int("chunk-min", 32, "minimum adaptive record bytes") chunkStart := flag.Int("chunk-start", 65535, "starting record bytes; app sets this equal to max") growAfter := flag.Int("chunk-grow-after", 64, "full successful records before increasing chunk size") timeout := flag.Duration("chunk-timeout", 2*time.Second, "framed transaction timeout") reconnectEvery := flag.Int("chunk-reconnect-every", 32, "reconnect a TCP/53 lane after this many transactions; 0 keeps it open") adaptLog := flag.Bool("chunk-adapt-log", false, "log adaptive chunk changes") flag.Parse() if *serverHost == "" { fmt.Fprintln(os.Stderr, "--server-host is required") os.Exit(2) } if *serverPort < 1 || *serverPort > 65535 { fmt.Fprintln(os.Stderr, "invalid server port") os.Exit(2) } if *chunkMin < 32 || *chunkMax > protocol.VPNMaxFragment || *chunkMin > *chunkMax { fmt.Fprintf(os.Stderr, "chunks must satisfy 32 <= min <= max <= %d\n", protocol.VPNMaxFragment) os.Exit(2) } if *chunkStart < *chunkMin { *chunkStart = *chunkMin } if *chunkStart > *chunkMax { *chunkStart = *chunkMax } v4, err := netip.ParseAddr(*ipv4Text) if err != nil || !v4.Is4() { fmt.Fprintln(os.Stderr, "invalid --vpn-ipv4") os.Exit(2) } v6, err := netip.ParseAddr(*ipv6Text) if err != nil || !v6.Is6() { fmt.Fprintln(os.Stderr, "invalid --vpn-ipv6") os.Exit(2) } var tun *os.File if *tunFD >= 0 { tun = os.NewFile(uintptr(*tunFD), "tun") } else { if *tunFDSocket == "" { fmt.Fprintln(os.Stderr, "--tun-fd-socket is required on Android") os.Exit(2) } tun, err = receiveTunFD(*tunFDSocket, 10*time.Second) if err != nil { fmt.Fprintln(os.Stderr, "receive TUN fd:", err) os.Exit(1) } } addr := net.JoinHostPort(*serverHost, strconv.Itoa(*serverPort)) client, err := newVPNClient(tun, addr, *token, v4, v6, *mtu, *chunkStart, *chunkMin, *chunkMax, *growAfter, *reconnectEvery, *timeout, *adaptLog) if err != nil { fmt.Fprintln(os.Stderr, err) os.Exit(1) } sig := make(chan os.Signal, 1) signal.Notify(sig, syscall.SIGINT, syscall.SIGTERM) go func() { <-sig; client.close() }() if err := client.run(); err != nil && !errors.Is(err, os.ErrClosed) && !errors.Is(err, net.ErrClosed) { fmt.Fprintln(os.Stderr, "VPN stopped:", err) client.close() os.Exit(1) } client.close() }