Files
DragonTCP/core/cmd/dragontcp-vpn-client/main.go
T
2026-08-16 02:33:07 -03:00

514 lines
13 KiB
Go

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()
}