Files
voice-cat/tests/test_tls_loopback.cpp

187 lines
6.2 KiB
C++
Raw Permalink Normal View History

2026-06-15 23:48:44 +02:00
/*
* test_tls_loopback in-process TLS 1.3 server + client over a loopback TCP socket pair.
* Validates: cert generation, handshake, ServerIdentity fingerprint, framed message exchange.
*/
#include <atomic>
#include <cstdio>
#include <cstring>
#include <thread>
#include <vector>
#ifdef _WIN32
# include <winsock2.h>
# include <ws2tcpip.h>
# pragma comment(lib, "ws2_32.lib")
using sock_t = SOCKET;
static constexpr sock_t kBadSock = INVALID_SOCKET;
static void close_sock(sock_t s) { closesocket(s); }
static int last_err() { return WSAGetLastError(); }
#else
# include <arpa/inet.h>
# include <netinet/in.h>
# include <sys/socket.h>
# include <unistd.h>
using sock_t = int;
static constexpr sock_t kBadSock = -1;
static void close_sock(sock_t s) { ::close(s); }
static int last_err() { return errno; }
#endif
#include "crypto/crypto.h"
using namespace voicecat::crypto;
static int g_failures = 0;
#define CHECK(cond) \
do { if (!(cond)) { std::printf("FAIL [%s:%d]: %s\n", __FILE__, __LINE__, #cond); ++g_failures; } } while (0)
// Create a blocking loopback TCP socket pair: returns {server_fd, client_fd}
static std::pair<sock_t, sock_t> make_socket_pair(uint16_t port) {
sock_t listener = ::socket(AF_INET, SOCK_STREAM, 0);
if (listener == kBadSock) return {kBadSock, kBadSock};
int opt = 1;
setsockopt(listener, SOL_SOCKET, SO_REUSEADDR,
reinterpret_cast<const char*>(&opt), sizeof(opt));
sockaddr_in addr{};
addr.sin_family = AF_INET;
addr.sin_addr.s_addr = htonl(INADDR_LOOPBACK);
addr.sin_port = htons(port);
if (::bind(listener, reinterpret_cast<sockaddr*>(&addr), sizeof(addr)) != 0) {
close_sock(listener); return {kBadSock, kBadSock};
}
if (::listen(listener, 1) != 0) {
close_sock(listener); return {kBadSock, kBadSock};
}
sock_t client = ::socket(AF_INET, SOCK_STREAM, 0);
if (client == kBadSock) { close_sock(listener); return {kBadSock, kBadSock}; }
if (::connect(client, reinterpret_cast<sockaddr*>(&addr), sizeof(addr)) != 0) {
close_sock(listener); close_sock(client); return {kBadSock, kBadSock};
}
sockaddr_in peer{};
socklen_t plen = sizeof(peer);
sock_t server = ::accept(listener, reinterpret_cast<sockaddr*>(&peer), &plen);
close_sock(listener);
if (server == kBadSock) { close_sock(client); return {kBadSock, kBadSock}; }
return {server, client};
}
// Write all bytes to a TLS context.
static bool tls_write_all(TlsContext& tls, const uint8_t* data, size_t len) {
size_t off = 0;
while (off < len) {
int n = tls.write(data + off, len - off);
if (n <= 0) return false;
off += n;
}
return true;
}
// Read exactly len bytes from a TLS context.
static bool tls_read_exact(TlsContext& tls, uint8_t* buf, size_t len) {
size_t off = 0;
while (off < len) {
int n = tls.read(buf + off, len - off);
if (n <= 0) return false;
off += n;
}
return true;
}
int main() {
#ifdef _WIN32
WSADATA wsa{};
if (WSAStartup(MAKEWORD(2, 2), &wsa) != 0) {
std::printf("WSAStartup failed\n");
return 1;
}
#endif
// Generate server identity + cert
ServerIdentity identity = ServerIdentity::generate();
ServerCert cert = ServerCert::generate("test-server");
CHECK(!cert.pem_cert.empty());
CHECK(!cert.pem_key.empty());
auto [server_fd_native, client_fd_native] = make_socket_pair(19851);
CHECK(server_fd_native != kBadSock);
CHECK(client_fd_native != kBadSock);
if (server_fd_native == kBadSock || client_fd_native == kBadSock) {
std::printf("tls_loopback: socket pair failed (err=%d)\n", last_err());
return 1;
}
std::string server_error, client_error;
std::atomic<bool> server_ok{false}, client_ok{false};
static const char kMsg1[] = "hello from server";
static const char kMsg2[] = "hello from client";
constexpr size_t kMsg1Len = sizeof(kMsg1) - 1;
constexpr size_t kMsg2Len = sizeof(kMsg2) - 1;
char client_recv[64]{};
char server_recv[64]{};
// Server thread: handshake, send msg1, recv msg2
std::thread server_thr([&] {
TlsContext tls(TlsContext::Role::Server, &cert);
int fd = static_cast<int>(server_fd_native);
if (!tls.handshake(fd, server_error)) { close_sock(server_fd_native); return; }
server_ok.store(true);
tls_write_all(tls, reinterpret_cast<const uint8_t*>(kMsg1), kMsg1Len);
tls_read_exact(tls, reinterpret_cast<uint8_t*>(server_recv), kMsg2Len);
close_sock(server_fd_native);
});
// Client thread: handshake, recv msg1, send msg2
std::thread client_thr([&] {
TlsContext tls(TlsContext::Role::Client, nullptr);
int fd = static_cast<int>(client_fd_native);
if (!tls.handshake(fd, client_error)) { close_sock(client_fd_native); return; }
client_ok.store(true);
tls_read_exact(tls, reinterpret_cast<uint8_t*>(client_recv), kMsg1Len);
tls_write_all(tls, reinterpret_cast<const uint8_t*>(kMsg2), kMsg2Len);
close_sock(client_fd_native);
});
server_thr.join();
client_thr.join();
if (!server_error.empty()) std::printf("server TLS error: %s\n", server_error.c_str());
if (!client_error.empty()) std::printf("client TLS error: %s\n", client_error.c_str());
CHECK(server_ok.load());
CHECK(client_ok.load());
CHECK(std::memcmp(client_recv, kMsg1, kMsg1Len) == 0);
CHECK(std::memcmp(server_recv, kMsg2, kMsg2Len) == 0);
// Verify ServerIdentity round-trip
{
ServerIdentity id2 = ServerIdentity::generate();
CHECK(id2.pk != identity.pk); // different key
// Fingerprint is SHA-256 of pk — non-zero
bool nonzero = false;
for (auto b : id2.fingerprint) if (b) { nonzero = true; break; }
CHECK(nonzero);
// fingerprint_hex should be 32 colons + 64 hex chars = 95 chars (AA:BB:...)
std::string hex = id2.fingerprint_hex();
CHECK(hex.size() == 95);
}
#ifdef _WIN32
WSACleanup();
#endif
if (g_failures == 0) {
std::printf("tls_loopback: all checks passed\n");
return 0;
}
std::printf("tls_loopback: %d failure(s)\n", g_failures);
return 1;
}