srdusr
aboutsummaryrefslogtreecommitdiffstats
path: root/tests/test_tls.cpp
diff options
context:
space:
mode:
Diffstat (limited to 'tests/test_tls.cpp')
-rw-r--r--tests/test_tls.cpp132
1 files changed, 132 insertions, 0 deletions
diff --git a/tests/test_tls.cpp b/tests/test_tls.cpp
new file mode 100644
index 0000000..65784a9
--- /dev/null
+++ b/tests/test_tls.cpp
@@ -0,0 +1,132 @@
+#include <doctest/doctest.h>
+
+#include <vector>
+
+#include "wireframe/l7/tls.hpp"
+
+using namespace wireframe::net;
+
+namespace {
+
+void append_be16(std::vector<unsigned char>& out, std::uint16_t v) {
+ out.push_back(static_cast<unsigned char>(v >> 8));
+ out.push_back(static_cast<unsigned char>(v & 0xFF));
+}
+
+// Builds a real, well-formed TLS record containing a ClientHello with
+// (optionally) a single SNI host_name extension. Every length field is
+// computed from the actual bytes assembled, not hand-counted - the
+// same lesson from this session's earlier UDP-length test typo.
+//
+// `corrupt_sni_ext_len`, when set, writes an oversized SNI extension
+// length instead of the real one (computed here, not via post-hoc
+// offset math into the finished buffer - equally fragile).
+std::vector<unsigned char> build_client_hello(const std::string& sni,
+ bool corrupt_sni_ext_len = false) {
+ std::vector<unsigned char> body;
+ body.push_back(0x03);
+ body.push_back(0x03); // client_version: TLS 1.2
+ body.insert(body.end(), 32, 0x00); // random
+ body.push_back(0x00); // session_id length: 0
+ append_be16(body, 2); // cipher_suites length
+ body.push_back(0x00);
+ body.push_back(0x2F); // one arbitrary cipher suite
+ body.push_back(0x01); // compression_methods length: 1
+ body.push_back(0x00); // null compression
+
+ std::vector<unsigned char> extensions;
+ if (!sni.empty()) {
+ std::vector<unsigned char> server_name_list;
+ server_name_list.push_back(0x00); // name_type: host_name
+ append_be16(server_name_list, static_cast<std::uint16_t>(sni.size()));
+ server_name_list.insert(server_name_list.end(), sni.begin(), sni.end());
+
+ std::vector<unsigned char> sni_ext_data;
+ append_be16(sni_ext_data, static_cast<std::uint16_t>(server_name_list.size()));
+ sni_ext_data.insert(sni_ext_data.end(), server_name_list.begin(), server_name_list.end());
+
+ append_be16(extensions, kTlsExtensionServerName);
+ std::uint16_t ext_len = corrupt_sni_ext_len
+ ? static_cast<std::uint16_t>(0xFFFF)
+ : static_cast<std::uint16_t>(sni_ext_data.size());
+ append_be16(extensions, ext_len);
+ extensions.insert(extensions.end(), sni_ext_data.begin(), sni_ext_data.end());
+ }
+ append_be16(body, static_cast<std::uint16_t>(extensions.size()));
+ body.insert(body.end(), extensions.begin(), extensions.end());
+
+ std::vector<unsigned char> handshake;
+ handshake.push_back(kTlsHandshakeTypeClientHello);
+ std::uint32_t hs_len = static_cast<std::uint32_t>(body.size());
+ handshake.push_back(static_cast<unsigned char>((hs_len >> 16) & 0xFF));
+ handshake.push_back(static_cast<unsigned char>((hs_len >> 8) & 0xFF));
+ handshake.push_back(static_cast<unsigned char>(hs_len & 0xFF));
+ handshake.insert(handshake.end(), body.begin(), body.end());
+
+ std::vector<unsigned char> record;
+ record.push_back(kTlsContentTypeHandshake);
+ record.push_back(0x03);
+ record.push_back(0x01); // record-layer version (legacy compat value)
+ append_be16(record, static_cast<std::uint16_t>(handshake.size()));
+ record.insert(record.end(), handshake.begin(), handshake.end());
+
+ return record;
+}
+
+} // namespace
+
+TEST_CASE("parse_tls_client_hello extracts a real SNI extension") {
+ auto record = build_client_hello("example.com");
+ auto hello = parse_tls_client_hello(record);
+ REQUIRE(hello.has_value());
+ REQUIRE(hello->server_name.has_value());
+ CHECK(*hello->server_name == "example.com");
+}
+
+TEST_CASE("parse_tls_client_hello succeeds with no SNI when there's no extensions block") {
+ auto record = build_client_hello("");
+ auto hello = parse_tls_client_hello(record);
+ REQUIRE(hello.has_value());
+ CHECK_FALSE(hello->server_name.has_value());
+}
+
+TEST_CASE("parse_tls_client_hello rejects a non-Handshake record") {
+ auto record = build_client_hello("example.com");
+ record[0] = 0x17; // application_data, not handshake
+ CHECK_FALSE(parse_tls_client_hello(record).has_value());
+}
+
+TEST_CASE("parse_tls_client_hello rejects a non-ClientHello handshake type") {
+ auto record = build_client_hello("example.com");
+ record[5] = 0x02; // ServerHello, not ClientHello
+ CHECK_FALSE(parse_tls_client_hello(record).has_value());
+}
+
+TEST_CASE("parse_tls_client_hello rejects a truncated record") {
+ auto record = build_client_hello("example.com");
+ record.resize(record.size() - 5); // claims more than it has
+ CHECK_FALSE(parse_tls_client_hello(record).has_value());
+}
+
+TEST_CASE("parse_tls_client_hello rejects a buffer shorter than the record header") {
+ std::vector<unsigned char> bytes(4, 0);
+ CHECK_FALSE(parse_tls_client_hello(bytes).has_value());
+}
+
+TEST_CASE("parse_tls_client_hello stops gracefully on a malformed extension length") {
+ auto record = build_client_hello("example.com", /*corrupt_sni_ext_len=*/true);
+ auto hello = parse_tls_client_hello(record);
+ REQUIRE(hello.has_value()); // still a structurally valid ClientHello otherwise
+ CHECK_FALSE(hello->server_name.has_value()); // SNI extension was malformed, so skipped
+}
+
+TEST_CASE("TlsSniDissector claims port 443 and its summary matches parse_tls_client_hello") {
+ TlsSniDissector dissector;
+ CHECK(dissector.port() == kTlsPort);
+
+ auto record = build_client_hello("wireframe.test");
+ auto summary = dissector.summarize(record);
+ REQUIRE(summary.has_value());
+ CHECK(summary->substr(0, 3) == "TLS");
+ CHECK(summary->find("SNI=wireframe.test") != std::string::npos);
+}