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 batchDelay time.Duration reconnectEvery int upSizer, downSizer *adaptiveSizer control, upload, download *txnLane upPackets, downPackets, upBytes, downBytes atomic.Uint64 upBatches, downBatches, localDropped 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, batchDelay time.Duration, adaptLog bool) (*vpnClient, error) { sid, err := randomSID() if err != nil { return nil, err } if batchDelay < 0 { batchDelay = 0 } return &vpnClient{tun: tun, sid: sid, serverAddr: addr, token: token, ipv4: v4, ipv6: v6, mtu: mtu, timeout: timeout, batchDelay: batchDelay, 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) logLocalDrop(reason string) { n := v.localDropped.Add(1) // Link-local/control traffic can be noisy. Keep it visible without filling // the Android live log or making a harmless packet fatal to the VPN. if n <= 8 || n%256 == 0 { fmt.Printf("VPN DROP local packet (%s) dropped=%d\n", reason, n) } } func (v *vpnClient) tunReadLoop(out chan<- []byte, errs chan<- error) { buf := make([]byte, protocol.VPNMaxPacket) for { n, err := v.tun.Read(buf) if err != nil { errs <- err return } if n < 1 || n > protocol.VPNMaxPacket { continue } packet := append([]byte(nil), buf[:n]...) src, _, err := protocol.PacketAddresses(packet) if err != nil { v.logLocalDrop(err.Error()) continue } if src != v.ipv4 && src != v.ipv6 { v.logLocalDrop(fmt.Sprintf("source %s is not assigned VPN address", src)) continue } select { case out <- packet: case <-v.stopped: return } } } func batchWireSize(packets [][]byte) int { n := 1 for _, p := range packets { n += 2 + len(p) } return n } func (v *vpnClient) uploadLoop(in <-chan []byte, errs chan<- error) { var seq uint32 var carry []byte for { var first []byte if carry != nil { first, carry = carry, nil } else { select { case first = <-in: case <-v.stopped: return } } packets := [][]byte{first} encodedSize := 1 + 2 + len(first) timer := time.NewTimer(v.batchDelay) collect: for encodedSize < protocol.VPNMaxBatch { select { case p := <-in: need := 2 + len(p) if encodedSize+need > protocol.VPNMaxBatch { carry = p break collect } packets = append(packets, p) encodedSize += need case <-timer.C: break collect case <-v.stopped: if !timer.Stop() { select { case <-timer.C: default: } } return } } if !timer.Stop() { select { case <-timer.C: default: } } batch, err := protocol.BuildVPNBatch(packets) if err != nil { errs <- err return } offset := 0 for offset < len(batch) { limit := v.upSizer.Current() size := len(batch) - offset if size > limit { size = limit } req, e := protocol.BuildVPNPush(v.sid, seq, offset, len(batch), batch[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 > len(batch) { errs <- errors.New("bad server upload ACK") return } v.upSizer.Success(size, size == limit) offset = accepted } var rawBytes uint64 for _, p := range packets { rawBytes += uint64(len(p)) } v.upPackets.Add(uint64(len(packets))) v.upBytes.Add(rawBytes) v.upBatches.Add(1) seq++ } } func (v *vpnClient) downloadLoop(errs chan<- error) { var want uint32 ack := protocol.VPNNoAck offset := 0 var transfer []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 > protocol.VPNMaxBatch { errs <- errors.New("bad server download sequence") return } if offset == 0 { total = rtotal transfer = make([]byte, 0, total) } else if rtotal != total { errs <- errors.New("download transfer size changed") return } transfer = append(transfer, data...) offset += len(data) v.downSizer.Success(len(data), len(data) == limit) if offset < total { continue } if offset != total { errs <- errors.New("download transfer overflow") return } packets, e := protocol.ParseVPNBatch(transfer) if e != nil { // Compatibility with the first packet-VPN build, which used one raw // IP packet as each transfer object. if len(transfer) > 0 && (transfer[0]>>4 == 4 || transfer[0]>>4 == 6) { packets = [][]byte{transfer} } else { errs <- e return } } var rawBytes uint64 for _, packet := range packets { n, e := v.tun.Write(packet) if e != nil { errs <- e return } if n != len(packet) { errs <- io.ErrShortWrite return } rawBytes += uint64(n) } v.downPackets.Add(uint64(len(packets))) v.downBytes.Add(rawBytes) v.downBatches.Add(1) ack = want want++ offset = 0 transfer = nil total = 0 } } func (v *vpnClient) run() error { if err := v.open(); err != nil { return err } fmt.Println("VPN READY") errs := make(chan error, 3) packets := make(chan []byte, 256) go v.tunReadLoop(packets, errs) go v.uploadLoop(packets, 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_batches=%d down_batches=%d up_bytes=%d down_bytes=%d local_dropped=%d upload_chunk=%d download_chunk=%d pollers=1\n", v.upPackets.Load(), v.downPackets.Load(), v.upBatches.Load(), v.downBatches.Load(), v.upBytes.Load(), v.downBytes.Load(), v.localDropped.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", protocol.VPNMaxFragment, "maximum adaptive record bytes (up to 1 MiB)") chunkMin := flag.Int("chunk-min", 32, "minimum adaptive record bytes") chunkStart := flag.Int("chunk-start", protocol.VPNMaxFragment, "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") batchDelay := flag.Duration("batch-delay", time.Millisecond, "maximum delay used to combine adjacent TUN packets into one transfer object") 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, *batchDelay, *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() }