Bläddra i källkod

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.
yhirose 2 veckor sedan
förälder
incheckning
4cb363e3f2
2 ändrade filer med 89 tillägg och 4 borttagningar
  1. 11 4
      httplib.h
  2. 78 0
      test/test.cc

+ 11 - 4
httplib.h

@@ -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);

+ 78 - 0
test/test.cc

@@ -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;