package main import ( "bufio" "encoding/binary" "errors" "fmt" "io" "log" "net" "strconv" "strings" "sync" "time" "unsafe" "golang.org/x/sys/unix" ) // dialFunc is the DNS-aware upstream dialer (same as MITM transport DialContext). type dialFunc func(network, addr string) (net.Conn, error) // serveTransparentAccept accepts iptables-REDIRECTED TCP connections, recovers // the original destination / SNI, then feeds them into the local HTTP MITM via // a synthetic CONNECT so UK-VPN routing stays intact (no Wi‑Fi http_proxy). func serveTransparentAccept(ln net.Listener, mitmAddr string, dial dialFunc) { log.Printf("transparent accept on %s → MITM %s", ln.Addr(), mitmAddr) for { c, err := ln.Accept() if err != nil { if errors.Is(err, net.ErrClosed) { return } log.Printf("transparent accept: %v", err) continue } go handleTransparent(c, mitmAddr, dial) } } func handleTransparent(client net.Conn, mitmAddr string, dial dialFunc) { defer client.Close() _ = client.SetDeadline(time.Now().Add(30 * time.Second)) bc := newBufConn(client) sni, helloErr := peekSNI(bc) dstHost, dstPort, dstErr := originalDst(client) host := strings.TrimSpace(sni) port := 443 if dstErr == nil && dstPort > 0 { port = dstPort } if host == "" { if dstErr != nil { log.Printf("transparent: no SNI (%v) and no original dst (%v)", helloErr, dstErr) return } host = dstHost } target := net.JoinHostPort(host, strconv.Itoa(port)) // CDN / media / mediaselector: splice without MITM so playback works. if shouldPassthrough(host) { _ = client.SetDeadline(time.Time{}) // Re-feed ClientHello: peekSNI left bytes in bc; rebuild a reader. left := bc.buffered() var clientR net.Conn = client if len(left) > 0 { clientR = &prefixConn{Conn: client, prefix: left} } upDial := dial if upDial == nil { upDial = func(network, addr string) (net.Conn, error) { return net.DialTimeout(network, addr, 15*time.Second) } } // Prefer the app's original destination IP (same CDN edge / geo) over re-resolve. passHost := host if dstErr == nil && net.ParseIP(dstHost) != nil && !hostIsLoopback(dstHost) { passHost = dstHost log.Printf("[PASS] %s → %s:%d (orig-dst)", host, dstHost, port) } handlePassthrough(clientR, passHost, port, upDial) return } log.Printf("[TPROXY] %s → CONNECT %s", client.RemoteAddr(), target) mitm, err := net.DialTimeout("tcp", mitmAddr, 5*time.Second) if err != nil { log.Printf("transparent dial mitm: %v", err) return } defer mitm.Close() // Pass SO_ORIGINAL_DST so MITM upstream dials the app's IP (NordVPN fake-IP // / same CDN edge). Re-resolving via public DNS breaks BBC geolocation. var b strings.Builder fmt.Fprintf(&b, "CONNECT %s HTTP/1.1\r\nHost: %s\r\n", target, target) if dstErr == nil && net.ParseIP(dstHost) != nil && !hostIsLoopback(dstHost) { fmt.Fprintf(&b, "X-Appproxy-Orig-Dst: %s\r\n", net.JoinHostPort(dstHost, strconv.Itoa(port))) } b.WriteString("\r\n") if _, err := io.WriteString(mitm, b.String()); err != nil { log.Printf("transparent CONNECT write: %v", err) return } br := bufio.NewReader(mitm) status, err := br.ReadString('\n') if err != nil { log.Printf("transparent CONNECT read: %v", err) return } if !strings.Contains(status, "200") { rest, _ := io.ReadAll(io.LimitReader(br, 512)) log.Printf("transparent CONNECT rejected: %s%s", status, rest) return } // Drain remaining response headers. for { line, err := br.ReadString('\n') if err != nil || line == "\r\n" || line == "\n" { break } } _ = client.SetDeadline(time.Time{}) _ = mitm.SetDeadline(time.Time{}) // Any buffered ClientHello bytes must go to the MITM first. var once sync.Once left := bc.buffered() if len(left) > 0 { if _, err := mitm.Write(left); err != nil { log.Printf("transparent hello write: %v", err) return } } errc := make(chan error, 2) go func() { _, err := io.Copy(mitm, bc) errc <- err once.Do(func() { if tc, ok := mitm.(*net.TCPConn); ok { _ = tc.CloseWrite() } }) }() go func() { _, err := io.Copy(client, br) errc <- err if tc, ok := client.(*net.TCPConn); ok { _ = tc.CloseWrite() } }() <-errc } // prefixConn emits prefix once, then reads from Conn. type prefixConn struct { net.Conn prefix []byte i int } func (p *prefixConn) Read(b []byte) (int, error) { if p.i < len(p.prefix) { n := copy(b, p.prefix[p.i:]) p.i += n return n, nil } return p.Conn.Read(b) } type bufConn struct { net.Conn r *bufio.Reader } func newBufConn(c net.Conn) *bufConn { return &bufConn{Conn: c, r: bufio.NewReaderSize(c, 4096)} } func (b *bufConn) Read(p []byte) (int, error) { return b.r.Read(p) } func (b *bufConn) buffered() []byte { n := b.r.Buffered() if n == 0 { return nil } buf, _ := b.r.Peek(n) out := make([]byte, len(buf)) copy(out, buf) // Consume so later Read does not duplicate. _, _ = b.r.Discard(n) return out } // peekSNI reads a TLS ClientHello (via Peek) and returns the SNI hostname. func peekSNI(bc *bufConn) (string, error) { hdr, err := bc.r.Peek(5) if err != nil { return "", err } if hdr[0] != 0x16 { // handshake return "", fmt.Errorf("not TLS handshake (type=%d)", hdr[0]) } recLen := int(hdr[3])<<8 | int(hdr[4]) need := 5 + recLen if need > 16*1024 { return "", fmt.Errorf("ClientHello too large (%d)", need) } // Wait until full record is buffered. deadline := time.Now().Add(5 * time.Second) for bc.r.Buffered() < need && time.Now().Before(deadline) { _ = bc.Conn.SetReadDeadline(time.Now().Add(200 * time.Millisecond)) _, _ = bc.r.Peek(need) } _ = bc.Conn.SetReadDeadline(time.Time{}) if bc.r.Buffered() < need { need = bc.r.Buffered() } data, err := bc.r.Peek(need) if err != nil { return "", err } return parseClientHelloSNI(data) } func parseClientHelloSNI(rec []byte) (string, error) { if len(rec) < 5+4 { return "", fmt.Errorf("short record") } if rec[0] != 0x16 || rec[5] != 0x01 { // handshake + client_hello return "", fmt.Errorf("not ClientHello") } // Skip: record hdr(5) + hs type(1) + hs len(3) + client version(2) + random(32) i := 5 + 1 + 3 + 2 + 32 if i >= len(rec) { return "", fmt.Errorf("truncated ClientHello") } // session id sidLen := int(rec[i]) i += 1 + sidLen if i+2 > len(rec) { return "", fmt.Errorf("truncate cipher suites") } csLen := int(rec[i])<<8 | int(rec[i+1]) i += 2 + csLen if i+1 > len(rec) { return "", fmt.Errorf("truncate compression") } compLen := int(rec[i]) i += 1 + compLen if i+2 > len(rec) { return "", nil // no extensions } extLen := int(rec[i])<<8 | int(rec[i+1]) i += 2 end := i + extLen if end > len(rec) { end = len(rec) } for i+4 <= end { typ := int(rec[i])<<8 | int(rec[i+1]) l := int(rec[i+2])<<8 | int(rec[i+3]) i += 4 if i+l > end { break } if typ == 0 { // server_name return parseSNIExtension(rec[i : i+l]) } i += l } return "", nil } func parseSNIExtension(b []byte) (string, error) { if len(b) < 2 { return "", nil } listLen := int(b[0])<<8 | int(b[1]) i := 2 end := 2 + listLen if end > len(b) { end = len(b) } for i+3 <= end { nameType := b[i] nameLen := int(b[i+1])<<8 | int(b[i+2]) i += 3 if i+nameLen > end { break } if nameType == 0 { return string(b[i : i+nameLen]), nil } i += nameLen } return "", nil } func originalDst(conn net.Conn) (host string, port int, err error) { tcp, ok := conn.(*net.TCPConn) if !ok { return "", 0, fmt.Errorf("not TCP") } rc, err := tcp.SyscallConn() if err != nil { return "", 0, err } var ( ip net.IP prt int err2 error ) cerr := rc.Control(func(fd uintptr) { ip, prt, err2 = getOrigDstFD(int(fd)) }) if cerr != nil { return "", 0, cerr } if err2 != nil { return "", 0, err2 } return ip.String(), prt, nil } func getOrigDstFD(fd int) (net.IP, int, error) { // Try IPv6 structure first (works for IPv4-mapped too on many kernels). if ip, port, err := getOrigDstIPv6(fd); err == nil { return ip, port, nil } return getOrigDstIPv4(fd) } func getOrigDstIPv4(fd int) (net.IP, int, error) { const soOriginalDst = 80 var addr unix.RawSockaddrInet4 sz := uint32(unsafe.Sizeof(addr)) _, _, errno := unix.Syscall6( unix.SYS_GETSOCKOPT, uintptr(fd), uintptr(unix.IPPROTO_IP), uintptr(soOriginalDst), uintptr(unsafe.Pointer(&addr)), uintptr(unsafe.Pointer(&sz)), 0, ) if errno != 0 { return nil, 0, errno } ip := net.IPv4(addr.Addr[0], addr.Addr[1], addr.Addr[2], addr.Addr[3]) port := int(binary.BigEndian.Uint16((*[2]byte)(unsafe.Pointer(&addr.Port))[:])) return ip, port, nil } func getOrigDstIPv6(fd int) (net.IP, int, error) { const soOriginalDst = 80 var addr unix.RawSockaddrInet6 sz := uint32(unsafe.Sizeof(addr)) _, _, errno := unix.Syscall6( unix.SYS_GETSOCKOPT, uintptr(fd), uintptr(unix.IPPROTO_IPV6), uintptr(soOriginalDst), uintptr(unsafe.Pointer(&addr)), uintptr(unsafe.Pointer(&sz)), 0, ) if errno != 0 { return nil, 0, errno } ip := make(net.IP, 16) copy(ip, addr.Addr[:]) port := int(binary.BigEndian.Uint16((*[2]byte)(unsafe.Pointer(&addr.Port))[:])) return ip, port, nil }