mirror of
https://github.com/yhirose/cpp-httplib.git
synced 2026-10-01 05:02:29 +07:00
Run pre_request_handler for WebSocket routes
The WebSocket upgrade path matched the route and switched protocols without setting req.matched_route or calling pre_request_handler, so a check placed there (e.g. authentication) never ran for WebSocket routes. Set matched_route and run the handler before the upgrade; if it handles the request, reply with a regular HTTP response instead of 101. Also write rejected upgrade responses (from pre_routing_handler too) with write_response_with_content, so they carry Content-Length. Without it, a client reading the body waited until the keep-alive timeout.
This commit is contained in:
@@ -14411,18 +14411,25 @@ Server::process_request(Stream &strm, const std::string &remote_addr,
|
||||
};
|
||||
|
||||
// WebSocket upgrade
|
||||
// Check pre_routing_handler_ before upgrading so that authentication
|
||||
// and other middleware can reject the request with an HTTP response
|
||||
// (e.g., 401) before the protocol switches.
|
||||
// Run pre_routing_handler_ and pre_request_handler_ before upgrading so
|
||||
// that authentication and other middleware can reject the request with an
|
||||
// HTTP response (e.g., 401) before the protocol switches.
|
||||
if (detail::is_websocket_upgrade(req)) {
|
||||
if (pre_routing_handler_ &&
|
||||
pre_routing_handler_(req, res) == HandlerResponse::Handled) {
|
||||
if (res.status == -1) { res.status = StatusCode::OK_200; }
|
||||
return write_response(strm, close_connection, req, res);
|
||||
return write_response_with_content(strm, close_connection, req, res);
|
||||
}
|
||||
// Find matching WebSocket handler
|
||||
for (const auto &entry : websocket_handlers_) {
|
||||
if (entry.matcher->match(req)) {
|
||||
req.matched_route = entry.matcher->pattern();
|
||||
if (pre_request_handler_ &&
|
||||
pre_request_handler_(req, res) == HandlerResponse::Handled) {
|
||||
if (res.status == -1) { res.status = StatusCode::OK_200; }
|
||||
return write_response_with_content(strm, close_connection, req, res);
|
||||
}
|
||||
|
||||
// Compute accept key
|
||||
auto client_key = req.get_header_value("Sec-WebSocket-Key");
|
||||
auto accept_key = detail::websocket_accept_key(client_key);
|
||||
|
||||
@@ -23462,6 +23462,20 @@ TEST(WebSocketPreRoutingTest, RejectWithoutAuth) {
|
||||
ws::WebSocketClient client1("ws://localhost:" + std::to_string(port) + "/ws");
|
||||
EXPECT_FALSE(client1.connect());
|
||||
|
||||
// The rejection is framed like any other HTTP response
|
||||
{
|
||||
Client cli("localhost", port);
|
||||
Headers headers = {{"Upgrade", "websocket"},
|
||||
{"Connection", "Upgrade"},
|
||||
{"Sec-WebSocket-Key", "dGhlIHNhbXBsZSBub25jZQ=="},
|
||||
{"Sec-WebSocket-Version", "13"}};
|
||||
auto res = cli.Get("/ws", headers);
|
||||
ASSERT_TRUE(res);
|
||||
EXPECT_EQ(StatusCode::Unauthorized_401, res->status);
|
||||
EXPECT_EQ("12", res->get_header_value("Content-Length"));
|
||||
EXPECT_EQ("Unauthorized", res->body);
|
||||
}
|
||||
|
||||
// With Authorization header - should succeed
|
||||
Headers headers = {{"Authorization", "Bearer token123"}};
|
||||
ws::WebSocketClient client2("ws://localhost:" + std::to_string(port) + "/ws",
|
||||
@@ -23477,6 +23491,70 @@ TEST(WebSocketPreRoutingTest, RejectWithoutAuth) {
|
||||
t.join();
|
||||
}
|
||||
|
||||
TEST(WebSocketPreRequestTest, RejectWithoutAuth) {
|
||||
Server svr;
|
||||
|
||||
std::atomic<int> pre_request_calls{0};
|
||||
std::atomic<bool> route_matched{false};
|
||||
svr.set_pre_request_handler([&](const Request &req, Response &res) {
|
||||
pre_request_calls++;
|
||||
if (req.matched_route == "/ws/:id") { route_matched = true; }
|
||||
if (req.get_header_value("Authorization") != "Bearer token123") {
|
||||
res.status = StatusCode::Unauthorized_401;
|
||||
res.set_content("Unauthorized", "text/plain");
|
||||
return Server::HandlerResponse::Handled;
|
||||
}
|
||||
return Server::HandlerResponse::Unhandled;
|
||||
});
|
||||
|
||||
std::atomic<bool> handler_called{false};
|
||||
svr.WebSocket("/ws/:id", [&](const Request &req, ws::WebSocket &ws) {
|
||||
handler_called = true;
|
||||
ws.send(req.matched_route + " " + req.path_params.at("id"));
|
||||
});
|
||||
|
||||
auto port = svr.bind_to_any_port("localhost");
|
||||
std::thread t([&]() { svr.listen_after_bind(); });
|
||||
svr.wait_until_ready();
|
||||
|
||||
// Without Authorization header - should be rejected before upgrade
|
||||
ws::WebSocketClient client1("ws://localhost:" + std::to_string(port) +
|
||||
"/ws/1");
|
||||
EXPECT_FALSE(client1.connect());
|
||||
EXPECT_FALSE(handler_called);
|
||||
EXPECT_EQ(1, pre_request_calls);
|
||||
EXPECT_TRUE(route_matched);
|
||||
|
||||
// The rejection is an ordinary HTTP response, not a protocol switch
|
||||
{
|
||||
Client cli("localhost", port);
|
||||
Headers headers = {{"Upgrade", "websocket"},
|
||||
{"Connection", "Upgrade"},
|
||||
{"Sec-WebSocket-Key", "dGhlIHNhbXBsZSBub25jZQ=="},
|
||||
{"Sec-WebSocket-Version", "13"}};
|
||||
auto res = cli.Get("/ws/1", headers);
|
||||
ASSERT_TRUE(res);
|
||||
EXPECT_EQ(StatusCode::Unauthorized_401, res->status);
|
||||
EXPECT_EQ("12", res->get_header_value("Content-Length"));
|
||||
EXPECT_EQ("Unauthorized", res->body);
|
||||
}
|
||||
EXPECT_FALSE(handler_called);
|
||||
|
||||
// With Authorization header - should succeed
|
||||
Headers headers = {{"Authorization", "Bearer token123"}};
|
||||
ws::WebSocketClient client2(
|
||||
"ws://localhost:" + std::to_string(port) + "/ws/2", headers);
|
||||
ASSERT_TRUE(client2.connect());
|
||||
std::string msg;
|
||||
ASSERT_TRUE(client2.read(msg));
|
||||
EXPECT_EQ("/ws/:id 2", msg);
|
||||
EXPECT_TRUE(handler_called);
|
||||
client2.close();
|
||||
|
||||
svr.stop();
|
||||
t.join();
|
||||
}
|
||||
|
||||
TEST(WebSocketServerTimeoutTest, HandlerSendsWhileNothingArrives) {
|
||||
Server svr;
|
||||
|
||||
|
||||
Reference in New Issue
Block a user