187 lines
6.2 KiB
C++
187 lines
6.2 KiB
C++
|
|
/*
|
||
|
|
* 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;
|
||
|
|
}
|