514 lines
13 KiB
Go
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()
|
|
}
|