websocket: rebuild handshake on Request/header pipeline

Replace the hand-built upgrade request string in
perform_websocket_handshake with the same Request,
write_request_line, and check_and_write_headers path used by
ClientImpl::open_stream. Add WebSocketClient::prepare_default_headers
to inject Host and User-Agent defaults; protocol-mandatory headers
(Upgrade, Connection, Sec-WebSocket-Key/Version) are always
overwritten.
This commit is contained in:
hyuk.kim
2026-08-01 17:15:58 +09:00
parent f24e79aab9
commit 2c28d2fa3e
2 changed files with 191 additions and 37 deletions
+53 -37
View File
@@ -3154,6 +3154,10 @@ private:
std::string make_host_and_port_string(const std::string &host, int port,
bool is_ssl);
template <typename T>
bool check_and_write_headers(Stream &strm, Headers &headers, T header_writer,
Error &error);
std::string trim_copy(const std::string &s);
void divide(
@@ -3992,6 +3996,7 @@ public:
private:
void shutdown_and_close();
bool create_stream(std::unique_ptr<Stream> &strm);
void prepare_default_headers(Request &req);
std::string host_;
int port_;
@@ -9049,21 +9054,8 @@ inline bool is_field_valid(const std::string &name, const std::string &value) {
} // namespace fields
inline bool perform_websocket_handshake(Stream &strm, const std::string &host,
int port, bool is_ssl,
const std::string &path,
const Headers &headers,
inline bool perform_websocket_handshake(Stream &strm, Request &req,
std::string &selected_subprotocol) {
// Validate path and host
if (!fields::is_field_value(path) || !fields::is_field_value(host)) {
return false;
}
// Validate user-provided headers
for (const auto &h : headers) {
if (!fields::is_field_valid(h.first, h.second)) { return false; }
}
// Generate random Sec-WebSocket-Key
thread_local std::mt19937 rng(std::random_device{}());
std::string key_bytes(16, '\0');
@@ -9073,19 +9065,21 @@ inline bool perform_websocket_handshake(Stream &strm, const std::string &host,
}
auto client_key = base64_encode(key_bytes);
// Build upgrade request
std::string req_str = "GET " + path + " HTTP/1.1\r\n";
req_str += "Host: " + make_host_and_port_string(host, port, is_ssl) + "\r\n";
req_str += "Upgrade: websocket\r\n";
req_str += "Connection: Upgrade\r\n";
req_str += "Sec-WebSocket-Key: " + client_key + "\r\n";
req_str += "Sec-WebSocket-Version: 13\r\n";
for (const auto &h : headers) {
req_str += h.first + ": " + h.second + "\r\n";
}
req_str += "\r\n";
req.headers.erase("Upgrade");
req.headers.erase("Connection");
req.headers.erase("Sec-WebSocket-Key");
req.headers.erase("Sec-WebSocket-Version");
req.headers.emplace("Upgrade", "websocket");
req.headers.emplace("Connection", "Upgrade");
req.headers.emplace("Sec-WebSocket-Key", client_key);
req.headers.emplace("Sec-WebSocket-Version", "13");
if (strm.write(req_str.data(), req_str.size()) < 0) { return false; }
if (write_request_line(strm, req.method, req.path) < 0) { return false; }
auto error = Error::Success;
if (!check_and_write_headers(strm, req.headers, write_headers, error)) {
return false;
}
// Verify 101 response and Sec-WebSocket-Accept header
auto expected_accept = websocket_accept_key(client_key);
@@ -20828,11 +20822,37 @@ inline bool WebSocketClient::create_stream(std::unique_ptr<Stream> &strm) {
return true;
}
inline void WebSocketClient::prepare_default_headers(Request &req) {
#ifdef CPPHTTPLIB_SSL_ENABLED
auto is_ssl = is_ssl_;
#else
auto is_ssl = false;
#endif
if (!req.has_header("Host")) {
if (address_family_ == AF_UNIX) {
req.headers.emplace("Host", "localhost");
} else {
req.headers.emplace(
"Host", detail::make_host_and_port_string(host_, port_, is_ssl));
}
}
#ifndef CPPHTTPLIB_NO_DEFAULT_USER_AGENT
if (!req.has_header("User-Agent")) {
auto agent = std::string("cpp-httplib/") + CPPHTTPLIB_VERSION;
req.set_header("User-Agent", agent);
}
#endif
}
inline bool WebSocketClient::connect() {
if (!is_valid_) { return false; }
shutdown_and_close();
// Check is custom IP specified for host_
// Check is custom IP specified for host_.
// host_ stays the identity used for the Host header and for SNI, while ip
// only redirects where the socket connects.
std::string ip;
auto it = addr_map_.find(host_);
if (it != addr_map_.end()) { ip = it->second; }
@@ -20852,23 +20872,19 @@ inline bool WebSocketClient::connect() {
return false;
}
#ifdef CPPHTTPLIB_SSL_ENABLED
auto is_ssl = is_ssl_;
#else
auto is_ssl = false;
#endif
Request req;
req.method = "GET";
req.path = path_;
req.headers = headers_;
prepare_default_headers(req);
std::string selected_subprotocol;
if (!detail::perform_websocket_handshake(*strm, host_, port_, is_ssl, path_,
headers_, selected_subprotocol)) {
if (!detail::perform_websocket_handshake(*strm, req, selected_subprotocol)) {
shutdown_and_close();
return false;
}
subprotocol_ = std::move(selected_subprotocol);
Request req;
req.method = "GET";
req.path = path_;
ws_ = std::unique_ptr<WebSocket>(new WebSocket(std::move(strm), req, false,
websocket_ping_interval_sec_,
websocket_max_missed_pongs_));
+138
View File
@@ -19818,6 +19818,144 @@ TEST(WebSocketTest, HostHeaderInHandshake) {
t.join();
}
// Run a handshake against a throwaway server and hand the request the server
// received back to the caller, so tests can assert on the headers the client
// actually put on the wire.
static void capture_websocket_handshake_request(
const Headers &client_headers,
std::function<void(const Request &, int port)> verify) {
Server svr;
std::mutex mtx;
Request received;
svr.WebSocket("/ws", [&](const Request &req, ws::WebSocket &ws) {
{
std::lock_guard<std::mutex> lock(mtx);
received = req;
}
std::string msg;
while (ws.read(msg)) {
ws.send(msg);
}
});
auto port = svr.bind_to_any_port("localhost");
std::thread t([&]() { svr.listen_after_bind(); });
auto se = detail::scope_exit([&] {
svr.stop();
t.join();
});
svr.wait_until_ready();
ws::WebSocketClient client("ws://localhost:" + std::to_string(port) + "/ws",
client_headers);
ASSERT_TRUE(client.connect());
ASSERT_TRUE(client.send("hello"));
std::string msg;
ASSERT_TRUE(client.read(msg));
client.close();
{
std::lock_guard<std::mutex> lock(mtx);
verify(received, port);
}
}
TEST(WebSocketTest, DefaultHeadersInHandshake) {
capture_websocket_handshake_request({}, [](const Request &req, int port) {
EXPECT_EQ("localhost:" + std::to_string(port),
req.get_header_value("Host"));
EXPECT_EQ(std::string("cpp-httplib/") + CPPHTTPLIB_VERSION,
req.get_header_value("User-Agent"));
EXPECT_FALSE(req.has_header("Accept"));
EXPECT_FALSE(req.has_header("Accept-Encoding"));
EXPECT_FALSE(req.has_header("Content-Length"));
});
}
TEST(WebSocketTest, UserHeadersOverrideGeneratedOnesInHandshake) {
capture_websocket_handshake_request(
{{"Host", "example.com"},
{"User-Agent", "custom-agent"},
{"X-Custom", "value"}},
[](const Request &req, int) {
EXPECT_EQ("example.com", req.get_header_value("Host"));
EXPECT_EQ(1U, req.get_header_value_count("Host"));
EXPECT_EQ("custom-agent", req.get_header_value("User-Agent"));
EXPECT_EQ(1U, req.get_header_value_count("User-Agent"));
EXPECT_EQ("value", req.get_header_value("X-Custom"));
});
}
TEST(WebSocketTest, MandatoryHeadersInHandshakeAreEnforced) {
capture_websocket_handshake_request(
{{"Upgrade", "bogus"},
{"Connection", "close"},
{"Sec-WebSocket-Key", "AAAAAAAAAAAAAAAAAAAAAA=="},
{"Sec-WebSocket-Version", "8"}},
[](const Request &req, int) {
EXPECT_EQ("websocket", req.get_header_value("Upgrade"));
EXPECT_EQ(1U, req.get_header_value_count("Upgrade"));
EXPECT_EQ("Upgrade", req.get_header_value("Connection"));
EXPECT_EQ(1U, req.get_header_value_count("Connection"));
EXPECT_EQ("13", req.get_header_value("Sec-WebSocket-Version"));
EXPECT_EQ(1U, req.get_header_value_count("Sec-WebSocket-Version"));
EXPECT_EQ(1U, req.get_header_value_count("Sec-WebSocket-Key"));
EXPECT_NE("AAAAAAAAAAAAAAAAAAAAAA==",
req.get_header_value("Sec-WebSocket-Key"));
});
}
TEST(WebSocketTest, HostHeaderOverUnixSocket) {
// The socket path doubles as the URL host, so it must not contain '/'.
const char *shard = getenv("GTEST_SHARD_INDEX");
const std::string sock_path =
shard ? std::string("httplib-ws-") + shard + ".sock"
: std::string("httplib-ws.sock");
std::remove(sock_path.c_str());
Server svr;
std::mutex mtx;
std::string received_host;
svr.WebSocket("/ws", [&](const Request &req, ws::WebSocket &ws) {
{
std::lock_guard<std::mutex> lock(mtx);
received_host = req.get_header_value("Host");
}
std::string msg;
while (ws.read(msg)) {
ws.send(msg);
}
});
svr.set_address_family(AF_UNIX);
std::thread t([&]() { ASSERT_TRUE(svr.listen(sock_path, 80)); });
auto se = detail::scope_exit([&] {
svr.stop();
t.join();
std::remove(sock_path.c_str());
});
svr.wait_until_ready();
ws::WebSocketClient client("ws://" + sock_path + "/ws");
client.set_address_family(AF_UNIX);
ASSERT_TRUE(client.connect());
ASSERT_TRUE(client.send("hello"));
std::string msg;
ASSERT_TRUE(client.read(msg));
client.close();
{
std::lock_guard<std::mutex> lock(mtx);
// There is no host:port for a Unix socket, so the same "localhost"
// placeholder the HTTP client uses is expected.
EXPECT_EQ("localhost", received_host);
}
}
#ifdef CPPHTTPLIB_OPENSSL_SUPPORT
class WebSocketSSLIntegrationTest : public ::testing::Test {
protected: