/* * 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 #include #include #include #include #ifdef _WIN32 # include # include # 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 # include # include # include 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 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(&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(&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(&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(&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 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(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(kMsg1), kMsg1Len); tls_read_exact(tls, reinterpret_cast(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(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(client_recv), kMsg1Len); tls_write_all(tls, reinterpret_cast(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; }