mirror of
https://github.com/yhirose/cpp-httplib.git
synced 2026-09-30 12:42:26 +07:00
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:
@@ -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
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user