Просмотр исходного кода

Build open_stream's request head in memory before sending

open_stream wrote the request line and each header straight to the
socket, so a header rejected by check_and_write_headers left the request
line and the headers before it on the wire. Build them in a BufferStream
first and flush once, as write_request and the WebSocket handshake do.
yhirose 2 недель назад
Родитель
Сommit
6d59d1e2df
2 измененных файлов с 65 добавлено и 9 удалено
  1. 23 9
      httplib.h
  2. 42 0
      test/test.cc

+ 23 - 9
httplib.h

@@ -15119,16 +15119,30 @@ ClientImpl::open_stream(const std::string &method, const std::string &path,
   prepare_default_headers(req, true, content_type);
 
   auto &strm = *handle.stream_;
-  if (detail::write_request_line(strm, req.method, req.path) < 0) {
-    handle.error = Error::Write;
-    handle.response.reset();
-    return handle;
-  }
 
-  if (!detail::check_and_write_headers(strm, req.headers, header_writer_,
-                                       handle.error)) {
-    handle.response.reset();
-    return handle;
+  // Build the request line and headers in memory first, like write_request()
+  // does, so that a rejected header leaves nothing on the wire.
+  {
+    detail::BufferStream bstrm;
+
+    if (detail::write_request_line(bstrm, req.method, req.path) < 0) {
+      handle.error = Error::Write;
+      handle.response.reset();
+      return handle;
+    }
+
+    if (!detail::check_and_write_headers(bstrm, req.headers, header_writer_,
+                                         handle.error)) {
+      handle.response.reset();
+      return handle;
+    }
+
+    const auto &data = bstrm.get_buffer();
+    if (!detail::write_data(strm, data.data(), data.size())) {
+      handle.error = Error::Write;
+      handle.response.reset();
+      return handle;
+    }
   }
 
   if (!body.empty()) {

+ 42 - 0
test/test.cc

@@ -17460,6 +17460,48 @@ TEST(ClientRejectedRequestTest, DoesNotWaitForResponse) {
   detail::close_socket(srv);
 }
 
+TEST(ClientRejectedRequestTest, OpenStreamSendsNothingOnInvalidHeader) {
+  auto srv = ::socket(AF_INET, SOCK_STREAM, 0);
+  default_socket_options(srv);
+
+  sockaddr_in addr{};
+  addr.sin_family = AF_INET;
+  addr.sin_port = htons(static_cast<uint16_t>(PORT + 1));
+  ::inet_pton(AF_INET, "127.0.0.1", &addr.sin_addr);
+  ASSERT_EQ(0, ::bind(srv, reinterpret_cast<sockaddr *>(&addr), sizeof(addr)));
+  ASSERT_EQ(0, ::listen(srv, 1));
+
+  std::string received;
+  auto server_thread = std::thread([&] {
+    auto sock = ::accept(srv, nullptr, nullptr);
+    if (sock == INVALID_SOCKET) { return; }
+    detail::set_socket_opt_time(sock, SOL_SOCKET, SO_RCVTIMEO, 2, 0);
+
+    char buf[2048];
+    ssize_t n;
+    while ((n = ::recv(sock, buf, sizeof(buf), 0)) > 0) {
+      received.append(buf, static_cast<size_t>(n));
+    }
+    detail::close_socket(sock);
+  });
+
+  {
+    auto cli = Client("127.0.0.1", PORT + 1);
+
+    // "Z" sorts after the default headers, so writing straight to the socket
+    // would have sent the request line and those headers before the rejection.
+    auto handle =
+        cli.open_stream("GET", "/", Params{}, Headers{{"Z", "B\r\nEvil: 1"}});
+    EXPECT_FALSE(handle.is_valid());
+    EXPECT_EQ(Error::InvalidHeaders, handle.error);
+  }
+
+  server_thread.join();
+  detail::close_socket(srv);
+
+  EXPECT_TRUE(received.empty()) << received;
+}
+
 TEST(PathParamsTest, StaticMatch) {
   const auto pattern = "/users/all";
   detail::PathParamsMatcher matcher(pattern);