This commit is contained in:
2026-08-16 02:33:07 -03:00
parent 14beee38b0
commit 6dac260155
33 changed files with 2065 additions and 2301 deletions
+513
View File
@@ -0,0 +1,513 @@
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()
}
+743
View File
@@ -0,0 +1,743 @@
package main
import (
"crypto/subtle"
"encoding/hex"
"errors"
"flag"
"fmt"
"io"
"net"
"net/netip"
"os"
"os/exec"
"os/signal"
"strconv"
"strings"
"sync"
"sync/atomic"
"syscall"
"time"
"unsafe"
"dragontcpvpn/internal/protocol"
)
const (
defaultVPNv4Prefix = "10.123.0.0/16"
defaultVPNv6Prefix = "fd7a:4472:6167:6f6e::/64"
)
type debugStats struct {
enabled bool
packets bool
started time.Time
activeConns atomic.Int64
activeSessions atomic.Int64
upPackets atomic.Uint64
downPackets atomic.Uint64
upBytes atomic.Uint64
downBytes atomic.Uint64
dropped atomic.Uint64
errors atomic.Uint64
}
func (d *debugStats) logf(format string, args ...any) {
if d != nil && d.enabled {
fmt.Printf("[DEBUG] "+format+"\n", args...)
}
}
func (d *debugStats) packetf(format string, args ...any) {
if d != nil && d.packets {
fmt.Printf("[PACKET] "+format+"\n", args...)
}
}
func (d *debugStats) errorf(format string, args ...any) {
if d != nil {
d.errors.Add(1)
if d.enabled {
fmt.Printf("[ERROR] "+format+"\n", args...)
}
}
}
func tokenEqual(a, b string) bool {
if len(a) != len(b) {
return false
}
return subtle.ConstantTimeCompare([]byte(a), []byte(b)) == 1
}
type vpnSession struct {
sid protocol.VPNSessionID
ipv4 netip.Addr
ipv6 netip.Addr
mtu int
maxChunk int
maxPackets int
manager *vpnManager
mu sync.Mutex
notify chan struct{}
packets map[uint32][]byte
nextDown uint32
closed bool
lastSeen time.Time
upMu sync.Mutex
expectedUp uint32
currentSeq uint32
currentTotal int
currentBuf []byte
haveCurrent bool
lastComplete uint32
lastCompleteTotal int
haveLastComplete bool
}
func newVPNSession(m *vpnManager, sid protocol.VPNSessionID, v4, v6 netip.Addr, mtu, maxChunk, maxPackets int) *vpnSession {
return &vpnSession{
sid: sid, ipv4: v4, ipv6: v6, mtu: mtu, maxChunk: maxChunk, maxPackets: maxPackets,
manager: m, notify: make(chan struct{}), packets: make(map[uint32][]byte, maxPackets), lastSeen: time.Now(),
}
}
func (s *vpnSession) signalLocked() {
close(s.notify)
s.notify = make(chan struct{})
}
func (s *vpnSession) touchLocked() { s.lastSeen = time.Now() }
func (s *vpnSession) touch() { s.mu.Lock(); s.touchLocked(); s.mu.Unlock() }
func (s *vpnSession) enqueue(packet []byte) bool {
if len(packet) == 0 || len(packet) > 65535 {
return false
}
s.mu.Lock()
defer s.mu.Unlock()
if s.closed {
return false
}
if len(s.packets) >= s.maxPackets {
if s.manager.debug != nil {
s.manager.debug.dropped.Add(1)
}
return false
}
seq := s.nextDown
s.nextDown++
s.packets[seq] = append([]byte(nil), packet...)
s.touchLocked()
s.signalLocked()
if s.manager.debug != nil {
s.manager.debug.downPackets.Add(1)
s.manager.debug.downBytes.Add(uint64(len(packet)))
s.manager.debug.packetf("QUEUE sid=%s seq=%d bytes=%d", shortSID(s.sid), seq, len(packet))
}
return true
}
func (s *vpnSession) push(seq uint32, offset, total int, data []byte) (int, error) {
s.upMu.Lock()
defer s.upMu.Unlock()
if total < 1 || total > 65535 || len(data) < 1 || len(data) > s.maxChunk || offset < 0 || offset+len(data) > total {
return 0, errors.New("invalid packet fragment")
}
if s.haveLastComplete && seq == s.lastComplete {
s.touch()
return s.lastCompleteTotal, nil
}
if seq < s.expectedUp {
return 0, fmt.Errorf("old upload sequence %d", seq)
}
if seq > s.expectedUp {
return 0, fmt.Errorf("upload sequence %d expected %d", seq, s.expectedUp)
}
if !s.haveCurrent {
if offset != 0 {
return 0, errors.New("first fragment offset must be zero")
}
s.haveCurrent = true
s.currentSeq = seq
s.currentTotal = total
s.currentBuf = make([]byte, 0, total)
}
if s.currentSeq != seq || s.currentTotal != total {
return 0, errors.New("packet fragment metadata changed")
}
// Idempotent retry: if this exact offset was already accepted, acknowledge
// the existing bytes instead of appending duplicate data.
if offset < len(s.currentBuf) {
end := offset + len(data)
if end <= len(s.currentBuf) && string(s.currentBuf[offset:end]) == string(data) {
return len(s.currentBuf), nil
}
return 0, errors.New("retry fragment does not match accepted data")
}
if offset != len(s.currentBuf) {
return 0, fmt.Errorf("fragment offset %d expected %d", offset, len(s.currentBuf))
}
s.currentBuf = append(s.currentBuf, data...)
accepted := len(s.currentBuf)
if accepted < total {
s.touch()
return accepted, nil
}
packet := append([]byte(nil), s.currentBuf...)
s.haveCurrent = false
s.currentBuf = nil
if err := s.manager.acceptClientPacket(s, packet); err != nil {
return 0, err
}
s.lastComplete = seq
s.lastCompleteTotal = total
s.haveLastComplete = true
s.expectedUp++
s.touch()
if s.manager.debug != nil {
s.manager.debug.upPackets.Add(1)
s.manager.debug.upBytes.Add(uint64(len(packet)))
s.manager.debug.packetf("UP sid=%s seq=%d bytes=%d", shortSID(s.sid), seq, len(packet))
}
return accepted, nil
}
func (s *vpnSession) pull(ack, want uint32, offset, limit int, wait time.Duration) ([]byte, int, bool, error) {
if offset < 0 || limit < 1 || limit > s.maxChunk {
return nil, 0, false, errors.New("invalid pull")
}
timer := time.NewTimer(wait)
defer timer.Stop()
for {
s.mu.Lock()
s.touchLocked()
if ack != protocol.VPNNoAck {
for seq := range s.packets {
if seq <= ack {
delete(s.packets, seq)
}
}
}
if packet, ok := s.packets[want]; ok {
if offset >= len(packet) {
s.mu.Unlock()
return nil, len(packet), false, errors.New("pull offset beyond packet")
}
end := offset + limit
if end > len(packet) {
end = len(packet)
}
out := append([]byte(nil), packet[offset:end]...)
total := len(packet)
s.mu.Unlock()
return out, total, false, nil
}
if s.closed {
s.mu.Unlock()
return nil, 0, false, net.ErrClosed
}
ch := s.notify
s.mu.Unlock()
select {
case <-ch:
case <-timer.C:
return nil, 0, true, nil
}
}
}
func (s *vpnSession) close() {
s.mu.Lock()
if !s.closed {
s.closed = true
s.signalLocked()
}
s.mu.Unlock()
}
type vpnManager struct {
mu sync.RWMutex
sessions map[protocol.VPNSessionID]*vpnSession
byIPv4 map[netip.Addr]*vpnSession
byIPv6 map[netip.Addr]*vpnSession
maxChunk int
maxPackets int
pollWait time.Duration
timeout time.Duration
tun *os.File
tunWriteMu sync.Mutex
mockEcho bool
allowPrivate bool
debug *debugStats
v4Prefix netip.Prefix
v6Prefix netip.Prefix
}
func newVPNManager(tun *os.File, mockEcho bool, maxChunk, maxPackets int, pollWait, timeout time.Duration, allowPrivate bool, debug *debugStats) *vpnManager {
v4p := netip.MustParsePrefix(defaultVPNv4Prefix)
v6p := netip.MustParsePrefix(defaultVPNv6Prefix)
m := &vpnManager{
sessions: make(map[protocol.VPNSessionID]*vpnSession), byIPv4: make(map[netip.Addr]*vpnSession), byIPv6: make(map[netip.Addr]*vpnSession),
maxChunk: maxChunk, maxPackets: maxPackets, pollWait: pollWait, timeout: timeout, tun: tun, mockEcho: mockEcho, allowPrivate: allowPrivate, debug: debug,
v4Prefix: v4p, v6Prefix: v6p,
}
if tun != nil {
go m.tunReadLoop()
}
go m.cleanupLoop()
return m
}
func (m *vpnManager) addOrGet(sid protocol.VPNSessionID, v4, v6 netip.Addr, mtu int) (*vpnSession, error) {
if !m.v4Prefix.Contains(v4) || v4 == netip.MustParseAddr("10.123.0.1") {
return nil, errors.New("client IPv4 outside DragonTCP subnet")
}
if !m.v6Prefix.Contains(v6) || v6 == netip.MustParseAddr("fd7a:4472:6167:6f6e::1") {
return nil, errors.New("client IPv6 outside DragonTCP subnet")
}
if mtu < 576 || mtu > 9000 {
return nil, errors.New("invalid client MTU")
}
m.mu.Lock()
defer m.mu.Unlock()
if old := m.sessions[sid]; old != nil {
if old.ipv4 != v4 || old.ipv6 != v6 {
return nil, errors.New("session address mismatch")
}
old.touch()
return old, nil
}
if m.byIPv4[v4] != nil || m.byIPv6[v6] != nil {
return nil, errors.New("client VPN address already in use")
}
s := newVPNSession(m, sid, v4, v6, mtu, m.maxChunk, m.maxPackets)
m.sessions[sid] = s
m.byIPv4[v4] = s
m.byIPv6[v6] = s
if m.debug != nil {
m.debug.activeSessions.Add(1)
m.debug.logf("SESSION OPEN sid=%s ipv4=%s ipv6=%s mtu=%d", shortSID(sid), v4, v6, mtu)
}
return s, nil
}
func (m *vpnManager) get(sid protocol.VPNSessionID) *vpnSession {
m.mu.RLock()
s := m.sessions[sid]
m.mu.RUnlock()
return s
}
func (m *vpnManager) remove(sid protocol.VPNSessionID) {
m.mu.Lock()
s := m.sessions[sid]
if s != nil {
delete(m.sessions, sid)
delete(m.byIPv4, s.ipv4)
delete(m.byIPv6, s.ipv6)
}
m.mu.Unlock()
if s != nil {
s.close()
if m.debug != nil {
m.debug.activeSessions.Add(-1)
m.debug.logf("SESSION CLOSE sid=%s", shortSID(sid))
}
}
}
func (m *vpnManager) cleanupLoop() {
ticker := time.NewTicker(30 * time.Second)
defer ticker.Stop()
for range ticker.C {
cutoff := time.Now().Add(-m.timeout)
var stale []protocol.VPNSessionID
m.mu.RLock()
for sid, s := range m.sessions {
s.mu.Lock()
last := s.lastSeen
closed := s.closed
s.mu.Unlock()
if closed || last.Before(cutoff) {
stale = append(stale, sid)
}
}
m.mu.RUnlock()
for _, sid := range stale {
m.remove(sid)
}
}
}
func packetAddresses(packet []byte) (src, dst netip.Addr, err error) {
if len(packet) < 1 {
return src, dst, errors.New("empty IP packet")
}
switch packet[0] >> 4 {
case 4:
if len(packet) < 20 {
return src, dst, errors.New("short IPv4 packet")
}
total := int(packet[2])<<8 | int(packet[3])
if total < 20 || total > len(packet) {
return src, dst, errors.New("invalid IPv4 total length")
}
var a, b [4]byte
copy(a[:], packet[12:16])
copy(b[:], packet[16:20])
return netip.AddrFrom4(a), netip.AddrFrom4(b), nil
case 6:
if len(packet) < 40 {
return src, dst, errors.New("short IPv6 packet")
}
total := 40 + (int(packet[4])<<8 | int(packet[5]))
if total > len(packet) {
return src, dst, errors.New("invalid IPv6 payload length")
}
var a, b [16]byte
copy(a[:], packet[8:24])
copy(b[:], packet[24:40])
return netip.AddrFrom16(a), netip.AddrFrom16(b), nil
default:
return src, dst, errors.New("unsupported IP version")
}
}
func destinationAllowed(dst netip.Addr, allowPrivate bool) bool {
if dst.IsUnspecified() || dst.IsMulticast() {
return false
}
if allowPrivate {
return true
}
if dst.IsLoopback() || dst.IsLinkLocalUnicast() || dst.IsPrivate() {
return false
}
return true
}
func (m *vpnManager) acceptClientPacket(s *vpnSession, packet []byte) error {
src, dst, err := packetAddresses(packet)
if err != nil {
return err
}
if src != s.ipv4 && src != s.ipv6 {
return fmt.Errorf("source %s does not match session address", src)
}
if !destinationAllowed(dst, m.allowPrivate) {
return fmt.Errorf("destination %s is blocked; use --allow-private to permit it", dst)
}
if m.mockEcho {
s.enqueue(packet)
return nil
}
if m.tun == nil {
return errors.New("VPN TUN is unavailable")
}
m.tunWriteMu.Lock()
n, err := m.tun.Write(packet)
m.tunWriteMu.Unlock()
if err != nil {
return err
}
if n != len(packet) {
return io.ErrShortWrite
}
return nil
}
func (m *vpnManager) tunReadLoop() {
buf := make([]byte, 65535)
for {
n, err := m.tun.Read(buf)
if err != nil {
if m.debug != nil {
m.debug.errorf("TUN read: %v", err)
}
return
}
if n < 1 {
continue
}
packet := append([]byte(nil), buf[:n]...)
_, dst, e := packetAddresses(packet)
if e != nil {
continue
}
m.mu.RLock()
var s *vpnSession
if dst.Is4() {
s = m.byIPv4[dst]
} else {
s = m.byIPv6[dst]
}
m.mu.RUnlock()
if s != nil {
s.enqueue(packet)
}
}
}
func shortSID(sid protocol.VPNSessionID) string { return hex.EncodeToString(sid[:4]) }
func processVPN(conn net.Conn, requestID uint32, payload []byte, token string, m *vpnManager) error {
switch payload[0] {
case protocol.VPNCmdOpen:
sid, tok, v4, v6, mtu, err := protocol.ParseVPNOpen(payload)
if err != nil {
return protocol.WriteResponseFrame(conn, requestID, protocol.VPNError(err.Error()))
}
if !tokenEqual(tok, token) {
return protocol.WriteResponseFrame(conn, requestID, protocol.VPNError("authentication failed"))
}
_, err = m.addOrGet(sid, v4, v6, mtu)
if err != nil {
return protocol.WriteResponseFrame(conn, requestID, protocol.VPNError(err.Error()))
}
return protocol.WriteResponseFrame(conn, requestID, protocol.BuildVPNOpened(m.maxChunk))
case protocol.VPNCmdPush:
sid, seq, offset, total, data, err := protocol.ParseVPNPush(payload)
if err != nil {
return protocol.WriteResponseFrame(conn, requestID, protocol.VPNError(err.Error()))
}
s := m.get(sid)
if s == nil {
return protocol.WriteResponseFrame(conn, requestID, protocol.VPNError("unknown VPN session"))
}
accepted, err := s.push(seq, offset, total, data)
if err != nil {
return protocol.WriteResponseFrame(conn, requestID, protocol.VPNError(err.Error()))
}
return protocol.WriteResponseFrame(conn, requestID, protocol.BuildVPNAck(seq, accepted))
case protocol.VPNCmdPull:
sid, ack, want, offset, limit, err := protocol.ParseVPNPull(payload)
if err != nil {
return protocol.WriteResponseFrame(conn, requestID, protocol.VPNError(err.Error()))
}
s := m.get(sid)
if s == nil {
return protocol.WriteResponseFrame(conn, requestID, protocol.VPNError("unknown VPN session"))
}
if limit > s.maxChunk {
limit = s.maxChunk
}
data, total, wait, err := s.pull(ack, want, offset, limit, m.pollWait)
if err != nil {
return protocol.WriteResponseFrame(conn, requestID, protocol.VPNError(err.Error()))
}
if wait {
return protocol.WriteResponseFrame(conn, requestID, []byte{protocol.VPNRespWait})
}
m.debug.packetf("DOWN sid=%s seq=%d offset=%d bytes=%d total=%d", shortSID(sid), want, offset, len(data), total)
return protocol.WriteResponseFrame(conn, requestID, protocol.BuildVPNData(want, offset, total, data))
case protocol.VPNCmdClose:
sid, err := protocol.ParseVPNClose(payload)
if err != nil {
return protocol.WriteResponseFrame(conn, requestID, protocol.VPNError(err.Error()))
}
m.remove(sid)
return protocol.WriteResponseFrame(conn, requestID, []byte{protocol.VPNRespClosed})
default:
return protocol.WriteResponseFrame(conn, requestID, protocol.VPNError("unknown VPN command"))
}
}
func handleConn(conn net.Conn, token string, m *vpnManager, slots chan struct{}, debug *debugStats) {
defer func() { <-slots; debug.activeConns.Add(-1); _ = conn.Close() }()
protocol.TuneTCP(conn)
for {
_ = conn.SetDeadline(time.Now().Add(30 * time.Second))
requestID, _, payload, err := protocol.ReadRequestFrame(conn)
if err != nil {
if !errors.Is(err, io.EOF) && !errors.Is(err, net.ErrClosed) {
debug.errorf("peer=%v read: %v", conn.RemoteAddr(), err)
}
return
}
if !protocol.IsVPNCommand(payload) {
_ = protocol.WriteResponseFrame(conn, requestID, protocol.VPNError("this binary accepts DragonTCP VPN packet commands only"))
continue
}
if err := processVPN(conn, requestID, payload, token, m); err != nil {
return
}
}
}
// Linux TUN setup.
type ifreq struct {
Name [16]byte
Flags uint16
_ [22]byte
}
const tunSetIFF = 0x400454ca
const iffTun = 0x0001
const iffNoPI = 0x1000
func openTun(name string) (*os.File, error) {
fd, err := syscall.Open("/dev/net/tun", syscall.O_RDWR|syscall.O_CLOEXEC, 0)
if err != nil {
return nil, err
}
var req ifreq
copy(req.Name[:], []byte(name))
req.Flags = iffTun | iffNoPI
_, _, errno := syscall.Syscall(syscall.SYS_IOCTL, uintptr(fd), uintptr(tunSetIFF), uintptr(unsafe.Pointer(&req)))
if errno != 0 {
syscall.Close(fd)
return nil, errno
}
return os.NewFile(uintptr(fd), name), nil
}
func run(cmd string, args ...string) error {
c := exec.Command(cmd, args...)
out, err := c.CombinedOutput()
if err != nil {
return fmt.Errorf("%s %s: %v: %s", cmd, strings.Join(args, " "), err, strings.TrimSpace(string(out)))
}
return nil
}
func runOptional(debug *debugStats, cmd string, args ...string) {
if err := run(cmd, args...); err != nil {
debug.logf("optional command failed: %v", err)
}
}
func ensureRule(debug *debugStats, binary string, argsCheck, argsAdd []string) {
if err := exec.Command(binary, argsCheck...).Run(); err == nil {
return
}
if err := run(binary, argsAdd...); err != nil {
debug.logf("NAT rule warning: %v", err)
}
}
func setupLinuxVPN(tunName string, mtu int, autoNAT bool, debug *debugStats) (*os.File, error) {
tun, err := openTun(tunName)
if err != nil {
return nil, fmt.Errorf("open /dev/net/tun: %w", err)
}
fail := func(e error) (*os.File, error) { tun.Close(); return nil, e }
if err := run("ip", "link", "set", "dev", tunName, "mtu", strconv.Itoa(mtu)); err != nil {
return fail(err)
}
if err := run("ip", "addr", "replace", "10.123.0.1/16", "dev", tunName); err != nil {
return fail(err)
}
// IPv6 may be disabled on some hosts; report clearly instead of silently bypassing it.
if err := run("ip", "-6", "addr", "replace", "fd7a:4472:6167:6f6e::1/64", "dev", tunName); err != nil {
return fail(err)
}
if err := run("ip", "link", "set", "dev", tunName, "up"); err != nil {
return fail(err)
}
if err := os.WriteFile("/proc/sys/net/ipv4/ip_forward", []byte("1\n"), 0644); err != nil {
return fail(fmt.Errorf("enable IPv4 forwarding: %w", err))
}
if err := os.WriteFile("/proc/sys/net/ipv6/conf/all/forwarding", []byte("1\n"), 0644); err != nil {
return fail(fmt.Errorf("enable IPv6 forwarding: %w", err))
}
if autoNAT {
if _, err := exec.LookPath("iptables"); err != nil {
return fail(errors.New("iptables not found; install iptables or start with --auto-nat=false and configure NAT yourself"))
}
ensureRule(debug, "iptables", []string{"-t", "nat", "-C", "POSTROUTING", "-s", "10.123.0.0/16", "-j", "MASQUERADE"}, []string{"-t", "nat", "-A", "POSTROUTING", "-s", "10.123.0.0/16", "-j", "MASQUERADE"})
ensureRule(debug, "iptables", []string{"-C", "FORWARD", "-i", tunName, "-j", "ACCEPT"}, []string{"-A", "FORWARD", "-i", tunName, "-j", "ACCEPT"})
ensureRule(debug, "iptables", []string{"-C", "FORWARD", "-o", tunName, "-m", "conntrack", "--ctstate", "RELATED,ESTABLISHED", "-j", "ACCEPT"}, []string{"-A", "FORWARD", "-o", tunName, "-m", "conntrack", "--ctstate", "RELATED,ESTABLISHED", "-j", "ACCEPT"})
if _, err := exec.LookPath("ip6tables"); err == nil {
ensureRule(debug, "ip6tables", []string{"-t", "nat", "-C", "POSTROUTING", "-s", "fd7a:4472:6167:6f6e::/64", "-j", "MASQUERADE"}, []string{"-t", "nat", "-A", "POSTROUTING", "-s", "fd7a:4472:6167:6f6e::/64", "-j", "MASQUERADE"})
ensureRule(debug, "ip6tables", []string{"-C", "FORWARD", "-i", tunName, "-j", "ACCEPT"}, []string{"-A", "FORWARD", "-i", tunName, "-j", "ACCEPT"})
ensureRule(debug, "ip6tables", []string{"-C", "FORWARD", "-o", tunName, "-m", "conntrack", "--ctstate", "RELATED,ESTABLISHED", "-j", "ACCEPT"}, []string{"-A", "FORWARD", "-o", tunName, "-m", "conntrack", "--ctstate", "RELATED,ESTABLISHED", "-j", "ACCEPT"})
} else {
debug.logf("WARNING: ip6tables not found; IPv6 Internet access needs manual routing/NAT")
}
}
return tun, nil
}
func main() {
host := flag.String("host", "0.0.0.0", "listen host")
port := flag.Int("port", 53, "listen TCP port")
token := flag.String("token", "change-this-token", "shared token")
maxConnections := flag.Int("max-connections", 20000, "maximum simultaneous TCP/53 connections")
maxChunk := flag.Int("chunk-max", 65535, "maximum VPN fragment payload bytes (32-65535)")
maxPackets := flag.Int("vpn-buffered-packets", 2048, "maximum queued return IP packets per client")
pollWait := flag.Duration("poll-wait", 100*time.Millisecond, "long-poll wait for a return packet")
sessionTimeout := flag.Duration("session-timeout", 5*time.Minute, "idle VPN session timeout")
tunName := flag.String("tun", "dragontcp0", "Linux TUN interface name")
mtu := flag.Int("mtu", 1280, "server TUN MTU")
autoNAT := flag.Bool("auto-nat", true, "configure IPv4/IPv6 forwarding and iptables MASQUERADE")
allowPrivate := flag.Bool("allow-private", false, "allow VPN clients to access private/link-local destinations")
mockEcho := flag.Bool("mock-echo", false, "test mode: echo client IP packets back instead of using Linux TUN/NAT")
debugOn := flag.Bool("debug", false, "debug sessions and statistics")
debugPackets := flag.Bool("debug-packets", false, "very verbose per-IP-packet logging")
statsEvery := flag.Duration("debug-stats-interval", 10*time.Second, "debug statistics interval; 0 disables")
flag.Parse()
if *maxChunk < 32 || *maxChunk > protocol.VPNMaxFragment {
fmt.Fprintf(os.Stderr, "--chunk-max must be 32-%d\n", protocol.VPNMaxFragment)
os.Exit(2)
}
if *mtu < 576 || *mtu > 9000 {
fmt.Fprintln(os.Stderr, "--mtu must be 576-9000")
os.Exit(2)
}
debug := &debugStats{enabled: *debugOn, packets: *debugPackets, started: time.Now()}
var tun *os.File
var err error
if !*mockEcho {
tun, err = setupLinuxVPN(*tunName, *mtu, *autoNAT, debug)
if err != nil {
fmt.Fprintln(os.Stderr, "VPN setup failed:", err)
os.Exit(1)
}
defer tun.Close()
}
manager := newVPNManager(tun, *mockEcho, *maxChunk, *maxPackets, *pollWait, *sessionTimeout, *allowPrivate, debug)
addr := net.JoinHostPort(*host, strconv.Itoa(*port))
ln, err := net.Listen("tcp", addr)
if err != nil {
fmt.Fprintln(os.Stderr, err)
os.Exit(1)
}
defer ln.Close()
fmt.Printf("DragonTCP VPN server listening on %s\n", addr)
if *mockEcho {
fmt.Println("mode=mock-echo (no Internet forwarding)")
} else {
fmt.Printf("tun=%s mtu=%d IPv4=10.123.0.1/16 IPv6=fd7a:4472:6167:6f6e::1/64 auto_nat=%t\n", *tunName, *mtu, *autoNAT)
}
fmt.Printf("chunk_max=%d poll_wait=%s buffered_packets=%d\n", *maxChunk, pollWait.String(), *maxPackets)
if debug.enabled && *statsEvery > 0 {
go func() {
t := time.NewTicker(*statsEvery)
defer t.Stop()
for range t.C {
fmt.Printf("[DEBUG] STATS uptime=%s conns=%d sessions=%d up_packets=%d down_packets=%d up_bytes=%d down_bytes=%d dropped=%d errors=%d\n", time.Since(debug.started).Round(time.Second), debug.activeConns.Load(), debug.activeSessions.Load(), debug.upPackets.Load(), debug.downPackets.Load(), debug.upBytes.Load(), debug.downBytes.Load(), debug.dropped.Load(), debug.errors.Load())
}
}()
}
sig := make(chan os.Signal, 1)
signal.Notify(sig, syscall.SIGINT, syscall.SIGTERM)
go func() { <-sig; fmt.Println("Stopping DragonTCP VPN server..."); ln.Close() }()
slots := make(chan struct{}, *maxConnections)
for {
conn, err := ln.Accept()
if err != nil {
break
}
select {
case slots <- struct{}{}:
debug.activeConns.Add(1)
go handleConn(conn, *token, manager, slots, debug)
default:
_ = conn.Close()
}
}
}
+3
View File
@@ -0,0 +1,3 @@
module dragontcpvpn
go 1.22
+211
View File
@@ -0,0 +1,211 @@
package protocol
import (
"encoding/binary"
"errors"
"fmt"
"io"
"net"
"sync"
"time"
)
const (
XORKey byte = 0xAD
// MaxChunkPayload is the hard application-record payload ceiling.
// The adaptive chunk protocol may use any size from 32 bytes through 1 MiB.
MaxChunkPayload = 1024 * 1024
// Framed CPUSH/DATA messages include text metadata in addition to chunk
// bytes, so keep the frame ceiling comfortably above MaxChunkPayload.
MaxHandshake = 2 * 1024 * 1024
)
// 64 KiB balances throughput with memory use at high connection counts.
var BufferPool = sync.Pool{
New: func() any {
b := make([]byte, 64*1024)
return &b
},
}
func ReadRequestFrame(r io.Reader) (uint32, uint32, []byte, error) {
var header [14]byte
if _, err := io.ReadFull(r, header[:]); err != nil {
return 0, 0, nil, err
}
if header[0] != 'U' || header[1] != 'P' {
return 0, 0, nil, errors.New("bad request magic")
}
requestID := binary.BigEndian.Uint32(header[2:6])
reserved := binary.BigEndian.Uint32(header[6:10])
length := binary.BigEndian.Uint32(header[10:14])
if length > MaxHandshake {
return 0, 0, nil, errors.New("handshake payload too large")
}
payload := make([]byte, int(length))
if _, err := io.ReadFull(r, payload); err != nil {
return 0, 0, nil, err
}
XorInPlace(payload)
return requestID, reserved, payload, nil
}
func WriteRequestFrame(w io.Writer, requestID uint32, payload []byte) error {
if len(payload) > MaxHandshake {
return errors.New("request frame payload too large")
}
packet := make([]byte, 14+len(payload))
packet[0], packet[1] = 'U', 'P'
binary.BigEndian.PutUint32(packet[2:6], requestID)
binary.BigEndian.PutUint32(packet[6:10], 0)
binary.BigEndian.PutUint32(packet[10:14], uint32(len(payload)))
copy(packet[14:], payload)
XorInPlace(packet[14:])
return writeAll(w, packet)
}
func ReadResponseFrame(r io.Reader) (uint32, []byte, error) {
var header [10]byte
if _, err := io.ReadFull(r, header[:]); err != nil {
return 0, nil, err
}
if header[0] != 'O' || header[1] != 'K' {
return 0, nil, fmt.Errorf("bad response magic: %q", header[:2])
}
requestID := binary.BigEndian.Uint32(header[2:6])
length := binary.BigEndian.Uint32(header[6:10])
if length > MaxHandshake {
return 0, nil, errors.New("handshake response too large")
}
payload := make([]byte, int(length))
if _, err := io.ReadFull(r, payload); err != nil {
return 0, nil, err
}
XorInPlace(payload)
return requestID, payload, nil
}
func WriteResponseFrame(w io.Writer, requestID uint32, payload []byte) error {
if len(payload) > MaxHandshake {
return errors.New("response frame payload too large")
}
packet := make([]byte, 10+len(payload))
packet[0], packet[1] = 'O', 'K'
binary.BigEndian.PutUint32(packet[2:6], requestID)
binary.BigEndian.PutUint32(packet[6:10], uint32(len(payload)))
copy(packet[10:], payload)
XorInPlace(packet[10:])
return writeAll(w, packet)
}
func writeAll(w io.Writer, b []byte) error {
for len(b) > 0 {
n, err := w.Write(b)
if err != nil {
return err
}
b = b[n:]
}
return nil
}
func CopyXOR(dst net.Conn, src net.Conn) error {
ptr := BufferPool.Get().(*[]byte)
buf := *ptr
defer BufferPool.Put(ptr)
for {
n, err := src.Read(buf)
if n > 0 {
chunk := buf[:n]
XorInPlace(chunk)
if err2 := writeAll(dst, chunk); err2 != nil {
return err2
}
// No restore pass is needed. The next Read overwrites these bytes.
}
if err != nil {
if errors.Is(err, io.EOF) {
return nil
}
return err
}
}
}
func relayPair(a, b net.Conn, copier func(net.Conn, net.Conn) error) {
done := make(chan struct{}, 2)
go func() {
_ = copier(b, a)
if cw, ok := b.(interface{ CloseWrite() error }); ok {
_ = cw.CloseWrite()
}
done <- struct{}{}
}()
go func() {
_ = copier(a, b)
if cw, ok := a.(interface{ CloseWrite() error }); ok {
_ = cw.CloseWrite()
}
done <- struct{}{}
}()
// Preserve normal TCP half-close semantics. The old implementation set a
// 2-second deadline on both connections after the first copy direction
// ended, which truncated slow or large responses. Wait for the remaining
// direction to drain naturally instead.
<-done
<-done
}
func RelayXOR(a, b net.Conn) {
relayPair(a, b, CopyXOR)
}
// RelayRaw allows Go/Linux to use the optimized TCP io.Copy path. On Linux,
// TCP-to-TCP copies can use splice, eliminating the userspace XOR/copy loop.
func RelayRaw(a, b net.Conn) {
relayPair(a, b, func(dst, src net.Conn) error {
_, err := io.Copy(dst, src)
return err
})
}
func TuneTCP(conn net.Conn) {
if tcp, ok := conn.(*net.TCPConn); ok {
_ = tcp.SetNoDelay(true)
_ = tcp.SetKeepAlive(true)
_ = tcp.SetKeepAlivePeriod(30 * time.Second)
}
}
// TuneTCPBuffer optionally requests larger kernel socket buffers. A value <= 0
// leaves Linux/Android autotuning untouched, which is the recommended default
// for large connection counts. For a small number of high-BDP mobile links,
// values such as 1048576 or 4194304 can improve throughput.
func TuneTCPBuffer(conn net.Conn, size int) {
if size <= 0 {
return
}
if tcp, ok := conn.(*net.TCPConn); ok {
_ = tcp.SetReadBuffer(size)
_ = tcp.SetWriteBuffer(size)
}
}
+267
View File
@@ -0,0 +1,267 @@
package protocol
import (
"encoding/binary"
"errors"
"fmt"
"net/netip"
)
const (
VPNCmdOpen byte = 0x30
VPNCmdPush byte = 0x31
VPNCmdPull byte = 0x32
VPNCmdClose byte = 0x33
VPNRespOpened byte = 0x40
VPNRespAck byte = 0x41
VPNRespData byte = 0x42
VPNRespWait byte = 0x43
VPNRespClosed byte = 0x44
VPNRespError byte = 0x7f
VPNNoAck uint32 = 0xffffffff
VPNMaxFragment = 65535
)
type VPNSessionID [16]byte
func VPNError(message string) []byte {
b := []byte(message)
if len(b) > 4096 {
b = b[:4096]
}
out := make([]byte, 1+len(b))
out[0] = VPNRespError
copy(out[1:], b)
return out
}
func ParseVPNError(payload []byte) error {
if len(payload) == 0 {
return errors.New("empty DragonTCP VPN response")
}
if payload[0] == VPNRespError {
return errors.New(string(payload[1:]))
}
return nil
}
// OPEN request:
// cmd(1) sid(16) tokenLen(2) token(N) ipv4(4) ipv6(16) mtu(2)
func BuildVPNOpen(sid VPNSessionID, token string, ipv4, ipv6 netip.Addr, mtu int) ([]byte, error) {
if len(token) > 4096 {
return nil, errors.New("token too long")
}
if !ipv4.Is4() || !ipv6.Is6() {
return nil, errors.New("invalid VPN client addresses")
}
if mtu < 576 || mtu > 65535 {
return nil, errors.New("invalid VPN MTU")
}
out := make([]byte, 1+16+2+len(token)+4+16+2)
out[0] = VPNCmdOpen
copy(out[1:17], sid[:])
binary.BigEndian.PutUint16(out[17:19], uint16(len(token)))
pos := 19
copy(out[pos:pos+len(token)], token)
pos += len(token)
v4 := ipv4.As4()
copy(out[pos:pos+4], v4[:])
pos += 4
v6 := ipv6.As16()
copy(out[pos:pos+16], v6[:])
pos += 16
binary.BigEndian.PutUint16(out[pos:pos+2], uint16(mtu))
return out, nil
}
func ParseVPNOpen(payload []byte) (sid VPNSessionID, token string, ipv4, ipv6 netip.Addr, mtu int, err error) {
if len(payload) < 1+16+2+4+16+2 || payload[0] != VPNCmdOpen {
err = errors.New("bad VPN OPEN")
return
}
copy(sid[:], payload[1:17])
tokenLen := int(binary.BigEndian.Uint16(payload[17:19]))
need := 1 + 16 + 2 + tokenLen + 4 + 16 + 2
if tokenLen < 0 || len(payload) != need {
err = errors.New("bad VPN OPEN length")
return
}
pos := 19
token = string(payload[pos : pos+tokenLen])
pos += tokenLen
var a4 [4]byte
copy(a4[:], payload[pos:pos+4])
ipv4 = netip.AddrFrom4(a4)
pos += 4
var a6 [16]byte
copy(a6[:], payload[pos:pos+16])
ipv6 = netip.AddrFrom16(a6)
pos += 16
mtu = int(binary.BigEndian.Uint16(payload[pos : pos+2]))
return
}
func BuildVPNOpened(maxChunk int) []byte {
if maxChunk > VPNMaxFragment {
maxChunk = VPNMaxFragment
}
if maxChunk < 1 {
maxChunk = 1
}
out := make([]byte, 3)
out[0] = VPNRespOpened
binary.BigEndian.PutUint16(out[1:3], uint16(maxChunk))
return out
}
func ParseVPNOpened(payload []byte) (int, error) {
if err := ParseVPNError(payload); err != nil {
return 0, err
}
if len(payload) != 3 || payload[0] != VPNRespOpened {
return 0, errors.New("bad VPN OPENED response")
}
return int(binary.BigEndian.Uint16(payload[1:3])), nil
}
// PUSH request: cmd(1) sid(16) seq(4) offset(2) total(2) data(N)
func BuildVPNPush(sid VPNSessionID, seq uint32, offset, total int, data []byte) ([]byte, error) {
if total < 1 || total > 65535 || offset < 0 || offset > total || len(data) < 1 || offset+len(data) > total || len(data) > VPNMaxFragment {
return nil, errors.New("invalid VPN PUSH fragment")
}
out := make([]byte, 25+len(data))
out[0] = VPNCmdPush
copy(out[1:17], sid[:])
binary.BigEndian.PutUint32(out[17:21], seq)
binary.BigEndian.PutUint16(out[21:23], uint16(offset))
binary.BigEndian.PutUint16(out[23:25], uint16(total))
copy(out[25:], data)
return out, nil
}
func ParseVPNPush(payload []byte) (sid VPNSessionID, seq uint32, offset, total int, data []byte, err error) {
if len(payload) < 26 || payload[0] != VPNCmdPush {
err = errors.New("bad VPN PUSH")
return
}
copy(sid[:], payload[1:17])
seq = binary.BigEndian.Uint32(payload[17:21])
offset = int(binary.BigEndian.Uint16(payload[21:23]))
total = int(binary.BigEndian.Uint16(payload[23:25]))
data = payload[25:]
if total < 1 || offset < 0 || offset > total || len(data) < 1 || offset+len(data) > total {
err = errors.New("bad VPN PUSH fragment bounds")
}
return
}
func BuildVPNAck(seq uint32, accepted int) []byte {
out := make([]byte, 7)
out[0] = VPNRespAck
binary.BigEndian.PutUint32(out[1:5], seq)
binary.BigEndian.PutUint16(out[5:7], uint16(accepted))
return out
}
func ParseVPNAck(payload []byte) (seq uint32, accepted int, err error) {
if e := ParseVPNError(payload); e != nil {
err = e
return
}
if len(payload) != 7 || payload[0] != VPNRespAck {
err = errors.New("bad VPN ACK")
return
}
seq = binary.BigEndian.Uint32(payload[1:5])
accepted = int(binary.BigEndian.Uint16(payload[5:7]))
return
}
// PULL request: cmd(1) sid(16) ack(4) want(4) offset(2) limit(2)
func BuildVPNPull(sid VPNSessionID, ack, want uint32, offset, limit int) ([]byte, error) {
if offset < 0 || offset > 65535 || limit < 1 || limit > VPNMaxFragment {
return nil, errors.New("invalid VPN PULL")
}
out := make([]byte, 29)
out[0] = VPNCmdPull
copy(out[1:17], sid[:])
binary.BigEndian.PutUint32(out[17:21], ack)
binary.BigEndian.PutUint32(out[21:25], want)
binary.BigEndian.PutUint16(out[25:27], uint16(offset))
binary.BigEndian.PutUint16(out[27:29], uint16(limit))
return out, nil
}
func ParseVPNPull(payload []byte) (sid VPNSessionID, ack, want uint32, offset, limit int, err error) {
if len(payload) != 29 || payload[0] != VPNCmdPull {
err = errors.New("bad VPN PULL")
return
}
copy(sid[:], payload[1:17])
ack = binary.BigEndian.Uint32(payload[17:21])
want = binary.BigEndian.Uint32(payload[21:25])
offset = int(binary.BigEndian.Uint16(payload[25:27]))
limit = int(binary.BigEndian.Uint16(payload[27:29]))
if limit < 1 {
err = errors.New("bad VPN PULL limit")
}
return
}
// DATA response: cmd(1) seq(4) offset(2) total(2) data(N)
func BuildVPNData(seq uint32, offset, total int, data []byte) []byte {
out := make([]byte, 9+len(data))
out[0] = VPNRespData
binary.BigEndian.PutUint32(out[1:5], seq)
binary.BigEndian.PutUint16(out[5:7], uint16(offset))
binary.BigEndian.PutUint16(out[7:9], uint16(total))
copy(out[9:], data)
return out
}
func ParseVPNData(payload []byte) (seq uint32, offset, total int, data []byte, wait bool, err error) {
if e := ParseVPNError(payload); e != nil {
err = e
return
}
if len(payload) == 1 && payload[0] == VPNRespWait {
wait = true
return
}
if len(payload) < 10 || payload[0] != VPNRespData {
err = fmt.Errorf("bad VPN DATA response type/length")
return
}
seq = binary.BigEndian.Uint32(payload[1:5])
offset = int(binary.BigEndian.Uint16(payload[5:7]))
total = int(binary.BigEndian.Uint16(payload[7:9]))
data = payload[9:]
if total < 1 || offset < 0 || offset+len(data) > total || len(data) < 1 {
err = errors.New("bad VPN DATA bounds")
}
return
}
func BuildVPNClose(sid VPNSessionID) []byte {
out := make([]byte, 17)
out[0] = VPNCmdClose
copy(out[1:17], sid[:])
return out
}
func ParseVPNClose(payload []byte) (sid VPNSessionID, err error) {
if len(payload) != 17 || payload[0] != VPNCmdClose {
return sid, errors.New("bad VPN CLOSE")
}
copy(sid[:], payload[1:17])
return sid, nil
}
func IsVPNCommand(payload []byte) bool {
if len(payload) == 0 {
return false
}
return payload[0] >= VPNCmdOpen && payload[0] <= VPNCmdClose
}
+44
View File
@@ -0,0 +1,44 @@
//go:build arm || 386
package protocol
import "unsafe"
const xorWordMask32 uint32 = 0xADADADAD
// XorInPlace is the 32-bit optimized path used by ARMv7/386 builds.
// It aligns once, then processes 32 bytes per iteration with native uint32
// operations instead of a byte-at-a-time loop.
func XorInPlace(b []byte) {
n := len(b)
if n == 0 {
return
}
i := 0
for i < n && (uintptr(unsafe.Pointer(&b[i]))&3) != 0 {
b[i] ^= XORKey
i++
}
for ; i+32 <= n; i += 32 {
p := unsafe.Pointer(&b[i])
*(*uint32)(unsafe.Add(p, 0)) ^= xorWordMask32
*(*uint32)(unsafe.Add(p, 4)) ^= xorWordMask32
*(*uint32)(unsafe.Add(p, 8)) ^= xorWordMask32
*(*uint32)(unsafe.Add(p, 12)) ^= xorWordMask32
*(*uint32)(unsafe.Add(p, 16)) ^= xorWordMask32
*(*uint32)(unsafe.Add(p, 20)) ^= xorWordMask32
*(*uint32)(unsafe.Add(p, 24)) ^= xorWordMask32
*(*uint32)(unsafe.Add(p, 28)) ^= xorWordMask32
}
for ; i+4 <= n; i += 4 {
p := (*uint32)(unsafe.Pointer(&b[i]))
*p ^= xorWordMask32
}
for ; i < n; i++ {
b[i] ^= XORKey
}
}
+50
View File
@@ -0,0 +1,50 @@
//go:build amd64 || arm64
package protocol
import "unsafe"
const xorWordMask uint64 = 0xADADADADADADADAD
// XorInPlace is optimized for 64-bit targets (amd64/arm64).
//
// It aligns the input once, then XORs 64 bytes per loop iteration using
// eight native 64-bit operations. This removes the encoding/binary call
// overhead from the hot relay path and lets the compiler generate a tight
// load/xor/store loop.
func XorInPlace(b []byte) {
n := len(b)
if n == 0 {
return
}
i := 0
// Align the pointer for native uint64 accesses. This is normally already
// aligned for pooled relay buffers, but also makes this safe for subslices.
for i < n && (uintptr(unsafe.Pointer(&b[i]))&7) != 0 {
b[i] ^= XORKey
i++
}
for ; i+64 <= n; i += 64 {
p := unsafe.Pointer(&b[i])
*(*uint64)(unsafe.Add(p, 0)) ^= xorWordMask
*(*uint64)(unsafe.Add(p, 8)) ^= xorWordMask
*(*uint64)(unsafe.Add(p, 16)) ^= xorWordMask
*(*uint64)(unsafe.Add(p, 24)) ^= xorWordMask
*(*uint64)(unsafe.Add(p, 32)) ^= xorWordMask
*(*uint64)(unsafe.Add(p, 40)) ^= xorWordMask
*(*uint64)(unsafe.Add(p, 48)) ^= xorWordMask
*(*uint64)(unsafe.Add(p, 56)) ^= xorWordMask
}
for ; i+8 <= n; i += 8 {
p := (*uint64)(unsafe.Pointer(&b[i]))
*p ^= xorWordMask
}
for ; i < n; i++ {
b[i] ^= XORKey
}
}
+10
View File
@@ -0,0 +1,10 @@
//go:build !amd64 && !arm64 && !arm && !386
package protocol
// Generic fallback for 32-bit and uncommon architectures.
func XorInPlace(b []byte) {
for i := range b {
b[i] ^= XORKey
}
}