diff --git a/docs/networking/advanced-networking.md b/docs/networking/advanced-networking.md index 29a21a3f88..5d691bd02d 100644 --- a/docs/networking/advanced-networking.md +++ b/docs/networking/advanced-networking.md @@ -74,7 +74,7 @@ server.use(custom_cors); server.get("/test-cors", [](const glz::request& req, glz::response& res) { res.json({ {"message", "CORS test endpoint"}, - {"origin", req.headers.count("Origin") ? req.headers.at("Origin") : "none"}, + {"origin", std::string{req.headers.first_value("Origin").value_or("none")}}, {"method", glz::to_string(req.method)}, {"timestamp", std::time(nullptr)} }); @@ -207,12 +207,12 @@ secure_ws->on_validate([](const glz::request& req) -> bool { return false; } - return validate_jwt_token(auth_header->second); + return validate_jwt_token(auth_header->value); }); secure_ws->on_open([](auto conn, const glz::request& req) { // Store user info from validated token - auto user_data = extract_user_from_token(req.headers.at("Authorization")); + auto user_data = extract_user_from_token(*req.headers.first_value("Authorization")); conn->set_user_data(std::make_shared(user_data)); conn->send_text("Authenticated successfully"); @@ -384,7 +384,7 @@ User new_user{0, "John Doe", "john@example.com"}; auto create_response = client.post_json("https://api.example.com/users", new_user); // POST with custom headers -std::unordered_map headers = { +glz::http_headers headers = { {"Authorization", "Bearer " + token}, {"Content-Type", "application/json"} }; @@ -444,7 +444,7 @@ auto timing_middleware = [](const glz::request& req, glz::response& res) { auto response_time_middleware = [](const glz::request& req, glz::response& res) { auto start_header = res.response_headers.find("X-Request-Start"); if (start_header != res.response_headers.end()) { - auto start_ms = std::stoll(start_header->second); + auto start_ms = std::stoll(start_header->value); auto now_ms = std::chrono::duration_cast( std::chrono::high_resolution_clock::now().time_since_epoch()).count(); diff --git a/docs/networking/http-client.md b/docs/networking/http-client.md index 9c523d11f9..9ddd2c78d0 100644 --- a/docs/networking/http-client.md +++ b/docs/networking/http-client.md @@ -84,7 +84,7 @@ Notes: ```cpp std::expected get( std::string_view url, - const std::unordered_map& headers = {} + const glz::http_headers& headers = {} ); ``` @@ -93,7 +93,7 @@ std::expected get( std::expected post( std::string_view url, std::string_view body, - const std::unordered_map& headers = {} + const glz::http_headers& headers = {} ); ``` @@ -103,7 +103,7 @@ template std::expected post_json( std::string_view url, const T& data, - const std::unordered_map& headers = {} + const glz::http_headers& headers = {} ); ``` @@ -120,7 +120,7 @@ All asynchronous methods come in two variants: ```cpp std::future> get_async( std::string_view url, - const std::unordered_map& headers = {} + const glz::http_headers& headers = {} ); ``` @@ -129,7 +129,7 @@ std::future> get_async( template void get_async( std::string_view url, - const std::unordered_map& headers, + const glz::http_headers& headers, CompletionHandler&& handler ); ``` @@ -141,7 +141,7 @@ void get_async( std::future> post_async( std::string_view url, std::string_view body, - const std::unordered_map& headers = {} + const glz::http_headers& headers = {} ); ``` @@ -151,7 +151,7 @@ template void post_async( std::string_view url, std::string_view body, - const std::unordered_map& headers, + const glz::http_headers& headers, CompletionHandler&& handler ); ``` @@ -164,7 +164,7 @@ template std::future> post_json_async( std::string_view url, const T& data, - const std::unordered_map& headers = {} + const glz::http_headers& headers = {} ); ``` @@ -174,7 +174,7 @@ template void post_json_async( std::string_view url, const T& data, - const std::unordered_map& headers, + const glz::http_headers& headers, CompletionHandler&& handler ); ``` @@ -197,7 +197,7 @@ struct stream_request_params_v2 { stream_read_strategy strategy{stream_read_strategy::bulk_transfer}; size_t max_buffer_size{1024 * 1024}; std::string body; - std::unordered_map headers; + glz::http_headers headers; http_connect_handler on_connect; http_disconnect_handler on_disconnect; http_data_handler on_data; @@ -257,7 +257,7 @@ The `response` object contains: ```cpp struct response { uint16_t status_code; // HTTP status code - std::unordered_map response_headers; // Response headers + glz::http_headers response_headers; // Response headers std::string response_body; // Response body }; ``` @@ -294,7 +294,7 @@ int main() { if (response) { std::cout << "Status: " << response->status_code << std::endl; - std::cout << "Content-Type: " << response->response_headers["Content-Type"] << std::endl; + std::cout << "Content-Type: " << response->response_headers.first_value("Content-Type").value_or("") << std::endl; std::cout << "Body: " << response->response_body << std::endl; } else { std::cerr << "Error: " << response.error().message() << std::endl; @@ -312,7 +312,7 @@ int main() { int main() { glz::http_client client; - std::unordered_map headers = { + glz::http_headers headers = { {"Content-Type", "text/plain"}, {"Authorization", "Bearer your-token"} }; diff --git a/docs/networking/http-examples.md b/docs/networking/http-examples.md index 69700ac77e..6940b94645 100644 --- a/docs/networking/http-examples.md +++ b/docs/networking/http-examples.md @@ -590,9 +590,8 @@ int main() { std::string username = "User" + std::to_string(std::rand() % 1000); // In a real app, you'd extract this from authentication - auto auth_header = req.headers.find("x-username"); - if (auth_header != req.headers.end()) { - username = auth_header->second; + if (auto name = req.headers.first_value("x-username")) { + username = *name; } std::cout << "WebSocket connection from " << conn->remote_address() << " (username: " << username << ")" << std::endl; @@ -815,7 +814,7 @@ auto create_simple_auth_middleware(SimpleAuthService& auth_service) { return; } - std::string auth_value = auth_header->second; + std::string auth_value{auth_header->value}; if (!auth_value.starts_with("Bearer ")) { res.status(401).json({{"error", "Bearer token required"}}); return; @@ -871,8 +870,8 @@ int main() { // Protected endpoints server.get("/api/profile", [](const glz::request& req, glz::response& res) { // User info is available from auth middleware - std::string username = res.response_headers["X-Username"]; - std::string role = res.response_headers["X-User-Role"]; + std::string username{res.response_headers.first_value("X-Username").value_or("")}; + std::string role{res.response_headers.first_value("X-User-Role").value_or("")}; res.json({ {"username", username}, @@ -882,7 +881,7 @@ int main() { }); server.get("/api/admin", [](const glz::request& req, glz::response& res) { - std::string role = res.response_headers["X-User-Role"]; + std::string role{res.response_headers.first_value("X-User-Role").value_or("")}; if (role != "admin") { res.status(403).json({{"error", "Admin access required"}}); diff --git a/docs/networking/http-rest-support.md b/docs/networking/http-rest-support.md index fdef8db0c5..272682747e 100644 --- a/docs/networking/http-rest-support.md +++ b/docs/networking/http-rest-support.md @@ -107,7 +107,7 @@ struct request { http_method method; // GET, POST, etc. std::string target; // "/users/123" std::unordered_map params; // Route parameters - std::unordered_map headers; // HTTP headers + glz::http_headers headers; // HTTP headers std::string body; // Request body std::string remote_ip; // Client IP uint16_t remote_port; // Client port @@ -119,12 +119,13 @@ struct request { ```cpp struct response { int status_code = 200; - std::unordered_map response_headers; + glz::http_headers response_headers; std::string response_body; // Fluent interface response& status(int code); - response& header(std::string_view name, std::string_view value); + response& header(std::string_view name, std::string_view value); // replaces any existing field + response& add_header(std::string_view name, std::string_view value); // appends, for Set-Cookie and friends response& body(std::string_view content); response& content_type(std::string_view type); diff --git a/docs/networking/http-router.md b/docs/networking/http-router.md index cf4c1d881f..1cbdb39910 100644 --- a/docs/networking/http-router.md +++ b/docs/networking/http-router.md @@ -355,7 +355,7 @@ router.use([](const glz::request& req, glz::response& res) { return; } - if (!validate_token(auth_header->second)) { + if (!validate_token(auth_header->value)) { res.status(403).json({{"error", "Invalid token"}}); return; } diff --git a/docs/networking/url.md b/docs/networking/url.md index f26d9b05a5..e9148a7830 100644 --- a/docs/networking/url.md +++ b/docs/networking/url.md @@ -247,9 +247,9 @@ For `application/x-www-form-urlencoded` POST requests, use `parse_urlencoded` on ```cpp server.post("/login", [](const glz::request& req, glz::response& res) { // Check content type - auto ct = req.headers.find("content-type"); - if (ct == req.headers.end() || - ct->second.find("application/x-www-form-urlencoded") == std::string::npos) { + // A media type can carry parameters like "; charset=utf-8", so match on the prefix. + auto ct = req.headers.first_value("content-type"); + if (!ct || !ct->starts_with("application/x-www-form-urlencoded")) { res.status(415).json({{"error", "Unsupported content type"}}); return; } diff --git a/include/glaze/net/cors.hpp b/include/glaze/net/cors.hpp index 97229b6526..bd1a885038 100644 --- a/include/glaze/net/cors.hpp +++ b/include/glaze/net/cors.hpp @@ -144,7 +144,7 @@ namespace glz std::string origin; auto origin_it = req.headers.find("origin"); if (origin_it != req.headers.end()) { - origin = origin_it->second; + origin = origin_it->value; } // Check if this is a preflight request (OPTIONS method with specific headers) @@ -189,7 +189,7 @@ namespace glz } else if (auto requested_headers = req.headers.find("access-control-request-headers"); requested_headers != req.headers.end()) { - res.header("Access-Control-Allow-Headers", requested_headers->second); + res.header("Access-Control-Allow-Headers", requested_headers->value); } // Add max age diff --git a/include/glaze/net/http_client.hpp b/include/glaze/net/http_client.hpp index 583d16cea8..8ca615eeb4 100644 --- a/include/glaze/net/http_client.hpp +++ b/include/glaze/net/http_client.hpp @@ -23,6 +23,7 @@ #include #include "glaze/ext/glaze_asio.hpp" +#include "glaze/net/http_headers.hpp" #include "glaze/net/http_router.hpp" #include "glaze/util/env.hpp" #include "glaze/util/itoa.hpp" @@ -260,8 +261,7 @@ namespace glz // body or re-iterating the headers map. For large idempotent PUT bodies this is // the difference between O(1) and O(N) body copies per request through the chain. inline std::string build_http_request_bytes(const std::string& method, const url_parts& url, bool use_https, - const std::string& body, - const std::unordered_map& headers) + const std::string& body, const glz::http_headers& headers) { const bool is_default_port = (!use_https && url.port == 80) || (use_https && url.port == 443); @@ -937,7 +937,7 @@ namespace glz http_error_handler on_error; std::string method{"GET"}; std::string body{}; - std::unordered_map headers{}; + glz::http_headers headers{}; http_connect_handler on_connect{}; http_disconnect_handler on_disconnect{}; std::chrono::seconds timeout{std::chrono::seconds{30}}; @@ -955,7 +955,7 @@ namespace glz stream_read_strategy strategy{stream_read_strategy::bulk_transfer}; size_t max_buffer_size{1024 * 1024}; std::string body; - std::unordered_map headers; + glz::http_headers headers; http_connect_handler on_connect; http_disconnect_handler on_disconnect; http_data_handler on_data; @@ -1097,8 +1097,7 @@ namespace glz void clear_connection_pool() { connection_pool->clear(); } // Synchronous GET request - truly synchronous, no promises/futures - std::expected get(std::string_view url, - const std::unordered_map& headers = {}) + std::expected get(std::string_view url, const glz::http_headers& headers = {}) { auto url_result = parse_url(url); if (!url_result) { @@ -1110,7 +1109,7 @@ namespace glz // Synchronous POST request - truly synchronous, no promises/futures std::expected post(std::string_view url, const std::string& body, - const std::unordered_map& headers = {}) + const glz::http_headers& headers = {}) { auto url_result = parse_url(url); if (!url_result) { @@ -1122,7 +1121,7 @@ namespace glz // Synchronous PUT request - truly synchronous, no promises/futures std::expected put(std::string_view url, const std::string& body, - const std::unordered_map& headers = {}) + const glz::http_headers& headers = {}) { auto url_result = parse_url(url); if (!url_result) { @@ -1134,8 +1133,8 @@ namespace glz // Synchronous JSON POST request template - std::expected post_json( - std::string_view url, const T& data, const std::unordered_map& headers = {}) + std::expected post_json(std::string_view url, const T& data, + const glz::http_headers& headers = {}) { std::string json_str; auto ec = glz::write_json(data, json_str); @@ -1144,15 +1143,15 @@ namespace glz } auto merged_headers = headers; - merged_headers["content-type"] = "application/json"; + merged_headers.set("Content-Type", "application/json"); return post(url, json_str, merged_headers); } // Synchronous JSON PUT request template - std::expected put_json( - std::string_view url, const T& data, const std::unordered_map& headers = {}) + std::expected put_json(std::string_view url, const T& data, + const glz::http_headers& headers = {}) { std::string json_str; auto ec = glz::write_json(data, json_str); @@ -1161,7 +1160,7 @@ namespace glz } auto merged_headers = headers; - merged_headers["content-type"] = "application/json"; + merged_headers.set("Content-Type", "application/json"); return put(url, json_str, merged_headers); } @@ -1195,8 +1194,7 @@ namespace glz // Asynchronous GET request template - void get_async(std::string_view url, const std::unordered_map& headers, - CompletionHandler&& handler) + void get_async(std::string_view url, const glz::http_headers& headers, CompletionHandler&& handler) { auto url_result = parse_url(url); if (!url_result) { @@ -1209,8 +1207,8 @@ namespace glz } // Overload for get_async without completion handler (returns future) - std::future> get_async( - std::string_view url, const std::unordered_map& headers = {}) + std::future> get_async(std::string_view url, + const glz::http_headers& headers = {}) { std::promise> promise; auto future = promise.get_future(); @@ -1225,8 +1223,8 @@ namespace glz // Asynchronous POST request template - void post_async(std::string_view url, const std::string& body, - const std::unordered_map& headers, CompletionHandler&& handler) + void post_async(std::string_view url, const std::string& body, const glz::http_headers& headers, + CompletionHandler&& handler) { auto url_result = parse_url(url); if (!url_result) { @@ -1239,9 +1237,8 @@ namespace glz } // Overload for post_async without completion handler (returns future) - std::future> post_async( - std::string_view url, const std::string& body, - const std::unordered_map& headers = {}) + std::future> post_async(std::string_view url, const std::string& body, + const glz::http_headers& headers = {}) { std::promise> promise; auto future = promise.get_future(); @@ -1256,8 +1253,8 @@ namespace glz // Async JSON POST request template - void post_json_async(std::string_view url, const T& data, - const std::unordered_map& headers, CompletionHandler&& handler) + void post_json_async(std::string_view url, const T& data, const glz::http_headers& headers, + CompletionHandler&& handler) { std::string json_str; auto ec = glz::write_json(data, json_str); @@ -1269,15 +1266,15 @@ namespace glz } auto merged_headers = headers; - merged_headers["content-type"] = "application/json"; + merged_headers.set("Content-Type", "application/json"); post_async(url, json_str, merged_headers, std::forward(handler)); } // Overload for post_json_async without completion handler (returns future) template - std::future> post_json_async( - std::string_view url, const T& data, const std::unordered_map& headers = {}) + std::future> post_json_async(std::string_view url, const T& data, + const glz::http_headers& headers = {}) { std::promise> promise; auto future = promise.get_future(); @@ -1354,9 +1351,9 @@ namespace glz std::shared_ptr perform_stream_request( const std::string& method, const url_parts& url, const std::string& body, size_t max_buffer_size, - const std::unordered_map& headers, std::chrono::seconds timeout, - stream_read_strategy strategy, std::function status_is_error, http_data_handler on_data, - http_error_handler on_error, http_connect_handler on_connect, http_disconnect_handler on_disconnect) + const glz::http_headers& headers, std::chrono::seconds timeout, stream_read_strategy strategy, + std::function status_is_error, http_data_handler on_data, http_error_handler on_error, + http_connect_handler on_connect, http_disconnect_handler on_disconnect) { const bool use_https = (url.protocol == "https"); @@ -1477,7 +1474,7 @@ namespace glz #ifdef GLZ_ENABLE_SSL void perform_stream_ssl_handshake(const url_parts& url, const std::string& method, const std::string& body, - const std::unordered_map& headers, + const glz::http_headers& headers, std::shared_ptr connection, http_data_handler on_data, http_error_handler on_error, http_connect_handler on_connect, http_disconnect_handler on_disconnect) @@ -1509,9 +1506,8 @@ namespace glz // Needs to take the wrapped disconnect handler and use keep-alive void send_stream_request(const url_parts& url, const std::string& method, const std::string& body, - const std::unordered_map& headers, - std::shared_ptr connection, http_data_handler on_data, - http_error_handler on_error, http_connect_handler on_connect, + const glz::http_headers& headers, std::shared_ptr connection, + http_data_handler on_data, http_error_handler on_error, http_connect_handler on_connect, http_disconnect_handler on_disconnect) { const bool use_https = url.protocol == "https"; @@ -1600,8 +1596,7 @@ namespace glz std::string_view value = (value_start != std::string::npos) ? header_line.substr(value_start) : ""; - // Convert header name to lowercase for case-insensitive lookups (RFC 7230) - response_headers.response_headers[to_lower_case(name)] = std::string(value); + response_headers.response_headers.add(std::string(name), std::string(value)); } } @@ -1626,15 +1621,7 @@ namespace glz return; } - bool is_chunked = false; - auto it = response_headers.response_headers.find("transfer-encoding"); - if (it != response_headers.response_headers.end()) { - if (it->second.find("chunked") != std::string::npos) { - is_chunked = true; - } - } - - if (is_chunked) { + if (response_headers.response_headers.contains_token("transfer-encoding", "chunked")) { start_chunked_reading(connection, std::move(on_data), std::move(on_error), std::move(on_disconnect)); } @@ -1865,9 +1852,9 @@ namespace glz *connection->socket); } - std::expected perform_sync_request( - const std::string& method, const url_parts& url, const std::string& body, - const std::unordered_map& headers) + std::expected perform_sync_request(const std::string& method, const url_parts& url, + const std::string& body, + const glz::http_headers& headers) { const bool use_https = (url.protocol == "https"); @@ -1987,11 +1974,7 @@ namespace glz } // Parse headers from the view - std::unordered_map response_headers; - size_t content_length = 0; - bool has_content_length = false; - bool connection_close = false; - bool is_chunked = false; + glz::http_headers response_headers; while (!header_data.starts_with("\r\n")) { line_end = header_data.find("\r\n"); @@ -2008,26 +1991,20 @@ namespace glz size_t value_start = header_line.find_first_not_of(" \t", colon_pos + 1); std::string_view value = (value_start != std::string::npos) ? header_line.substr(value_start) : ""; - if (name.length() == 14 && glz::strncasecmp(name.data(), "Content-Length", 14) == 0) { - std::from_chars(value.data(), value.data() + value.size(), content_length); - has_content_length = true; - } - else if (name.length() == 17 && glz::strncasecmp(name.data(), "Transfer-Encoding", 17) == 0) { - if (value.find("chunked") != std::string_view::npos) { - is_chunked = true; - } - } - else if (name.length() == 10 && glz::strncasecmp(name.data(), "Connection", 10) == 0) { - if (value.find("close") != std::string_view::npos) { - connection_close = true; - } - } - - // Convert header name to lowercase for case-insensitive lookups (RFC 7230) - response_headers.emplace(to_lower_case(name), value); + response_headers.add(std::string(name), std::string(value)); } } + size_t content_length = 0; + const auto content_length_value = response_headers.first_value("Content-Length"); + const bool has_content_length = content_length_value.has_value(); + if (has_content_length) { + std::from_chars(content_length_value->data(), + content_length_value->data() + content_length_value->size(), content_length); + } + const bool is_chunked = response_headers.contains_token("Transfer-Encoding", "chunked"); + bool connection_close = response_headers.contains_token("Connection", "close"); + // Consume header data, leaving only the over-read body part. response_buffer.consume(header_bytes); @@ -2212,8 +2189,7 @@ namespace glz template void perform_request_async(const std::string& method, const url_parts& url, const std::string& body, - const std::unordered_map& headers, - CompletionHandler&& handler) + const glz::http_headers& headers, CompletionHandler&& handler) { const bool use_https = (url.protocol == "https"); @@ -2424,8 +2400,7 @@ namespace glz template void async_consume_trailers(std::shared_ptr socket_var, std::shared_ptr buffer, std::shared_ptr body, const url_parts& url, bool use_https, - int status_code, std::unordered_map response_headers, - CompletionHandler&& handler) + int status_code, glz::http_headers response_headers, CompletionHandler&& handler) { std::visit( [&, this](auto& sock) { @@ -2454,9 +2429,7 @@ namespace glz resp.response_headers = std::move(response_headers); resp.response_body = std::move(*body); - auto connection_header = resp.response_headers.find("connection"); - if (connection_header == resp.response_headers.end() || - connection_header->second.find("close") == std::string::npos) { + if (!resp.response_headers.contains_token("connection", "close")) { connection_pool->return_connection(url.host, url.port, use_https, std::move(*socket_var)); } @@ -2476,8 +2449,7 @@ namespace glz template void async_read_chunked_body(std::shared_ptr socket_var, std::shared_ptr buffer, std::shared_ptr body, const url_parts& url, bool use_https, - int status_code, std::unordered_map response_headers, - CompletionHandler&& handler) + int status_code, glz::http_headers response_headers, CompletionHandler&& handler) { // Read until we have the chunk size line std::visit( @@ -2532,8 +2504,7 @@ namespace glz template void async_read_chunk_data(std::shared_ptr socket_var, std::shared_ptr buffer, std::shared_ptr body, size_t chunk_size, const url_parts& url, - bool use_https, int status_code, - std::unordered_map response_headers, + bool use_https, int status_code, glz::http_headers response_headers, CompletionHandler&& handler) { // Check accumulated body size against limit before reading chunk data @@ -2595,8 +2566,7 @@ namespace glz template void async_read_eof_body(std::shared_ptr socket_var, std::shared_ptr buffer, const url_parts& url, bool use_https, int status_code, - std::unordered_map response_headers, - CompletionHandler&& handler) + glz::http_headers response_headers, CompletionHandler&& handler) { std::visit( [&, this](auto& sock) { @@ -2659,11 +2629,7 @@ namespace glz } // Parse all header fields from the view. - std::unordered_map response_headers; - size_t content_length = 0; - bool has_content_length = false; - bool connection_close = false; - bool is_chunked = false; + glz::http_headers response_headers; // The header section ends with an empty line ("\r\n"), which means our view will start with it. while (!header_section.starts_with("\r\n")) { line_end = header_section.find("\r\n"); @@ -2681,29 +2647,20 @@ namespace glz size_t value_start = header_line.find_first_not_of(" \t", colon_pos + 1); std::string_view value = (value_start != std::string::npos) ? header_line.substr(value_start) : ""; - // A case-insensitive comparison is more robust for header names. - if (name.size() == 14 && (name[0] == 'C' || name[0] == 'c') && - glz::strncasecmp(name.data(), "Content-Length", 14) == 0) { - std::from_chars(value.data(), value.data() + value.size(), content_length); - has_content_length = true; - } - else if (name.size() == 17 && (name[0] == 'T' || name[0] == 't') && - glz::strncasecmp(name.data(), "Transfer-Encoding", 17) == 0) { - if (value.find("chunked") != std::string_view::npos) { - is_chunked = true; - } - } - else if (name.size() == 10 && (name[0] == 'C' || name[0] == 'c') && - glz::strncasecmp(name.data(), "Connection", 10) == 0) { - if (value.find("close") != std::string_view::npos) { - connection_close = true; - } - } - // Convert header name to lowercase for case-insensitive lookups (RFC 7230) - response_headers.emplace(to_lower_case(name), value); + response_headers.add(std::string(name), std::string(value)); } } + size_t content_length = 0; + const auto content_length_value = response_headers.first_value("Content-Length"); + const bool has_content_length = content_length_value.has_value(); + if (has_content_length) { + std::from_chars(content_length_value->data(), content_length_value->data() + content_length_value->size(), + content_length); + } + const bool is_chunked = response_headers.contains_token("Transfer-Encoding", "chunked"); + const bool connection_close = response_headers.contains_token("Connection", "close"); + // Consume the entire header block from the streambuf. // This efficiently discards the header data we've just parsed, leaving only body data. buffer->consume(header_size); @@ -2764,10 +2721,7 @@ namespace glz // Pool the connection unless the server asked to close it. A truncated or // peer-closed connection already returned above, so anything reaching here // is a complete response on a still-open socket. - auto connection_header = resp.response_headers.find("connection"); - const bool server_wants_close = connection_header != resp.response_headers.end() && - connection_header->second.find("close") != std::string::npos; - if (server_wants_close) { + if (resp.response_headers.contains_token("connection", "close")) { detail::close_socket(*socket_var, connection_pool->graceful_ssl_shutdown()); } else { diff --git a/include/glaze/net/http_router.hpp b/include/glaze/net/http_router.hpp index cb836f8840..b7d506d7f0 100644 --- a/include/glaze/net/http_router.hpp +++ b/include/glaze/net/http_router.hpp @@ -17,6 +17,7 @@ #include "glaze/json/generic.hpp" #include "glaze/net/http.hpp" +#include "glaze/net/http_headers.hpp" #include "glaze/net/url.hpp" #include "glaze/util/key_transformers.hpp" @@ -48,7 +49,7 @@ namespace glz std::string path{}; // Path component only (without query string) std::unordered_map params{}; // Path parameters (e.g., :id) std::unordered_map query{}; // Query parameters (e.g., ?limit=10) - std::unordered_map headers{}; + glz::http_headers headers{}; std::string body{}; std::string remote_ip{}; uint16_t remote_port{}; @@ -75,7 +76,7 @@ namespace glz }; int status_code = 200; - std::unordered_map response_headers{}; + glz::http_headers response_headers{}; std::string response_body{}; uint8_t user_headers_set{}; @@ -85,35 +86,32 @@ namespace glz return *this; } + // Replaces any existing field with this name. inline response& header(std::string_view name, std::string_view value) { - // RFC 7230 3.2: a field-name or field-value carrying CR or LF would - // terminate the field on the wire, letting attacker-influenced data - // inject extra headers or a body (CWE-113). Reject such a field where it - // is set so it never enters the map and, crucially, the default-header - // bookkeeping below is skipped: otherwise a dropped Content-Length or - // Connection would still suppress its auto-generated counterpart and - // leave the message unframed. The wire serializers keep an independent - // drop as a backstop for headers that bypass this setter. - if (header_field_has_crlf(name, value)) [[unlikely]] { + if (reject_unwritable_field(name, value)) [[unlikely]] { return *this; } - // Convert header name to lowercase for case-insensitive lookups (RFC 7230) - std::string key(name); - for (auto& c : key) c = ascii_tolower(c); + response_headers.set(std::string(name), std::string(value)); + return *this; + } - // Track which default headers the user has set - if (key == "content-length") - user_headers_set |= has_content_length; - else if (key == "date") - user_headers_set |= has_date; - else if (key == "server") - user_headers_set |= has_server; - else if (key == "connection") - user_headers_set |= has_connection; + // Appends instead of replacing, for names that can repeat like Set-Cookie. + // Content-Length and Transfer-Encoding are replaced regardless: a second one leaves + // the body length ambiguous and opens response smuggling (RFC 9112 6.3). + inline response& add_header(std::string_view name, std::string_view value) + { + if (reject_unwritable_field(name, value)) [[unlikely]] { + return *this; + } - response_headers[std::move(key)] = std::string(value); + if (frames_the_body(name)) [[unlikely]] { + response_headers.set(std::string(name), std::string(value)); + } + else { + response_headers.add(std::string(name), std::string(value)); + } return *this; } @@ -168,7 +166,7 @@ namespace glz user_headers_set = 0; } - inline response& content_type(std::string_view type) { return header("content-type", type); } + inline response& content_type(std::string_view type) { return header("Content-Type", type); } // JSON response helper using Glaze template @@ -181,6 +179,38 @@ namespace glz } return *this; } + + private: + [[nodiscard]] static bool frames_the_body(std::string_view name) noexcept + { + return glz::striequal(name, "content-length") || glz::striequal(name, "transfer-encoding"); + } + + // RFC 7230 3.2: a field-name or field-value carrying CR or LF would + // terminate the field on the wire, letting attacker-influenced data + // inject extra headers or a body (CWE-113). Reject such a field where it + // is set so it never enters the container and, crucially, the default-header + // bookkeeping is skipped: otherwise a dropped Content-Length or + // Connection would still suppress its auto-generated counterpart and + // leave the message unframed. The wire serializers keep an independent + // drop as a backstop for headers that bypass these setters. + [[nodiscard]] bool reject_unwritable_field(std::string_view name, std::string_view value) noexcept + { + if (header_field_has_crlf(name, value)) { + return true; + } + + if (glz::striequal(name, "content-length")) + user_headers_set |= has_content_length; + else if (glz::striequal(name, "date")) + user_headers_set |= has_date; + else if (glz::striequal(name, "server")) + user_headers_set |= has_server; + else if (glz::striequal(name, "connection")) + user_headers_set |= has_connection; + + return false; + } }; using handler = std::function; diff --git a/include/glaze/net/http_server.hpp b/include/glaze/net/http_server.hpp index 8ff144ace9..c9e3eeea0c 100644 --- a/include/glaze/net/http_server.hpp +++ b/include/glaze/net/http_server.hpp @@ -26,6 +26,7 @@ #include "glaze/ext/glaze_asio.hpp" #include "glaze/net/cors.hpp" #include "glaze/net/http.hpp" +#include "glaze/net/http_headers.hpp" #include "glaze/net/http_router.hpp" #include "glaze/net/openapi.hpp" #include "glaze/net/websocket_connection.hpp" @@ -950,7 +951,7 @@ namespace glz virtual ~streaming_connection_interface() = default; // Send initial headers for streaming response - virtual void send_headers(int status_code, const std::unordered_map& headers = {}, + virtual void send_headers(int status_code, const glz::http_headers& headers = {}, data_sent_handler handler = {}) = 0; // Send a chunk of data @@ -1008,8 +1009,7 @@ namespace glz {} // Send initial headers for streaming response - void send_headers(int status_code, const std::unordered_map& headers = {}, - data_sent_handler handler = {}) override + void send_headers(int status_code, const glz::http_headers& headers = {}, data_sent_handler handler = {}) override { if (is_headers_sent_) return; is_headers_sent_ = true; @@ -1039,9 +1039,10 @@ namespace glz // dropped Transfer-Encoding/Connection would otherwise leave the stream // unframed (no chunked framing, no keep-alive signal). Treat a // present-but-invalid header as absent, matching what was written. - const auto present_and_valid = [&](const char* key) { - const auto it = headers.find(key); - return it != headers.end() && !header_field_has_crlf(it->first, it->second); + const auto present_and_valid = [&](std::string_view key) { + return std::ranges::any_of(headers.fields(key), [](const http_header& field) { + return !header_field_has_crlf(field.name, field.value); + }); }; if (!present_and_valid("transfer-encoding")) { @@ -1283,8 +1284,7 @@ namespace glz streaming_response(std::shared_ptr conn) : stream(std::move(conn)) {} // Send headers and start streaming - streaming_response& start_stream(int status_code = 200, - const std::unordered_map& headers = {}) + streaming_response& start_stream(int status_code = 200, const glz::http_headers& headers = {}) { if (stream) { stream->send_headers(status_code, headers); @@ -2555,7 +2555,7 @@ namespace glz std::string_view data(conn->read_buf.data(), conn->buf_len); parse_result result; - // Clear headers for re-parse safety (preserves bucket allocation) + // Clear headers for re-parse safety (preserves capacity) conn->request_.headers.clear(); // --- Parse request line --- @@ -2650,15 +2650,13 @@ namespace glz if (const auto last = value_sv.find_last_not_of(" \t"); last != std::string_view::npos) { value_sv.remove_suffix(value_sv.size() - (last + 1)); } - std::string key(name_sv); - for (auto& c : key) c = ascii_tolower(c); // RFC 7230 3.3.2: multiple Content-Length fields with differing values make the - // body framing unrecoverable. The header map keeps last-wins, so a second - // Content-Length would silently override the first; a proxy that frames the body - // by the first value then desyncs from this server (CL.CL request smuggling). + // body framing unrecoverable. Lookups resolve to the first field, so a second + // Content-Length carrying another length would leave a proxy that frames the body + // by the later value desynced from this server (CL.CL request smuggling). // Reject with 400 and close. Identical repeats are tolerated. - if (key == "content-length") { - if (auto existing = headers.find(key); existing != headers.end() && existing->second != value_sv) { + if (glz::striequal(name_sv, "content-length")) { + if (auto existing = headers.find(name_sv); existing != headers.end() && existing->value != value_sv) { result.status = parse_status::error; send_error_response_with_close(conn, 400, "Bad Request"); return result; @@ -2666,7 +2664,7 @@ namespace glz } // Duplicate or invalid Host fields are errors for any HTTP/1.x request - if (key == "host") { + if (glz::striequal(name_sv, "host")) { if (host_header_parsed || !detail::is_valid_authority(value_sv)) { result.status = parse_status::error; send_error_response_with_close(conn, 400, "Bad Request"); @@ -2675,7 +2673,7 @@ namespace glz host_header_parsed = true; } - headers[std::move(key)] = std::string(value_sv); + headers.add(std::string(name_sv), std::string(value_sv)); } pos = line_end + 2; @@ -2703,14 +2701,14 @@ namespace glz // keep-alive loop parses it as a second, smuggled request (TE/CL desync). // RFC 7230 3.3.1 requires answering an undecodable Transfer-Encoding with // 501 and closing the connection. - if (headers.find("transfer-encoding") != headers.end()) { + if (headers.contains("transfer-encoding")) { send_error_response_with_close(conn, 501, "Not Implemented"); return; } std::size_t content_length = 0; if (auto it = headers.find("content-length"); it != headers.end()) { - const auto& cl = it->second; + const auto& cl = it->value; auto [ptr, ec] = std::from_chars(cl.data(), cl.data() + cl.size(), content_length); if (ec != std::errc{} || ptr != cl.data() + cl.size()) { send_error_response_with_close(conn, 400, "Bad Request"); @@ -2758,19 +2756,7 @@ namespace glz } } - // Case-insensitive substring search - static bool ci_contains(std::string_view haystack, std::string_view needle) - { - if (needle.size() > haystack.size()) return false; - for (size_t i = 0; i <= haystack.size() - needle.size(); ++i) { - if (glz::striequal(haystack.substr(i, needle.size()), needle)) { - return true; - } - } - return false; - } - - inline bool determine_keep_alive(const std::unordered_map& headers, bool is_http_11, + inline bool determine_keep_alive(const glz::http_headers& headers, bool is_http_11, std::shared_ptr conn) { // If server has keep-alive disabled, always close @@ -2784,15 +2770,11 @@ namespace glz return false; } - // Check client's Connection header (case-insensitive) - auto conn_header_it = headers.find("connection"); - if (conn_header_it != headers.end()) { - if (ci_contains(conn_header_it->second, "close")) { - return false; - } - if (ci_contains(conn_header_it->second, "keep-alive")) { - return true; - } + if (headers.contains_token("connection", "close")) { + return false; + } + if (headers.contains_token("connection", "keep-alive")) { + return true; } // Default behavior based on HTTP version @@ -2891,7 +2873,7 @@ namespace glz if (auto request_method_it = request.headers.find("access-control-request-method"); request_method_it != request.headers.end()) { has_request_method_header = true; - std::string method_token{request_method_it->second}; + std::string method_token{request_method_it->value}; auto trim_pos = method_token.find_last_not_of(" \t\r\n"); if (trim_pos != std::string::npos) { method_token.erase(trim_pos + 1); @@ -3173,16 +3155,9 @@ namespace glz send_error_response_with_conn(conn, status_code, message); } - inline bool is_websocket_upgrade(const std::unordered_map& headers) + inline bool is_websocket_upgrade(const glz::http_headers& headers) { - auto upgrade_it = headers.find("upgrade"); - if (upgrade_it == headers.end()) return false; - - auto connection_it = headers.find("connection"); - if (connection_it == headers.end()) return false; - - if (!ci_contains(upgrade_it->second, "websocket")) return false; - return ci_contains(connection_it->second, "upgrade"); + return headers.contains_token("upgrade", "websocket") && headers.contains_token("connection", "upgrade"); } inline void handle_websocket_upgrade_with_conn(std::shared_ptr conn) diff --git a/include/glaze/net/websocket_client.hpp b/include/glaze/net/websocket_client.hpp index ecb5f85c16..aceaace694 100644 --- a/include/glaze/net/websocket_client.hpp +++ b/include/glaze/net/websocket_client.hpp @@ -12,7 +12,9 @@ #include #include "glaze/net/http_client.hpp" +#include "glaze/net/http_headers.hpp" #include "glaze/net/websocket_connection.hpp" +#include "glaze/util/compare.hpp" #include "glaze/util/itoa.hpp" namespace glz @@ -63,7 +65,7 @@ namespace glz std::shared_ptr ssl_socket_; #endif mutable std::mutex request_headers_mutex_; - std::vector> request_headers_; + glz::http_headers request_headers_; std::atomic last_header_validation_error_{header_validation_error::none}; size_t max_message_size{1024 * 1024 * 16}; // 16 MB limit @@ -73,26 +75,16 @@ namespace glz explicit impl(std::shared_ptr context) : ctx(std::move(context)) {} - static bool header_name_equal(std::string_view lhs, std::string_view rhs) - { - if (lhs.size() != rhs.size()) return false; - return std::equal(lhs.begin(), lhs.end(), rhs.begin(), [](char a, char b) { - return std::tolower(static_cast(a)) == std::tolower(static_cast(b)); - }); - } - static bool header_name_starts_with(std::string_view value, std::string_view prefix) { if (value.size() < prefix.size()) return false; - return std::equal(prefix.begin(), prefix.end(), value.begin(), [](char a, char b) { - return std::tolower(static_cast(a)) == std::tolower(static_cast(b)); - }); + return glz::striequal(value.substr(0, prefix.size()), prefix); } static bool is_reserved_handshake_header(std::string_view name) { - return header_name_equal(name, "Host") || header_name_equal(name, "Upgrade") || - header_name_equal(name, "Connection") || header_name_starts_with(name, "Sec-WebSocket-"); + return glz::striequal(name, "Host") || glz::striequal(name, "Upgrade") || + glz::striequal(name, "Connection") || header_name_starts_with(name, "Sec-WebSocket-"); } // Field-name and field-value validity (tchar names, no CR/LF/CTL/DEL @@ -127,7 +119,7 @@ namespace glz return true; } - std::vector> request_headers_snapshot() const + glz::http_headers request_headers_snapshot() const { std::lock_guard lock(request_headers_mutex_); return request_headers_; @@ -142,14 +134,7 @@ namespace glz } std::lock_guard lock(request_headers_mutex_); - for (auto& [existing_name, existing_value] : request_headers_) { - if (header_name_equal(existing_name, name)) { - existing_value = std::string(value); - last_header_validation_error_.store(header_validation_error::none, std::memory_order_relaxed); - return true; - } - } - request_headers_.emplace_back(std::string(name), std::string(value)); + request_headers_.set(std::string(name), std::string(value)); last_header_validation_error_.store(header_validation_error::none, std::memory_order_relaxed); return true; } @@ -363,95 +348,86 @@ namespace glz auto response_buf = std::make_shared(max_handshake_size); std::weak_ptr weak_self = weak_from_this(); - asio::async_read_until(*socket, *response_buf, "\r\n\r\n", - [weak_self, socket, response_buf, expected_key](std::error_code ec, std::size_t) { - auto self = weak_self.lock(); - if (!self) return; // Client was destroyed - - if (ec) { - if (self->on_error && *self->on_error) (*self->on_error)(ec); - return; - } - - std::istream response_stream(response_buf.get()); - std::string http_version; - unsigned int status_code; - std::string status_message; - - response_stream >> http_version >> status_code; - std::getline(response_stream, status_message); - - if (!response_stream || status_code != 101) { - if (self->on_error && *self->on_error) - (*self->on_error)(std::make_error_code(std::errc::protocol_error)); - return; - } - - // Parse headers to verify upgrade and accept key - std::string header; - bool upgrade_websocket = false; - bool connection_upgrade = false; - bool accept_key_valid = false; - - std::string expected_accept = ws_util::generate_accept_key(expected_key); - - while (std::getline(response_stream, header) && header != "\r") { - if (!header.empty() && header.back() == '\r') header.pop_back(); - - auto colon = header.find(':'); - if (colon != std::string::npos) { - std::string name = header.substr(0, colon); - std::string value = header.substr(colon + 1); - - while (!value.empty() && (value.front() == ' ' || value.front() == '\t')) - value.erase(0, 1); - while (!value.empty() && (value.back() == ' ' || value.back() == '\t')) - value.pop_back(); - - if (strncasecmp(name.c_str(), "Upgrade", 7) == 0 && - ws_util::header_contains(value, "websocket")) { - upgrade_websocket = true; - } - else if (strncasecmp(name.c_str(), "Connection", 10) == 0 && - ws_util::header_contains(value, "upgrade")) { - connection_upgrade = true; - } - else if (strncasecmp(name.c_str(), "Sec-WebSocket-Accept", 20) == 0) { - if (value == expected_accept) accept_key_valid = true; - } - } - } - - if (!upgrade_websocket || !connection_upgrade || !accept_key_valid) { - if (self->on_error && *self->on_error) - (*self->on_error)(std::make_error_code(std::errc::protocol_error)); - return; - } - - // Handshake successful. Transfer socket to websocket_connection. - auto ws_conn = std::make_shared>(socket); - ws_conn->set_client_mode(true); - ws_conn->set_max_message_size(self->max_message_size); - - if (self->on_message && *self->on_message) ws_conn->on_message(*self->on_message); - if (self->on_close && *self->on_close) ws_conn->on_close(*self->on_close); - if (self->on_error && *self->on_error) ws_conn->on_error(*self->on_error); - - { - std::lock_guard lock(self->connection_mutex); - self->connection = ws_conn; - } - - if (self->on_open && *self->on_open) (*self->on_open)(); - - if (response_buf->size() > 0) { - std::vector initial_data(response_buf->size()); - asio::buffer_copy(asio::buffer(initial_data), response_buf->data()); - ws_conn->set_initial_data(std::move(initial_data)); - } - - ws_conn->start_read(); - }); + asio::async_read_until( + *socket, *response_buf, "\r\n\r\n", + [weak_self, socket, response_buf, expected_key](std::error_code ec, std::size_t) { + auto self = weak_self.lock(); + if (!self) return; // Client was destroyed + + if (ec) { + if (self->on_error && *self->on_error) (*self->on_error)(ec); + return; + } + + std::istream response_stream(response_buf.get()); + std::string http_version; + unsigned int status_code; + std::string status_message; + + response_stream >> http_version >> status_code; + std::getline(response_stream, status_message); + + if (!response_stream || status_code != 101) { + if (self->on_error && *self->on_error) + (*self->on_error)(std::make_error_code(std::errc::protocol_error)); + return; + } + + // Parse headers to verify upgrade and accept key + std::string header; + glz::http_headers response_headers; + + std::string expected_accept = ws_util::generate_accept_key(expected_key); + + while (std::getline(response_stream, header) && header != "\r") { + if (!header.empty() && header.back() == '\r') header.pop_back(); + + auto colon = header.find(':'); + if (colon != std::string::npos) { + std::string name = header.substr(0, colon); + std::string value = header.substr(colon + 1); + + while (!value.empty() && (value.front() == ' ' || value.front() == '\t')) value.erase(0, 1); + while (!value.empty() && (value.back() == ' ' || value.back() == '\t')) value.pop_back(); + + response_headers.add(std::move(name), std::move(value)); + } + } + + const bool upgrade_websocket = response_headers.contains_token("Upgrade", "websocket"); + const bool connection_upgrade = response_headers.contains_token("Connection", "upgrade"); + const bool accept_key_valid = response_headers.first_value("Sec-WebSocket-Accept") == expected_accept; + + if (!upgrade_websocket || !connection_upgrade || !accept_key_valid) { + if (self->on_error && *self->on_error) + (*self->on_error)(std::make_error_code(std::errc::protocol_error)); + return; + } + + // Handshake successful. Transfer socket to websocket_connection. + auto ws_conn = std::make_shared>(socket); + ws_conn->set_client_mode(true); + ws_conn->set_max_message_size(self->max_message_size); + + if (self->on_message && *self->on_message) ws_conn->on_message(*self->on_message); + if (self->on_close && *self->on_close) ws_conn->on_close(*self->on_close); + if (self->on_error && *self->on_error) ws_conn->on_error(*self->on_error); + + { + std::lock_guard lock(self->connection_mutex); + self->connection = ws_conn; + } + + if (self->on_open && *self->on_open) (*self->on_open)(); + + if (response_buf->size() > 0) { + std::vector initial_data(response_buf->size()); + asio::buffer_copy(asio::buffer(initial_data), response_buf->data()); + ws_conn->set_initial_data(std::move(initial_data)); + } + + ws_conn->start_read(); + }); } void send_text(std::string_view msg) diff --git a/include/glaze/net/websocket_connection.hpp b/include/glaze/net/websocket_connection.hpp index f17ccce04a..d20cfd9fd4 100644 --- a/include/glaze/net/websocket_connection.hpp +++ b/include/glaze/net/websocket_connection.hpp @@ -18,6 +18,8 @@ #include #include +#include "glaze/util/compare.hpp" + // Optional OpenSSL support - detected at compile time #if defined(GLZ_ENABLE_OPENSSL) && __has_include() #include @@ -250,40 +252,6 @@ namespace glz return glz::write_base64(std::string_view{reinterpret_cast(hash), sizeof(hash)}); } - // Check if a string contains a value (case-insensitive, comma-separated) - inline bool header_contains(std::string_view header, std::string_view value) - { - while (!header.empty()) { - // Skip whitespace - while (!header.empty() && (header.front() == ' ' || header.front() == '\t')) { - header.remove_prefix(1); - } - - if (header.empty()) break; - - // Find the end of this token - auto comma_pos = header.find(','); - std::string_view token = header.substr(0, comma_pos); - - // Remove trailing whitespace from token - while (!token.empty() && (token.back() == ' ' || token.back() == '\t')) { - token.remove_suffix(1); - } - - // Case-insensitive comparison - if (token.size() == value.size() && - std::equal(token.begin(), token.end(), value.begin(), value.end(), [](char a, char b) { - return std::tolower(static_cast(a)) == std::tolower(static_cast(b)); - })) { - return true; - } - - if (comma_pos == std::string_view::npos) break; - header.remove_prefix(comma_pos + 1); - } - - return false; - } } // Forward declarations @@ -612,25 +580,18 @@ namespace glz inline void perform_handshake(const request& req) { // Validate WebSocket upgrade request - auto it = req.headers.find("upgrade"); - constexpr std::string_view websocket_str = "websocket"; - if (it == req.headers.end() || !std::equal(it->second.begin(), it->second.end(), websocket_str.begin(), - websocket_str.end(), [](char a, char b) { - return std::tolower(static_cast(a)) == - std::tolower(static_cast(b)); - })) { + if (!req.headers.contains_token("upgrade", "websocket")) { do_close(); return; } - it = req.headers.find("connection"); - if (it == req.headers.end() || !ws_util::header_contains(it->second, "upgrade")) { + if (!req.headers.contains_token("connection", "upgrade")) { do_close(); return; } - it = req.headers.find("sec-websocket-version"); - if (it == req.headers.end() || it->second != "13") { + auto it = req.headers.find("sec-websocket-version"); + if (it == req.headers.end() || it->value != "13") { do_close(); return; } @@ -649,8 +610,9 @@ namespace glz } } - // Generate accept key - std::string accept_key = ws_util::generate_accept_key(it->second); + std::string_view sec_websocket_key = it->value; + + std::string accept_key = ws_util::generate_accept_key(sec_websocket_key); // Send handshake response std::string response_str = diff --git a/tests/networking_tests/cors_test/cors_test.cpp b/tests/networking_tests/cors_test/cors_test.cpp index 1eed0b9a91..d544082ae1 100644 --- a/tests/networking_tests/cors_test/cors_test.cpp +++ b/tests/networking_tests/cors_test/cors_test.cpp @@ -3,6 +3,7 @@ #include "glaze/net/cors.hpp" +#include "glaze/net/http_headers.hpp" #include "ut/ut.hpp" using namespace ut; @@ -14,7 +15,7 @@ namespace auto middleware = glz::create_cors_middleware(config); glz::request req{}; req.method = glz::http_method::GET; - req.headers["origin"] = std::string(origin); + req.headers.set("origin", std::string(origin)); glz::response res{}; middleware(req, res); return res; @@ -29,8 +30,8 @@ suite cors_origin_tests = [] { expect(glz::is_origin_allowed(config, "https://evil.example")); auto res = run_cors(config, "https://evil.example"); - expect(res.response_headers["access-control-allow-origin"] == "*"); - expect(res.response_headers.count("access-control-allow-credentials") == 0); + expect(res.response_headers.first_value("access-control-allow-origin") == "*"); + expect(!res.response_headers.contains("access-control-allow-credentials")); }; "wildcard_with_credentials_rejects_unlisted_origin"_test = [] { @@ -40,8 +41,8 @@ suite cors_origin_tests = [] { expect(not glz::is_origin_allowed(config, "https://evil.example")); auto res = run_cors(config, "https://evil.example"); - expect(res.response_headers.count("access-control-allow-origin") == 0); - expect(res.response_headers.count("access-control-allow-credentials") == 0); + expect(!res.response_headers.contains("access-control-allow-origin")); + expect(!res.response_headers.contains("access-control-allow-credentials")); }; "credentials_allow_exact_listed_origin"_test = [] { @@ -52,8 +53,8 @@ suite cors_origin_tests = [] { expect(not glz::is_origin_allowed(config, "https://evil.example")); auto res = run_cors(config, "https://app.example"); - expect(res.response_headers["access-control-allow-origin"] == "https://app.example"); - expect(res.response_headers["access-control-allow-credentials"] == "true"); + expect(res.response_headers.first_value("access-control-allow-origin") == "https://app.example"); + expect(res.response_headers.first_value("access-control-allow-credentials") == "true"); }; "empty_origins_with_credentials_rejects"_test = [] { diff --git a/tests/networking_tests/http_chunked_test/http_chunked_test.cpp b/tests/networking_tests/http_chunked_test/http_chunked_test.cpp index 8aabdf1d11..53f1f11153 100644 --- a/tests/networking_tests/http_chunked_test/http_chunked_test.cpp +++ b/tests/networking_tests/http_chunked_test/http_chunked_test.cpp @@ -379,13 +379,13 @@ suite chunked_sync_tests = [] { auto custom = result->response_headers.find("x-custom-header"); expect(custom != result->response_headers.end()) << "Custom header should be present"; if (custom != result->response_headers.end()) { - expect(custom->second == "custom-value") << "Custom header value should match"; + expect(custom->value == "custom-value") << "Custom header value should match"; } auto te = result->response_headers.find("transfer-encoding"); expect(te != result->response_headers.end()) << "Transfer-Encoding header should be present"; if (te != result->response_headers.end()) { - expect(te->second.find("chunked") != std::string::npos) << "Transfer-Encoding should be chunked"; + expect(te->value.find("chunked") != std::string::npos) << "Transfer-Encoding should be chunked"; } } diff --git a/tests/networking_tests/http_client_ssl_test/http_client_ssl_test.cpp b/tests/networking_tests/http_client_ssl_test/http_client_ssl_test.cpp index 70bb456ff9..9d48542cb1 100644 --- a/tests/networking_tests/http_client_ssl_test/http_client_ssl_test.cpp +++ b/tests/networking_tests/http_client_ssl_test/http_client_ssl_test.cpp @@ -448,15 +448,15 @@ suite https_client_tests = [] { glz::http_client client; client.set_ssl_verify_mode(asio::ssl::verify_none); - std::unordered_map headers; - headers["X-Custom-Header"] = "CustomValue"; - headers["Authorization"] = "Bearer test-token"; + glz::http_headers headers; + headers.set("X-Custom-Header", "CustomValue"); + headers.set("Authorization", "Bearer test-token"); auto result = client.get(g_server.base_url() + "/headers", headers); expect(result.has_value()) << "HTTPS with custom headers should succeed"; if (result.has_value()) { expect(result->status_code == 200); - expect(result->response_body.find("x-custom-header") != std::string::npos); + expect(result->response_body.find("X-Custom-Header") != std::string::npos); } }; diff --git a/tests/networking_tests/http_client_test/http_client_test.cpp b/tests/networking_tests/http_client_test/http_client_test.cpp index d8aa39d814..d5a6988e88 100644 --- a/tests/networking_tests/http_client_test/http_client_test.cpp +++ b/tests/networking_tests/http_client_test/http_client_test.cpp @@ -161,7 +161,7 @@ class working_test_server response_body.append(req.body); if (auto it = req.headers.find("x-test-header"); it != req.headers.end()) { response_body.append(":"); - response_body.append(it->second); + response_body.append(it->value); } res.status(200).content_type("text/plain").body(response_body); }); @@ -174,7 +174,7 @@ class working_test_server } std::string response_body = "CT="; - response_body.append(content_type->second); + response_body.append(content_type->value); response_body.append(";BODY="); response_body.append(req.body); res.status(200).content_type("text/plain").body(response_body); @@ -326,7 +326,7 @@ class simple_test_client } std::expected options( - const std::string& url, const std::vector>& extra_headers) + const std::string& url, const glz::http_headers& extra_headers) { auto url_parts = parse_url(url); if (!url_parts) { @@ -342,7 +342,7 @@ class simple_test_client std::expected perform_request( const std::string& method, const url_parts& url, const std::string& body, - std::vector> extra_headers = {}) + glz::http_headers extra_headers = {}) { std::promise> promise; auto future = promise.get_future(); @@ -413,8 +413,7 @@ class simple_test_client std::find_if(value.rbegin(), value.rend(), [](unsigned char ch) { return !std::isspace(ch); }).base(), value.end()); - auto lower_name = glz::to_lower_case(name); - resp.response_headers[lower_name] = value; + resp.response_headers.add(std::move(name), std::move(value)); } // Read body @@ -482,7 +481,7 @@ suite working_http_tests = [] { expect(server.start()) << "Server should start\n"; simple_test_client client; - std::vector> headers = { + glz::http_headers headers = { {"Origin", "http://localhost"}, {"Access-Control-Request-Method", "GET"}, {"Access-Control-Request-Headers", "X-Test-Header"}, @@ -512,7 +511,7 @@ suite working_http_tests = [] { expect(server.start()) << "Server should start\n"; simple_test_client client; - std::vector> headers = { + glz::http_headers headers = { {"Origin", "http://app.allowed.local"}, {"Access-Control-Request-Method", "GET"}, }; @@ -523,19 +522,19 @@ suite working_http_tests = [] { auto origin_header = allowed->response_headers.find("access-control-allow-origin"); expect(origin_header != allowed->response_headers.end()) << "Allow-Origin header should be present\n"; if (origin_header != allowed->response_headers.end()) { - expect(origin_header->second == "http://app.allowed.local") + expect(origin_header->value == "http://app.allowed.local") << "Origin should be echoed for allowed pattern\n"; } } - headers[0].second = "http://special.local"; + headers.set("Origin", "http://special.local"); auto allowed_callback = client.options(server.base_url() + "/hello", headers); expect(allowed_callback.has_value()) << "Dynamic callback origin should succeed\n"; if (allowed_callback.has_value()) { expect(allowed_callback->status_code == 204); } - headers[0].second = "http://denied.local"; + headers.set("Origin", "http://denied.local"); auto denied = client.options(server.base_url() + "/hello", headers); expect(denied.has_value()) << "Request should return a response even when denied\n"; if (denied.has_value()) { @@ -571,19 +570,19 @@ suite working_http_tests = [] { auto methods_it = result->response_headers.find("access-control-allow-methods"); expect(methods_it != result->response_headers.end()) << "Allow-Methods header missing\n"; if (methods_it != result->response_headers.end()) { - expect(methods_it->second == "GET, HEAD, POST, PUT, DELETE, PATCH"); + expect(methods_it->value == "GET, HEAD, POST, PUT, DELETE, PATCH"); } auto headers_it = result->response_headers.find("access-control-allow-headers"); expect(headers_it != result->response_headers.end()); if (headers_it != result->response_headers.end()) { - expect(headers_it->second == "X-Test-Header"); + expect(headers_it->value == "X-Test-Header"); } auto max_age_it = result->response_headers.find("access-control-max-age"); expect(max_age_it != result->response_headers.end()); if (max_age_it != result->response_headers.end()) { - expect(max_age_it->second == "123"); + expect(max_age_it->value == "123"); } } @@ -611,7 +610,7 @@ suite working_http_tests = [] { auto headers_it = result->response_headers.find("access-control-allow-headers"); expect(headers_it != result->response_headers.end()); if (headers_it != result->response_headers.end()) { - expect(headers_it->second == "*") << "Expected * but got " << headers_it->second << "\n"; + expect(headers_it->value == "*") << "Expected * but got " << headers_it->value << "\n"; } } @@ -636,7 +635,7 @@ suite working_http_tests = [] { auto allow_it = result->response_headers.find("allow"); expect(allow_it != result->response_headers.end()) << "Allow header must be present\n"; if (allow_it != result->response_headers.end()) { - expect(allow_it->second.find("GET") != std::string::npos) + expect(allow_it->value.find("GET") != std::string::npos) << "Allow header should list the implemented method\n"; } } @@ -724,7 +723,7 @@ suite working_http_tests = [] { auto allow_it = put_result->response_headers.find("allow"); expect(allow_it != put_result->response_headers.end()) << "Allow header should be present\n"; if (allow_it != put_result->response_headers.end()) { - expect(allow_it->second.find("PUT") != std::string::npos) << "Allow should advertise PUT\n"; + expect(allow_it->value.find("PUT") != std::string::npos) << "Allow should advertise PUT\n"; } } @@ -782,26 +781,26 @@ suite working_http_tests = [] { auto origin_it = result->response_headers.find("access-control-allow-origin"); expect(origin_it != result->response_headers.end()) << "Allow-Origin header missing\n"; if (origin_it != result->response_headers.end()) { - expect(origin_it->second == "https://app.local") + expect(origin_it->value == "https://app.local") << "Allow-Origin should echo the specific origin, not '*', when credentials are allowed\n"; } auto credentials_it = result->response_headers.find("access-control-allow-credentials"); expect(credentials_it != result->response_headers.end()) << "Allow-Credentials should be present\n"; if (credentials_it != result->response_headers.end()) { - expect(credentials_it->second == "true"); + expect(credentials_it->value == "true"); } auto methods_it = result->response_headers.find("access-control-allow-methods"); expect(methods_it != result->response_headers.end()) << "Allow-Methods should be present\n"; if (methods_it != result->response_headers.end()) { - expect(methods_it->second == "GET, POST"); + expect(methods_it->value == "GET, POST"); } auto max_age_it = result->response_headers.find("access-control-max-age"); expect(max_age_it != result->response_headers.end()) << "Max-Age should be present\n"; if (max_age_it != result->response_headers.end()) { - expect(max_age_it->second == "7200"); + expect(max_age_it->value == "7200"); } } @@ -973,7 +972,7 @@ suite glz_http_client_tests = [] { glz::http_client client; - std::unordered_map headers{{"x-test-header", "header-value"}}; + glz::http_headers headers{{"x-test-header", "header-value"}}; auto result = client.put(server.base_url() + "/update", "payload", headers); expect(result.has_value()) << "PUT request should succeed"; @@ -997,7 +996,7 @@ suite glz_http_client_tests = [] { auto ec = glz::write_json(payload, expected_json); expect(!ec) << "Serializing payload should succeed"; - std::unordered_map extra_headers{{"x-extra", "value"}}; + glz::http_headers extra_headers{{"x-extra", "value"}}; auto result = client.put_json(server.base_url() + "/json", payload, extra_headers); expect(result.has_value()) << "PUT JSON request should succeed"; diff --git a/tests/networking_tests/http_response_splitting_test/http_response_splitting_test.cpp b/tests/networking_tests/http_response_splitting_test/http_response_splitting_test.cpp index 38bf91ae5b..42c4055370 100644 --- a/tests/networking_tests/http_response_splitting_test/http_response_splitting_test.cpp +++ b/tests/networking_tests/http_response_splitting_test/http_response_splitting_test.cpp @@ -124,7 +124,7 @@ suite client_request_serializer_crlf = [] { url.port = 80; url.path = "/"; - const std::unordered_map headers{ + const glz::http_headers headers{ {"X-Evil", "a\r\nSmuggled-Header: 1"}, {"X-Safe", "kept"}, }; @@ -210,8 +210,7 @@ suite http_response_splitting_suite = [] { const std::string response = send_raw_timed(port, payload); expect(response.find("200") != std::string::npos) << "Request should be served"; - // response::header() lowercases the field-name (RFC 7230 case-insensitive). - expect(response.find("x-echo: harmless-value") != std::string::npos) << "Benign header should round-trip"; + expect(response.find("X-Echo: harmless-value") != std::string::npos) << "Benign header should round-trip"; }; "dropping a reflected Content-Length still frames the response"_test = [&] { diff --git a/tests/networking_tests/http_server_api_tests/http_server_api_tests.cpp b/tests/networking_tests/http_server_api_tests/http_server_api_tests.cpp index ddf8a2cf50..ee85187d01 100644 --- a/tests/networking_tests/http_server_api_tests/http_server_api_tests.cpp +++ b/tests/networking_tests/http_server_api_tests/http_server_api_tests.cpp @@ -735,8 +735,8 @@ suite response_building_tests = [] { expect(&chained_res == &res) << "Response methods should return reference for chaining\n"; expect(res.status_code == 201) << "Status should be set correctly\n"; - expect(res.response_headers.at("x-custom") == "value") << "Custom header should be set\n"; - expect(res.response_headers.at("content-type") == "application/json") << "Content-Type should be set\n"; + expect(res.response_headers.first_value("x-custom") == "value") << "Custom header should be set\n"; + expect(res.response_headers.first_value("content-type") == "application/json") << "Content-Type should be set\n"; expect(res.response_body == "test body") << "Body should be set correctly\n"; }; @@ -747,7 +747,7 @@ suite response_building_tests = [] { res.json(data); expect(!res.response_body.empty()) << "JSON serialization should produce content\n"; - expect(res.response_headers.at("content-type") == "application/json") << "Should set JSON content type\n"; + expect(res.response_headers.first_value("content-type") == "application/json") << "Should set JSON content type\n"; // Verify serialization worked TestData deserialized; @@ -882,7 +882,7 @@ suite response_middleware_tests = [] { auto response_hook = [&captured_content_type](const glz::request&, const glz::response& res) { auto it = res.response_headers.find("content-type"); if (it != res.response_headers.end()) { - captured_content_type = it->second; + captured_content_type = it->value; } }; @@ -1153,7 +1153,7 @@ struct raw_http_client struct http_response { int status_code = 0; - std::unordered_map headers; + glz::http_headers headers; std::string body; }; @@ -1218,11 +1218,10 @@ struct raw_http_client value.erase(0, value.find_first_not_of(" \t")); value.erase(value.find_last_not_of(" \t") + 1); - // Convert name to lowercase for easier lookup - std::transform(name.begin(), name.end(), name.begin(), [](unsigned char c) { return std::tolower(c); }); - resp.headers[name] = value; + const bool is_content_length = glz::striequal(name, "content-length"); + resp.headers.add(std::move(name), value); - if (name == "content-length") { + if (is_content_length) { content_length = std::stoul(value); } } @@ -1322,7 +1321,7 @@ struct keepalive_test_server server_.get("/echo-connection", [](const glz::request& req, glz::response& res) { auto conn_it = req.headers.find("connection"); if (conn_it != req.headers.end()) { - res.body("Connection: " + conn_it->second); + res.body("Connection: " + conn_it->value); } else { res.body("Connection: (none)"); @@ -1377,13 +1376,13 @@ suite keepalive_behavior_tests = [] { auto conn_header = resp->headers.find("connection"); expect(conn_header != resp->headers.end()) << "Connection header should be present\n"; if (conn_header != resp->headers.end()) { - expect(conn_header->second == "keep-alive") << "Connection should be keep-alive\n"; + expect(conn_header->value == "keep-alive") << "Connection should be keep-alive\n"; } auto ka_header = resp->headers.find("keep-alive"); expect(ka_header != resp->headers.end()) << "Keep-Alive header should be present\n"; if (ka_header != resp->headers.end()) { - expect(ka_header->second.find("timeout=30") != std::string::npos) << "Should contain timeout value\n"; + expect(ka_header->value.find("timeout=30") != std::string::npos) << "Should contain timeout value\n"; } } @@ -1406,7 +1405,7 @@ suite keepalive_behavior_tests = [] { auto conn_header = resp->headers.find("connection"); expect(conn_header != resp->headers.end()) << "Connection header should be present\n"; if (conn_header != resp->headers.end()) { - expect(conn_header->second == "close") << "Connection should be close\n"; + expect(conn_header->value == "close") << "Connection should be close\n"; } } @@ -1431,7 +1430,7 @@ suite keepalive_behavior_tests = [] { auto conn_header = resp->headers.find("connection"); expect(conn_header != resp->headers.end()) << "Connection header should be present\n"; if (conn_header != resp->headers.end()) { - expect(conn_header->second == "close") << "Server should respect client's close request\n"; + expect(conn_header->value == "close") << "Server should respect client's close request\n"; } } @@ -1477,7 +1476,7 @@ suite keepalive_behavior_tests = [] { if (resp1.has_value()) { auto conn_header = resp1->headers.find("connection"); if (conn_header != resp1->headers.end()) { - expect(conn_header->second == "keep-alive") << "First request should keep alive\n"; + expect(conn_header->value == "keep-alive") << "First request should keep alive\n"; } } @@ -1487,7 +1486,7 @@ suite keepalive_behavior_tests = [] { if (resp2.has_value()) { auto conn_header = resp2->headers.find("connection"); if (conn_header != resp2->headers.end()) { - expect(conn_header->second == "close") << "Second request should close (at limit)\n"; + expect(conn_header->value == "close") << "Second request should close (at limit)\n"; } } @@ -1510,8 +1509,8 @@ suite keepalive_behavior_tests = [] { auto ka_header = resp->headers.find("keep-alive"); expect(ka_header != resp->headers.end()) << "Keep-Alive header should be present\n"; if (ka_header != resp->headers.end()) { - expect(ka_header->second.find("timeout=45") != std::string::npos) << "Should contain timeout\n"; - expect(ka_header->second.find("max=100") != std::string::npos) << "Should contain max\n"; + expect(ka_header->value.find("timeout=45") != std::string::npos) << "Should contain timeout\n"; + expect(ka_header->value.find("max=100") != std::string::npos) << "Should contain max\n"; } } diff --git a/tests/networking_tests/rest_test/rest_server/rest_server.cpp b/tests/networking_tests/rest_test/rest_server/rest_server.cpp index 66cd7d0519..f030919cb8 100644 --- a/tests/networking_tests/rest_test/rest_server/rest_server.cpp +++ b/tests/networking_tests/rest_test/rest_server/rest_server.cpp @@ -277,7 +277,7 @@ int main() server.get("/test-cors", [](const glz::request& req, glz::response& res) { // The CORS middleware will automatically add the appropriate headers res.json({{"message", "CORS test endpoint"}, - {"origin", req.headers.count("origin") ? req.headers.at("origin") : "none"}, + {"origin", std::string{req.headers.first_value("origin").value_or("none")}}, {"method", glz::to_string(req.method)}}); }); diff --git a/tests/networking_tests/websocket_test/websocket_client_test.cpp b/tests/networking_tests/websocket_test/websocket_client_test.cpp index 4ebade5665..401713d4bb 100644 --- a/tests/networking_tests/websocket_test/websocket_client_test.cpp +++ b/tests/networking_tests/websocket_test/websocket_client_test.cpp @@ -223,9 +223,13 @@ void run_counting_server(std::atomic& server_ready, std::atomic& sho } } -// Helper to run a server that returns a load-balancer-style connection header -void run_lb_header_handshake_server(std::atomic& server_ready, std::atomic& should_stop, - std::atomic& selected_port, std::string* host_header = nullptr) +constexpr std::string_view lb_upgrade_fields = + "Upgrade: websocket\r\n" + "Connection: keep-alive, Upgrade\r\n"; + +void run_handshake_response_server(std::atomic& server_ready, std::atomic& should_stop, + std::atomic& selected_port, std::string_view upgrade_fields, + std::string* host_header = nullptr) { try { asio::io_context io_ctx; @@ -269,12 +273,8 @@ void run_lb_header_handshake_server(std::atomic& server_ready, std::atomic if (websocket_key.empty()) return; const std::string accept_key = ws_util::generate_accept_key(websocket_key); - const std::string response = - "HTTP/1.1 101 Switching Protocols\r\n" - "Upgrade: websocket\r\n" - "Connection: keep-alive, Upgrade\r\n" - "Sec-WebSocket-Accept: " + - accept_key + "\r\n\r\n"; + const std::string response = "HTTP/1.1 101 Switching Protocols\r\n" + std::string(upgrade_fields) + + "Sec-WebSocket-Accept: " + accept_key + "\r\n\r\n"; asio::write(socket, asio::buffer(response), ec); if (ec) return; @@ -285,11 +285,32 @@ void run_lb_header_handshake_server(std::atomic& server_ready, std::atomic } } catch (const std::exception& e) { - std::cerr << "[lb_header_server] Exception: " << e.what() << "\n"; + std::cerr << "[handshake_response_server] Exception: " << e.what() << "\n"; server_ready = true; } } +std::string send_raw_upgrade_request(uint16_t port, const std::string& request) +{ + try { + asio::io_context io_ctx; + asio::ip::tcp::socket socket(io_ctx); + asio::error_code ec; + socket.connect(asio::ip::tcp::endpoint(asio::ip::make_address("127.0.0.1"), port), ec); + if (ec) return ""; + + asio::write(socket, asio::buffer(request), ec); + if (ec) return ""; + + asio::streambuf response_buf; + asio::read_until(socket, response_buf, "\r\n\r\n", ec); + return {asio::buffers_begin(response_buf.data()), asio::buffers_end(response_buf.data())}; + } + catch (const std::exception&) { + return ""; + } +} + // Helper to run a server that sends WebSocket frames in the same TCP segment // as the handshake response void run_initial_data_handshake_server(std::atomic& server_ready, std::atomic& should_stop, @@ -495,10 +516,8 @@ void run_auth_websocket_server(std::atomic& server_ready, std::atomic(); - ws_server->on_validate([expected_auth](const request& req) { - auto it = req.headers.find("authorization"); - return it != req.headers.end() && it->second == expected_auth; - }); + ws_server->on_validate( + [expected_auth](const request& req) { return req.headers.first_value("authorization") == expected_auth; }); ws_server->on_open([](auto /*conn*/, const request&) {}); ws_server->on_message([](auto conn, std::string_view message, ws_opcode opcode) { @@ -1341,8 +1360,8 @@ suite websocket_client_tests = [] { std::atomic stop_server{false}; std::atomic port{0}; - std::thread server_thread(run_lb_header_handshake_server, std::ref(server_ready), std::ref(stop_server), - std::ref(port), nullptr); + std::thread server_thread(run_handshake_response_server, std::ref(server_ready), std::ref(stop_server), + std::ref(port), lb_upgrade_fields, nullptr); expect(wait_for_condition([&] { return server_ready.load() && port.load() != 0; })) << "Server failed to start"; @@ -1387,14 +1406,99 @@ suite websocket_client_tests = [] { server_thread.join(); }; + "handshake_upgrade_field_name_must_match_in_full_test"_test = [] { + std::atomic server_ready{false}; + std::atomic stop_server{false}; + std::atomic port{0}; + + std::thread server_thread(run_handshake_response_server, std::ref(server_ready), std::ref(stop_server), + std::ref(port), + std::string_view{"Upgrade-Extra: websocket\r\n" + "Connection: Upgrade\r\n"}, + nullptr); + + expect(wait_for_condition([&] { return server_ready.load() && port.load() != 0; })) << "Server failed to start"; + + websocket_client client; + std::atomic open_called{false}; + std::atomic protocol_error{false}; + + client.on_open([&]() { + open_called = true; + client.context()->stop(); + }); + + client.on_message([](std::string_view, ws_opcode) {}); + + client.on_close([](ws_close_code, std::string_view) {}); + + client.on_error([&](std::error_code ec) { + if (ec == std::make_error_code(std::errc::protocol_error)) { + protocol_error = true; + } + client.context()->stop(); + }); + + client.connect("ws://127.0.0.1:" + std::to_string(port.load()) + "/ws"); + + std::thread client_thread([&client]() { client.context()->run(); }); + + expect(wait_for_condition([&] { return open_called.load() || protocol_error.load(); })) + << "Handshake neither completed nor failed"; + expect(!open_called.load()) << "Upgrade-Extra must not satisfy the Upgrade check\n"; + expect(protocol_error.load()) << "A response without a real Upgrade field must fail the handshake\n"; + + if (!client.context()->stopped()) { + client.context()->stop(); + } + if (client_thread.joinable()) { + client_thread.join(); + } + + stop_server = true; + server_thread.join(); + }; + + "upgrade_accepted_with_connection_split_over_two_fields_test"_test = [] { + std::atomic server_ready{false}; + std::atomic stop_server{false}; + std::atomic port{0}; + + std::thread server_thread(run_echo_server, std::ref(server_ready), std::ref(stop_server), std::ref(port)); + + expect(wait_for_condition([&] { return server_ready.load() && port.load() != 0; })) << "Server failed to start"; + + const std::string request = + "GET /ws HTTP/1.1\r\n" + "Host: 127.0.0.1:" + + std::to_string(port.load()) + + "\r\n" + "Upgrade: websocket\r\n" + "Connection: keep-alive\r\n" + "Connection: Upgrade\r\n" + "Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n" + "Sec-WebSocket-Version: 13\r\n" + "\r\n"; + + const std::string response = send_raw_upgrade_request(port.load(), request); + + expect(response.find("101") != std::string::npos) + << "Split Connection options must still be recognized as an upgrade, got: " << response << "\n"; + expect(response.find("s3pPLMBiTxaQ9kYGzzhZRbK+xOo=") != std::string::npos) + << "The RFC 6455 accept key for the sample nonce should come back\n"; + + stop_server = true; + server_thread.join(); + }; + "websocket_handshake_host_header_includes_port"_test = [] { std::atomic server_ready{false}; std::atomic stop_server{false}; std::atomic port{0}; std::string host_header; - std::thread server_thread(run_lb_header_handshake_server, std::ref(server_ready), std::ref(stop_server), - std::ref(port), &host_header); + std::thread server_thread(run_handshake_response_server, std::ref(server_ready), std::ref(stop_server), + std::ref(port), lb_upgrade_fields, &host_header); expect(wait_for_condition([&] { return server_ready.load() && port.load() != 0; })) << "Server failed to start";