From ca72bd8d4558870f8c24615678b181e52ebcea4f Mon Sep 17 00:00:00 2001 From: Kacy Fortner Date: Wed, 10 Jun 2026 03:50:00 +0000 Subject: [PATCH] =?UTF-8?q?feat:=20server-side=20mtls=20=E2=80=94=20certif?= =?UTF-8?q?icaterequest=20+=20client-cert=20verify?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit extracts the handshake portion of handleTlsSession into a new acceptServerHandshake() that returns a ServerSession. handleTlsSession keeps doing exactly what it did before — calls accept(.., null, ..) then enters the forwarding loop with the returned session. new MtlsOpts threads client-cert behavior through: - when non-null, the server emits CertificateRequest between EncryptedExtensions and Certificate - after the server's Finished, expects Certificate (possibly empty), optional CertificateVerify, then Finished - verifies the cert chain against trust_ca_pem, then verifies the CertVerify signature against the cert's pubkey using the .client context string - require_client_cert=true rejects an empty Certificate with MissingClientCert; require_client_cert=false accepts (warn mode) - returns the SAN URI (or subject CN) as session.peer_identity parseCertificateMessage now accepts an empty cert list per RFC 8446 §4.4.2 so the server can detect the empty-cert case. three e2e tests on a socketpair drive the new path through the real client_session.doHandshake: 1. valid client cert → handshake completes, identity surfaces 2. require + empty cert → server returns MissingClientCert 3. warn mode + empty cert → handshake completes, no identity --- src/tls/handshake/message_parse.zig | 11 +- src/tls/proxy/session_runtime.zig | 607 +++++++++++++++++++++++++--- 2 files changed, 567 insertions(+), 51 deletions(-) diff --git a/src/tls/handshake/message_parse.zig b/src/tls/handshake/message_parse.zig index 529cdc3..bde8c56 100644 --- a/src/tls/handshake/message_parse.zig +++ b/src/tls/handshake/message_parse.zig @@ -189,10 +189,12 @@ pub fn parseEncryptedExtensions(body: []const u8) ParseError!EncryptedExtensions } /// parse a TLS 1.3 Certificate message body. returns the leaf cert DER -/// (the first cert in the chain). PR 4 only supports single-cert chains — -/// matches what x509_gen issues and what the existing server emits. +/// (the first cert in the chain), or an empty slice when the message +/// carries no certs at all (TLS 1.3 client-auth: the client sends an +/// empty Certificate to mean "I have nothing to authenticate with"). +/// PR 4 only supports single-cert chains. pub fn parseCertificateMessage(body: []const u8) ParseError![]const u8 { - if (body.len < 1 + 3 + 3 + 2) return ParseError.Truncated; + if (body.len < 1 + 3) return ParseError.Truncated; var pos: usize = 0; const ctx_len = body[pos]; pos += 1; @@ -205,6 +207,9 @@ pub fn parseCertificateMessage(body: []const u8) ParseError![]const u8 { if (pos + list_len > body.len) return ParseError.Truncated; const list_end = pos + list_len; + // empty cert list: peer has nothing to present. valid per RFC 8446 §4.4.2. + if (list_len == 0) return body[pos..pos]; + if (pos + 3 > list_end) return ParseError.Truncated; const cert_len = readU24(body[pos..]); pos += 3; diff --git a/src/tls/proxy/session_runtime.zig b/src/tls/proxy/session_runtime.zig index 37d8391..1460946 100644 --- a/src/tls/proxy/session_runtime.zig +++ b/src/tls/proxy/session_runtime.zig @@ -5,24 +5,67 @@ const log = @import("../../lib/log.zig"); const http2_request = @import("../../network/proxy/http2_request.zig"); const backend_mod = @import("../backend.zig"); const handshake = @import("../handshake.zig"); +const message_build = @import("../handshake/message_build.zig"); +const message_parse = @import("../handshake/message_parse.zig"); const pem = @import("../pem.zig"); const record = @import("../record.zig"); const socket_support = @import("socket_support.zig"); +const x509_verify = @import("../x509_verify.zig"); const X25519 = std.crypto.dh.X25519; const Sha384 = std.crypto.hash.sha2.Sha384; +const EcdsaP256 = std.crypto.sign.ecdsa.EcdsaP256Sha256; const hash_len = Sha384.digest_length; -pub fn handleTlsSession( +/// optional mTLS configuration for the server-side handshake. when set, +/// the server emits a CertificateRequest, requires (or merely inspects) +/// a client cert in response, and verifies it against `trust_ca_pem`. +pub const MtlsOpts = struct { + /// when true an empty/absent client Certificate fails the handshake. + /// when false the verified peer identity is `null` after handshake + /// and the caller may decide what to do (PR 5 wires the "warn" mode). + require_client_cert: bool, + /// PEM bytes of the cluster CA. + trust_ca_pem: []const u8, + /// optional SAN URI to require on the client cert. + expected_identity: ?[]const u8 = null, + /// current unix seconds (test-injectable; production passes wall-clock). + now_unix: i64, +}; + +/// what the handshake produced. owned by the caller; `deinit` frees the +/// optional peer-identity buffer. +pub const ServerSession = struct { + selected_alpn: ?[]const u8, + app_keys: handshake.ApplicationKeys, + /// duplicate of the verified client SAN URI when mTLS succeeded with + /// a valid client cert; null otherwise. allocator-owned. + peer_identity: ?[]u8 = null, + + pub fn deinit(self: *ServerSession, alloc: std.mem.Allocator) void { + if (self.peer_identity) |p| alloc.free(p); + } +}; + +/// run the TLS 1.3 server handshake on `client_fd`. when `mtls_opts` is +/// non-null, the server emits CertificateRequest and verifies the client's +/// cert before sending its own Finished response — note that's a +/// per-side ordering quirk: TLS 1.3 has the server send all its +/// handshake messages and Finished first, then expects the client to +/// send Certificate + CertificateVerify + Finished. so client-cert +/// verify happens *after* the server's Finished, just like in regular +/// TLS 1.3 client auth. +pub fn acceptServerHandshake( io: std.Io, + alloc: std.mem.Allocator, client_fd: posix.fd_t, client_hello: []const u8, cert_pem: []u8, key_pem: []u8, - backend_info: backend_mod.Backend, + mtls_opts: ?MtlsOpts, handshake_complete: *bool, -) !void { +) !ServerSession { if (client_hello.len < 9) return error.InvalidClientHello; const rec_len = (@as(usize, client_hello[3]) << 8) | @as(usize, client_hello[4]); if (client_hello.len < 5 + rec_len) return error.InvalidClientHello; @@ -87,8 +130,15 @@ pub fn handleTlsSession( transcript.update(ee_buf[0..ee_len]); try sendEncryptedHandshake(client_fd, ee_buf[0..ee_len], server_hs_traffic, &server_seq); - const cert_der = pem.parseCertDer(std.heap.page_allocator, cert_pem) catch return error.CertParseFailed; - defer std.heap.page_allocator.free(cert_der); + if (mtls_opts != null) { + var cr_buf: [64]u8 = undefined; + const cr_len = message_build.buildCertificateRequest(&cr_buf) catch return error.HandshakeFailed; + transcript.update(cr_buf[0..cr_len]); + try sendEncryptedHandshake(client_fd, cr_buf[0..cr_len], server_hs_traffic, &server_seq); + } + + const cert_der = pem.parseCertDer(alloc, cert_pem) catch return error.CertParseFailed; + defer alloc.free(cert_der); var cert_buf: [8192]u8 = undefined; const cert_len = handshake.buildCertificate(&cert_buf, cert_der) catch return error.HandshakeFailed; @@ -113,51 +163,19 @@ pub fn handleTlsSession( try sendEncryptedHandshake(client_fd, fin_buf[0..fin_len], server_hs_traffic, &server_seq); var client_seq: u64 = 0; - var client_finished_buf: [512]u8 = undefined; - const client_rec_n = socket_support.readWithTimeout(client_fd, &client_finished_buf, 10000) catch - return error.ReadFailed; - if (client_rec_n < record.record_header_size + record.aead_tag_size + 1) - return error.InvalidClientFinished; - - var client_data = client_finished_buf[0..client_rec_n]; - if (client_data[0] == 0x14) { - if (client_data.len < 6) return error.InvalidClientFinished; - const ccs_len: usize = 5 + @as(usize, (@as(u16, client_data[3]) << 8) | @as(u16, client_data[4])); - if (ccs_len > client_data.len) return error.InvalidClientFinished; - client_data = client_data[ccs_len..]; - if (client_data.len < record.record_header_size + record.aead_tag_size + 1) - return error.InvalidClientFinished; + var peer_identity_out: ?[]u8 = null; + errdefer if (peer_identity_out) |p| alloc.free(p); + + if (mtls_opts) |opts| { + // mTLS path: client may send Certificate + CertificateVerify + Finished, + // possibly across several records. read until we have a Finished. + try acceptClientAuthAndFinished(alloc, client_fd, client_hs_traffic, &client_seq, &transcript, hs_keys, opts, &peer_identity_out); + } else { + // non-mTLS path: single read, expects just a Finished (preceded by + // an optional ChangeCipherSpec record from middlebox-friendly clients). + try readPlainClientFinished(client_fd, client_hs_traffic, &client_seq, &transcript, hs_keys); } - const client_rec_header = client_data[0..5].*; - const client_ciphertext_len = (@as(usize, client_data[3]) << 8) | @as(usize, client_data[4]); - if (5 + client_ciphertext_len > client_data.len) return error.InvalidClientFinished; - const client_ciphertext = client_data[5 .. 5 + client_ciphertext_len]; - - const client_decrypted = record.decryptRecord( - client_hs_traffic.key, - client_hs_traffic.iv, - client_seq, - client_ciphertext, - client_rec_header, - ) catch return error.InvalidClientFinished; - client_seq += 1; - - if (client_decrypted.content_type != .handshake) return error.InvalidClientFinished; - if (client_decrypted.plaintext.len < 4 + hash_len) return error.InvalidClientFinished; - if (client_decrypted.plaintext[0] != 0x14) return error.InvalidClientFinished; - - const client_fin_transcript_hash = transcript.peek(); - const expected_verify = handshake.computeFinished( - hs_keys.client_handshake_traffic_secret, - client_fin_transcript_hash, - ); - - if (!std.mem.eql(u8, client_decrypted.plaintext[4 .. 4 + hash_len], &expected_verify)) - return error.FinishedVerifyFailed; - - transcript.update(client_decrypted.plaintext); - var app_transcript_hash: [hash_len]u8 = undefined; transcript.final(&app_transcript_hash); @@ -166,6 +184,41 @@ pub fn handleTlsSession( handshake_complete.* = true; + return .{ + .selected_alpn = selected_alpn, + .app_keys = app_keys, + .peer_identity = peer_identity_out, + }; +} + +/// run the server handshake, then forward bytes between the encrypted +/// client and the plaintext backend until either side closes. preserves +/// the existing non-mTLS code path — no `MtlsOpts` are threaded through +/// here yet (PR 5 wires the listener-level decision). +pub fn handleTlsSession( + io: std.Io, + client_fd: posix.fd_t, + client_hello: []const u8, + cert_pem: []u8, + key_pem: []u8, + backend_info: backend_mod.Backend, + handshake_complete: *bool, +) !void { + var session = try acceptServerHandshake( + io, + std.heap.page_allocator, + client_fd, + client_hello, + cert_pem, + key_pem, + null, + handshake_complete, + ); + defer session.deinit(std.heap.page_allocator); + + const selected_alpn = session.selected_alpn; + const app_keys = session.app_keys; + const backend_fd = socket_support.connectToBackend(backend_info) catch return error.BackendConnectFailed; defer linux_platform.posix.close(backend_fd); @@ -304,6 +357,230 @@ fn injectForwardedProtoHttp1(alloc: std.mem.Allocator, request: []const u8, prot return out.toOwnedSlice(alloc); } +/// non-mTLS path: the client sends just an (optionally CCS-prefixed) +/// encrypted Finished. existing behavior, factored out into a helper. +fn readPlainClientFinished( + client_fd: posix.fd_t, + keys: handshake.TrafficKeys, + client_seq: *u64, + transcript: *Sha384, + hs_keys: handshake.HandshakeKeys, +) !void { + var client_finished_buf: [512]u8 = undefined; + const client_rec_n = socket_support.readWithTimeout(client_fd, &client_finished_buf, 10000) catch + return error.ReadFailed; + if (client_rec_n < record.record_header_size + record.aead_tag_size + 1) + return error.InvalidClientFinished; + + var client_data = client_finished_buf[0..client_rec_n]; + if (client_data[0] == 0x14) { + if (client_data.len < 6) return error.InvalidClientFinished; + const ccs_len: usize = 5 + @as(usize, (@as(u16, client_data[3]) << 8) | @as(u16, client_data[4])); + if (ccs_len > client_data.len) return error.InvalidClientFinished; + client_data = client_data[ccs_len..]; + if (client_data.len < record.record_header_size + record.aead_tag_size + 1) + return error.InvalidClientFinished; + } + + const client_rec_header = client_data[0..5].*; + const client_ciphertext_len = (@as(usize, client_data[3]) << 8) | @as(usize, client_data[4]); + if (5 + client_ciphertext_len > client_data.len) return error.InvalidClientFinished; + const client_ciphertext = client_data[5 .. 5 + client_ciphertext_len]; + + const client_decrypted = record.decryptRecord( + keys.key, + keys.iv, + client_seq.*, + client_ciphertext, + client_rec_header, + ) catch return error.InvalidClientFinished; + client_seq.* += 1; + + if (client_decrypted.content_type != .handshake) return error.InvalidClientFinished; + if (client_decrypted.plaintext.len < 4 + hash_len) return error.InvalidClientFinished; + if (client_decrypted.plaintext[0] != 0x14) return error.InvalidClientFinished; + + const client_fin_transcript_hash = transcript.peek(); + const expected_verify = handshake.computeFinished( + hs_keys.client_handshake_traffic_secret, + client_fin_transcript_hash, + ); + + if (!std.mem.eql(u8, client_decrypted.plaintext[4 .. 4 + hash_len], &expected_verify)) + return error.FinishedVerifyFailed; + + transcript.update(client_decrypted.plaintext); +} + +/// mTLS path: read the client's Certificate, optionally CertificateVerify, +/// then Finished. messages may arrive over one or more encrypted records; +/// accumulate plaintext in `pending` and dispatch by handshake message +/// type. on a verified non-empty client cert, dup the SAN URI into +/// `peer_identity_out` (caller frees). +fn acceptClientAuthAndFinished( + alloc: std.mem.Allocator, + client_fd: posix.fd_t, + keys: handshake.TrafficKeys, + client_seq: *u64, + transcript: *Sha384, + hs_keys: handshake.HandshakeKeys, + opts: MtlsOpts, + peer_identity_out: *?[]u8, +) !void { + var pending: std.ArrayList(u8) = .empty; + defer pending.deinit(alloc); + + var client_cert_der: ?[]u8 = null; + defer if (client_cert_der) |d| alloc.free(d); + var client_cert_pem: ?[]u8 = null; + defer if (client_cert_pem) |p| alloc.free(p); + + var got_cert_verify_sig: ?[]u8 = null; + defer if (got_cert_verify_sig) |s| alloc.free(s); + var transcript_hash_at_cv: ?[hash_len]u8 = null; + + var saw_finished = false; + var saw_certificate = false; + + while (!saw_finished) { + const plaintext = try readOneEncryptedHandshakeRecordAlloc(alloc, client_fd, keys, client_seq); + defer alloc.free(plaintext); + try pending.appendSlice(alloc, plaintext); + + var pos: usize = 0; + while (true) { + const opt_msg = message_parse.nextMessage(pending.items, &pos) catch break; + const msg = opt_msg orelse break; + + switch (msg.msg_type) { + @intFromEnum(message_parse.HandshakeType.certificate) => { + if (saw_certificate) return error.InvalidClientFinished; + saw_certificate = true; + const der = message_parse.parseCertificateMessage(msg.body) catch return error.InvalidClientFinished; + if (der.len > 0) { + client_cert_der = try alloc.dupe(u8, der); + client_cert_pem = try derToPemLocal(alloc, der); + } + transcript.update(msg.raw); + }, + @intFromEnum(message_parse.HandshakeType.certificate_verify) => { + if (client_cert_der == null) return error.InvalidClientFinished; + const cv = message_parse.parseCertificateVerify(msg.body) catch return error.InvalidClientFinished; + if (cv.algorithm != 0x0403) return error.InvalidClientFinished; + transcript_hash_at_cv = transcript.peek(); + got_cert_verify_sig = try alloc.dupe(u8, cv.signature_der); + transcript.update(msg.raw); + }, + @intFromEnum(message_parse.HandshakeType.finished) => { + const fin = message_parse.parseFinished(msg.body) catch return error.InvalidClientFinished; + const expected = handshake.computeFinished(hs_keys.client_handshake_traffic_secret, transcript.peek()); + if (!std.mem.eql(u8, fin.verify_data, &expected)) return error.FinishedVerifyFailed; + transcript.update(msg.raw); + saw_finished = true; + break; + }, + else => return error.InvalidClientFinished, + } + } + + if (pos > 0) { + const remaining = pending.items.len - pos; + std.mem.copyForwards(u8, pending.items[0..remaining], pending.items[pos..]); + pending.shrinkRetainingCapacity(remaining); + } + } + + // verify the client cert (if any) before declaring the handshake done. + if (client_cert_pem) |cpem| { + x509_verify.verifyLeafAgainstCa(alloc, cpem, opts.trust_ca_pem, opts.expected_identity, opts.now_unix) catch return error.UntrustedClientCert; + + // verify the client's CertificateVerify signature against the + // cert's public key — the spec-required peer-of-keypair check. + const sig = got_cert_verify_sig orelse return error.InvalidClientFinished; + const cv_hash = transcript_hash_at_cv orelse return error.InvalidClientFinished; + try verifyClientCertVerify(client_cert_der.?, sig, cv_hash); + + // surface the identity (use the SAN URI if present, fall back to + // subject CN) so callers can audit / authorize. + var san_buf: [8][]const u8 = undefined; + const parsed = x509_verify.parseDer(client_cert_der.?, &san_buf) catch return error.InvalidClientFinished; + if (parsed.san_uris.len > 0) { + peer_identity_out.* = try alloc.dupe(u8, parsed.san_uris[0]); + } else { + peer_identity_out.* = try alloc.dupe(u8, parsed.subject_cn); + } + } else if (opts.require_client_cert) { + return error.MissingClientCert; + } +} + +/// read one encrypted handshake record off the wire, decrypt, return the +/// plaintext. helper for the mTLS client-auth read loop. +fn readOneEncryptedHandshakeRecordAlloc( + alloc: std.mem.Allocator, + fd: posix.fd_t, + keys: handshake.TrafficKeys, + seq: *u64, +) ![]u8 { + var hdr: [5]u8 = undefined; + var off: usize = 0; + while (off < 5) { + const n = posix.read(fd, hdr[off..]) catch return error.ReadFailed; + if (n == 0) return error.UnexpectedEof; + off += n; + } + const payload_len = (@as(usize, hdr[3]) << 8) | @as(usize, hdr[4]); + if (payload_len > record.max_ciphertext_size) return error.InvalidClientFinished; + const ciphertext = try alloc.alloc(u8, payload_len); + defer alloc.free(ciphertext); + off = 0; + while (off < payload_len) { + const n = posix.read(fd, ciphertext[off..]) catch return error.ReadFailed; + if (n == 0) return error.UnexpectedEof; + off += n; + } + const dec = record.decryptRecord(keys.key, keys.iv, seq.*, ciphertext, hdr) catch return error.DecryptFailed; + seq.* += 1; + if (dec.content_type != .handshake) return error.InvalidClientFinished; + return try alloc.dupe(u8, dec.plaintext); +} + +fn verifyClientCertVerify(client_cert_der: []const u8, sig_der: []const u8, transcript_hash: [hash_len]u8) !void { + var san_buf: [8][]const u8 = undefined; + const parsed = x509_verify.parseDer(client_cert_der, &san_buf) catch return error.UntrustedClientCert; + if (parsed.public_key_point.len != 65) return error.UntrustedClientCert; + var sec1: [65]u8 = undefined; + @memcpy(&sec1, parsed.public_key_point[0..65]); + const pub_key = EcdsaP256.PublicKey.fromSec1(&sec1) catch return error.UntrustedClientCert; + + var signed_content: [64 + 33 + 1 + hash_len]u8 = undefined; + defer std.crypto.secureZero(u8, &signed_content); + @memset(signed_content[0..64], 0x20); + const ctx = "TLS 1.3, client CertificateVerify"; + @memcpy(signed_content[64 .. 64 + ctx.len], ctx); + signed_content[64 + ctx.len] = 0x00; + @memcpy(signed_content[64 + ctx.len + 1 ..], &transcript_hash); + + var sig_buf: [EcdsaP256.Signature.der_encoded_length_max]u8 = undefined; + if (sig_der.len > sig_buf.len) return error.UntrustedClientCert; + @memcpy(sig_buf[0..sig_der.len], sig_der); + const sig = EcdsaP256.Signature.fromDer(sig_buf[0..sig_der.len]) catch return error.UntrustedClientCert; + sig.verify(&signed_content, pub_key) catch return error.UntrustedClientCert; +} + +fn derToPemLocal(alloc: std.mem.Allocator, der: []const u8) ![]u8 { + var out: std.ArrayList(u8) = .empty; + errdefer out.deinit(alloc); + try out.appendSlice(alloc, "-----BEGIN CERTIFICATE-----\n"); + const enc_size = std.base64.standard.Encoder.calcSize(der.len); + const tmp = try alloc.alloc(u8, enc_size); + defer alloc.free(tmp); + _ = std.base64.standard.Encoder.encode(tmp, der); + try out.appendSlice(alloc, tmp); + try out.appendSlice(alloc, "\n-----END CERTIFICATE-----\n"); + return out.toOwnedSlice(alloc); +} + pub fn sendEncryptedHandshake(fd: posix.fd_t, msg: []const u8, keys: handshake.TrafficKeys, seq: *u64) !void { var ct_buf: [record.max_ciphertext_size]u8 = undefined; const ct_len = record.encryptRecord( @@ -356,3 +633,237 @@ test "injectForwardedProtoHttp1 rewrites the first request headers" { try std.testing.expect(std.mem.indexOf(u8, rewritten, "X-Forwarded-Proto: https\r\n") != null); try std.testing.expect(std.mem.indexOf(u8, rewritten, "X-Forwarded-Proto: http\r\n") == null); } + +// --- mTLS round-trip tests --- +// +// these spin up a unix socketpair, run `acceptServerHandshake` with +// `MtlsOpts` on one side and the client driver (`client_session`) on the +// other, and assert the negotiated peer identity is the client cert's +// SAN URI. + +const x509_gen = @import("../x509_gen.zig"); +const csr = @import("../csr.zig"); +const client_session = @import("../client_session.zig"); + +const mtls_test_now: i64 = 1_700_000_000; +const mtls_test_window: i64 = 24 * 3600; + +const MtlsCerts = struct { + ca_pem: []u8, + server_pem: []u8, + server_key_pem: []u8, + client_pem: []u8, + client_key_pem: []u8, + + fn deinit(self: *MtlsCerts, alloc: std.mem.Allocator) void { + alloc.free(self.ca_pem); + alloc.free(self.server_pem); + alloc.free(self.server_key_pem); + alloc.free(self.client_pem); + alloc.free(self.client_key_pem); + } +}; + +fn mintMtlsCerts(alloc: std.mem.Allocator) !MtlsCerts { + const ca = try x509_gen.generateCa(std.testing.io, alloc, "mtls-ca", mtls_test_now - 3600, mtls_test_now + mtls_test_window); + errdefer alloc.free(ca.cert_pem); + + const server = try x509_gen.issueLeaf(std.testing.io, alloc, ca.key_pair, "mtls-ca", "server", "spiffe://yoq/service/server", mtls_test_now - 60, mtls_test_now + mtls_test_window); + errdefer alloc.free(server.cert_pem); + const server_key_pem = try csr.derKeyToPem(alloc, &server.key_pair.secret_key.toBytes()); + errdefer alloc.free(server_key_pem); + + const client = try x509_gen.issueLeaf(std.testing.io, alloc, ca.key_pair, "mtls-ca", "client", "spiffe://yoq/service/client", mtls_test_now - 60, mtls_test_now + mtls_test_window); + errdefer alloc.free(client.cert_pem); + const client_key_pem = try csr.derKeyToPem(alloc, &client.key_pair.secret_key.toBytes()); + + return .{ + .ca_pem = ca.cert_pem, + .server_pem = server.cert_pem, + .server_key_pem = server_key_pem, + .client_pem = client.cert_pem, + .client_key_pem = client_key_pem, + }; +} + +const MtlsServerThreadArgs = struct { + fd: posix.fd_t, + certs: *const MtlsCerts, + require_client_cert: bool, + peer_identity_out: *?[]u8, + err_out: *?anyerror, + alloc: std.mem.Allocator, +}; + +fn runMtlsServer(args: MtlsServerThreadArgs) void { + runMtlsServerImpl(args) catch |err| { + args.err_out.* = err; + }; +} + +fn runMtlsServerImpl(args: MtlsServerThreadArgs) !void { + // pre-read the ClientHello record off the wire so acceptServerHandshake + // can consume it (matching the production listener convention). + var ch_buf: [4096]u8 = undefined; + var ch_len: usize = 0; + while (ch_len < 5) { + const n = try posix.read(args.fd, ch_buf[ch_len..]); + if (n == 0) return error.UnexpectedEof; + ch_len += n; + } + const promised = (@as(usize, ch_buf[3]) << 8) | @as(usize, ch_buf[4]); + while (ch_len < 5 + promised) { + const n = try posix.read(args.fd, ch_buf[ch_len..]); + if (n == 0) return error.UnexpectedEof; + ch_len += n; + } + + var handshake_complete = false; + var session = try acceptServerHandshake( + std.testing.io, + args.alloc, + args.fd, + ch_buf[0..ch_len], + args.certs.server_pem, + args.certs.server_key_pem, + .{ + .require_client_cert = args.require_client_cert, + .trust_ca_pem = args.certs.ca_pem, + .now_unix = mtls_test_now, + }, + &handshake_complete, + ); + + // hand the identity back to the test before deinit frees it. + if (session.peer_identity) |p| { + args.peer_identity_out.* = try args.alloc.dupe(u8, p); + } + session.deinit(args.alloc); +} + +test "mTLS round-trip — valid client cert is accepted and identity surfaces" { + const alloc = std.testing.allocator; + + var fds: [2]i32 = undefined; + const rc = std.os.linux.socketpair(std.posix.AF.UNIX, std.posix.SOCK.STREAM, 0, &fds); + if (rc != 0) return error.SkipZigTest; + defer linux_platform.posix.close(fds[0]); + defer linux_platform.posix.close(fds[1]); + + var certs = try mintMtlsCerts(alloc); + defer certs.deinit(alloc); + + var peer_id: ?[]u8 = null; + defer if (peer_id) |p| alloc.free(p); + var server_err: ?anyerror = null; + + const args: MtlsServerThreadArgs = .{ + .fd = fds[1], + .certs = &certs, + .require_client_cert = true, + .peer_identity_out = &peer_id, + .err_out = &server_err, + .alloc = alloc, + }; + const t = try std.Thread.spawn(.{}, runMtlsServer, .{args}); + + var sess = try client_session.doHandshake(std.testing.io, alloc, fds[0], .{ + .ca_cert_pem = certs.ca_pem, + .expected_server_identity = "spiffe://yoq/service/server", + .client_cert_pem = certs.client_pem, + .client_key_pem = certs.client_key_pem, + .now_unix = mtls_test_now, + }); + defer sess.deinit(); + + t.join(); + try std.testing.expect(server_err == null); + try std.testing.expect(peer_id != null); + try std.testing.expectEqualStrings("spiffe://yoq/service/client", peer_id.?); +} + +test "mTLS — require_client_cert rejects an empty client cert" { + const alloc = std.testing.allocator; + + var fds: [2]i32 = undefined; + const rc = std.os.linux.socketpair(std.posix.AF.UNIX, std.posix.SOCK.STREAM, 0, &fds); + if (rc != 0) return error.SkipZigTest; + defer linux_platform.posix.close(fds[0]); + defer linux_platform.posix.close(fds[1]); + + var certs = try mintMtlsCerts(alloc); + defer certs.deinit(alloc); + + var peer_id: ?[]u8 = null; + defer if (peer_id) |p| alloc.free(p); + var server_err: ?anyerror = null; + + const args: MtlsServerThreadArgs = .{ + .fd = fds[1], + .certs = &certs, + .require_client_cert = true, + .peer_identity_out = &peer_id, + .err_out = &server_err, + .alloc = alloc, + }; + const t = try std.Thread.spawn(.{}, runMtlsServer, .{args}); + + // client doesn't pass client_cert_pem — the driver sends an empty + // Certificate, the server rejects. + const handshake_err = client_session.doHandshake(std.testing.io, alloc, fds[0], .{ + .ca_cert_pem = certs.ca_pem, + .now_unix = mtls_test_now, + }); + // the client doesn't see the rejection directly (server returns the + // error after the client's Finished); the client typically reports + // success or a later read failure. either is acceptable — we assert + // on the server side instead. + if (handshake_err) |*ok| { + var ok_mut = ok.*; + ok_mut.deinit(); + } else |_| {} + + _ = std.os.linux.shutdown(fds[0], 2); // SHUT_RDWR + t.join(); + + try std.testing.expect(peer_id == null); + try std.testing.expect(server_err != null); + try std.testing.expectEqual(@as(anyerror, error.MissingClientCert), server_err.?); +} + +test "mTLS — warn mode accepts an empty client cert (no identity)" { + const alloc = std.testing.allocator; + + var fds: [2]i32 = undefined; + const rc = std.os.linux.socketpair(std.posix.AF.UNIX, std.posix.SOCK.STREAM, 0, &fds); + if (rc != 0) return error.SkipZigTest; + defer linux_platform.posix.close(fds[0]); + defer linux_platform.posix.close(fds[1]); + + var certs = try mintMtlsCerts(alloc); + defer certs.deinit(alloc); + + var peer_id: ?[]u8 = null; + defer if (peer_id) |p| alloc.free(p); + var server_err: ?anyerror = null; + + const args: MtlsServerThreadArgs = .{ + .fd = fds[1], + .certs = &certs, + .require_client_cert = false, // warn mode + .peer_identity_out = &peer_id, + .err_out = &server_err, + .alloc = alloc, + }; + const t = try std.Thread.spawn(.{}, runMtlsServer, .{args}); + + var sess = try client_session.doHandshake(std.testing.io, alloc, fds[0], .{ + .ca_cert_pem = certs.ca_pem, + .now_unix = mtls_test_now, + }); + defer sess.deinit(); + + t.join(); + try std.testing.expect(server_err == null); + try std.testing.expect(peer_id == null); // no identity surfaced +}