Prechádzať zdrojové kódy

reject an unoffered subprotocol in read_websocket_upgrade_response (#2595)

* reject an unoffered subprotocol in the ws client handshake

* test the subprotocol check through WebSocketClient so the split build compiles

* Trim comments in the subprotocol check

---------

Co-authored-by: yhirose <yuji.hirose.bug@gmail.com>
metsw24-max 3 dní pred
rodič
commit
4fcbc08f2d
2 zmenil súbory, kde vykonal 71 pridanie a 4 odobranie
  1. 25 4
      httplib.h
  2. 46 0
      test/test.cc

+ 25 - 4
httplib.h

@@ -8272,9 +8272,11 @@ struct WebSocketUpgradeResponse {
   std::string selected_subprotocol;
 };
 
-inline bool read_websocket_upgrade_response(Stream &strm,
-                                            const std::string &expected_accept,
-                                            WebSocketUpgradeResponse &upgrade) {
+inline bool
+read_websocket_upgrade_response(Stream &strm,
+                                const std::string &expected_accept,
+                                const std::string &offered_subprotocols,
+                                WebSocketUpgradeResponse &upgrade) {
   // Read status line
   const auto bufsiz = 2048;
   char buf[bufsiz];
@@ -8331,6 +8333,22 @@ inline bool read_websocket_upgrade_response(Stream &strm,
     upgrade.selected_subprotocol = proto_it->second;
   }
 
+  // Verify the subprotocol is one the client offered (RFC 6455 4.1)
+  if (!upgrade.selected_subprotocol.empty()) {
+    auto was_offered = false;
+    split(offered_subprotocols.data(),
+          offered_subprotocols.data() + offered_subprotocols.size(), ',',
+          [&](const char *b, const char *e) {
+            if (std::string(b, e) == upgrade.selected_subprotocol) {
+              was_offered = true;
+            }
+          });
+    if (!was_offered) {
+      upgrade.error = Error::WebSocketHandshake;
+      return false;
+    }
+  }
+
   return true;
 }
 
@@ -10328,7 +10346,10 @@ inline bool perform_websocket_handshake(Stream &strm, Request &req,
 
   // Verify 101 response and Sec-WebSocket-Accept header
   auto expected_accept = websocket_accept_key(client_key);
-  return read_websocket_upgrade_response(strm, expected_accept, upgrade);
+  auto offered_subprotocols =
+      get_combined_header_value(req.headers, "Sec-WebSocket-Protocol");
+  return read_websocket_upgrade_response(strm, expected_accept,
+                                         offered_subprotocols, upgrade);
 }
 
 inline bool is_ip_address(const std::string &host) {

+ 46 - 0
test/test.cc

@@ -25001,6 +25001,52 @@ TEST(WebSocketTest, ClientRejectsResponseWithoutUpgradeToken) {
   EXPECT_FALSE(client.is_open());
 }
 
+TEST(WebSocketTest, ClientRejectsUnofferedSubprotocol) {
+  Server svr;
+  svr.Get("/ws", [](const Request &req, Response &res) {
+    res.status = StatusCode::SwitchingProtocol_101;
+    res.set_header("Upgrade", "websocket");
+    res.set_header("Connection", "Upgrade");
+    res.set_header("Sec-WebSocket-Accept",
+                   detail::websocket_accept_key(
+                       req.get_header_value("Sec-WebSocket-Key")));
+    res.set_header("Sec-WebSocket-Protocol", "admin");
+  });
+
+  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();
+
+  const auto url = "ws://localhost:" + std::to_string(port) + "/ws";
+
+  // Server selects a subprotocol the client never offered
+  {
+    Headers headers = {{"Sec-WebSocket-Protocol", "chat"}};
+    ws::WebSocketClient client(url, headers);
+
+    auto res = client.connect();
+    EXPECT_FALSE(res);
+    EXPECT_EQ(Error::WebSocketHandshake, res.error());
+    EXPECT_FALSE(client.is_open());
+    EXPECT_TRUE(client.subprotocol().empty());
+  }
+
+  // Client offered none but the server named one anyway
+  {
+    ws::WebSocketClient client(url);
+
+    auto res = client.connect();
+    EXPECT_FALSE(res);
+    EXPECT_EQ(Error::WebSocketHandshake, res.error());
+    EXPECT_FALSE(client.is_open());
+    EXPECT_TRUE(client.subprotocol().empty());
+  }
+}
+
 TEST(WebSocketTest, HostHeaderOverUnixSocket) {
   // The socket path doubles as the URL host, so it must not contain '/'.
   const char *shard = getenv("GTEST_SHARD_INDEX");