srdusr
aboutsummaryrefslogtreecommitdiffstats
path: root/tests/test_tls.cpp
blob: 7a3ad271046435d6aacb8599dbb7cddbdac5dbe6 (plain) (blame)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
#include <doctest/doctest.h>

#include <vector>

#include "packeteer/l7/tls.hpp"

using namespace packeteer::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("packeteer.test");
    auto summary = dissector.summarize(record);
    REQUIRE(summary.has_value());
    CHECK(summary->substr(0, 3) == "TLS");
    CHECK(summary->find("SNI=packeteer.test") != std::string::npos);
}