From 906500cb36619b5cbfa5bbbf9056edf6c2baef5a Mon Sep 17 00:00:00 2001 From: Benjamin Oldenburg Date: Mon, 20 Jul 2026 01:19:56 +0700 Subject: [PATCH] fix(cpp-boost-beast): validate request components --- .../http-client-impl-source.mustache | 212 +++++++++++++++++- .../CppBoostBeastClientCodegenTest.java | 5 +- .../generated/api/HttpClientImpl.cpp | 212 +++++++++++++++++- .../tests/api/http_client_test.cpp | 68 +++++- 4 files changed, 475 insertions(+), 22 deletions(-) diff --git a/modules/openapi-generator/src/main/resources/cpp-boost-beast-client/http-client-impl-source.mustache b/modules/openapi-generator/src/main/resources/cpp-boost-beast-client/http-client-impl-source.mustache index 8b08d4db7b54..770c99f0820d 100644 --- a/modules/openapi-generator/src/main/resources/cpp-boost-beast-client/http-client-impl-source.mustache +++ b/modules/openapi-generator/src/main/resources/cpp-boost-beast-client/http-client-impl-source.mustache @@ -26,6 +26,196 @@ namespace { using OperationCompletion = std::function; +bool isAsciiAlphaNumeric(const unsigned char character) { + return (character >= '0' && character <= '9') || + (character >= 'A' && character <= 'Z') || + (character >= 'a' && character <= 'z'); +} + +bool isRfcTokenCharacter(const unsigned char character) { + if (isAsciiAlphaNumeric(character)) { + return true; + } + + switch (character) { + case '!': + case '#': + case '$': + case '%': + case '&': + case '\'': + case '*': + case '+': + case '-': + case '.': + case '^': + case '_': + case '`': + case '|': + case '~': + return true; + default: + return false; + } +} + +bool isRfcToken(const std::string &token) { + if (token.empty()) { + return false; + } + + for (const unsigned char character : token) { + if (!isRfcTokenCharacter(character)) { + return false; + } + } + return true; +} + +bool containsControlCharacter(const std::string &value) { + for (const unsigned char character : value) { + if (character < 0x20 || character == 0x7f) { + return true; + } + } + return false; +} + +bool isHexDigit(const unsigned char character) { + return (character >= '0' && character <= '9') || + (character >= 'A' && character <= 'F') || + (character >= 'a' && character <= 'f'); +} + +bool isOriginFormTargetCharacter(const unsigned char character) { + if (isAsciiAlphaNumeric(character)) { + return true; + } + + switch (character) { + case '!': + case '$': + case '&': + case '\'': + case '(': + case ')': + case '*': + case '+': + case ',': + case '-': + case '.': + case '/': + case ':': + case ';': + case '=': + case '?': + case '@': + case '_': + case '~': + return true; + default: + return false; + } +} + +void validateHost(const std::string &host) { + if (host.empty()) { + throw std::invalid_argument("host must not be empty"); + } + if (containsControlCharacter(host) || + host.find_first_of(" /\\?#") != std::string::npos) { + throw std::invalid_argument("host contains an invalid character"); + } +} + +void validatePort(const std::string &port) { + if (port.empty()) { + throw std::invalid_argument("port must not be empty"); + } + if (containsControlCharacter(port) || + port.find_first_of(" /\\:?#") != std::string::npos) { + throw std::invalid_argument("port contains an invalid character"); + } +} + +void validateRequestTarget(const std::string &target) { + if (target.empty() || target.front() != '/') { + throw std::invalid_argument("target must use HTTP origin-form"); + } + + for (std::size_t index = 0; index < target.size(); ++index) { + const unsigned char character = + static_cast(target[index]); + if (character == '%') { + if (index + 2 >= target.size() || + !isHexDigit(static_cast(target[index + 1])) || + !isHexDigit(static_cast(target[index + 2]))) { + throw std::invalid_argument("target contains an invalid percent escape"); + } + index += 2; + } else if (!isOriginFormTargetCharacter(character)) { + throw std::invalid_argument("target contains an invalid character"); + } + } +} + +bool asciiCaseInsensitiveEqual(const std::string &headerName, + const char *reservedHeaderName) { + const std::size_t reservedHeaderNameLength = + std::char_traits::length(reservedHeaderName); + if (headerName.size() != reservedHeaderNameLength) { + return false; + } + + for (std::size_t index = 0; index < headerName.size(); ++index) { + const unsigned char headerCharacter = + static_cast(headerName[index]); + const unsigned char reservedCharacter = + static_cast(reservedHeaderName[index]); + const unsigned char lowerHeaderCharacter = + headerCharacter >= 'A' && headerCharacter <= 'Z' + ? static_cast(headerCharacter + ('a' - 'A')) + : headerCharacter; + const unsigned char lowerReservedCharacter = + reservedCharacter >= 'A' && reservedCharacter <= 'Z' + ? static_cast(reservedCharacter + ('a' - 'A')) + : reservedCharacter; + if (lowerHeaderCharacter != lowerReservedCharacter) { + return false; + } + } + return true; +} + +void validateCallerHeader(const std::string &headerName, + const std::string &headerValue) { + if (!isRfcToken(headerName)) { + throw std::invalid_argument("header name must be a non-empty RFC token"); + } + + const char *const reservedHeaderNames[] = { + "Host", + "Content-Length", + "Transfer-Encoding", + "Connection", + "Proxy-Connection", + "Keep-Alive", + "Upgrade", + "TE", + "Trailer", + "Expect", + "Proxy-Authorization"}; + for (const char *reservedHeaderName : reservedHeaderNames) { + if (asciiCaseInsensitiveEqual(headerName, reservedHeaderName)) { + throw std::invalid_argument("header name is reserved by the transport"); + } + } + + if (containsControlCharacter(headerValue)) { + throw std::invalid_argument("header value must not contain control characters"); + } +} + void throwOperationError(const boost::beast::error_code &operationError, const char *operationName) { if (operationError) { @@ -161,12 +351,8 @@ HttpClientImpl::HttpClientImpl( m_operationTimeout(operationTimeout), m_responseBodyLimit(responseBodyLimit) { - if (m_host.empty()) { - throw std::invalid_argument("host must not be empty"); - } - if (m_port.empty()) { - throw std::invalid_argument("port must not be empty"); - } + validateHost(m_host); + validatePort(m_port); if (m_httpVersion != 10 && m_httpVersion != 11) { throw std::invalid_argument("httpVersion must be 10 or 11"); } @@ -388,8 +574,16 @@ HttpClientImpl::prepareRequest(const std::string &verb, const std::string &target, const std::string &body, const std::map &headers) { - HttpRequest request{ - boost::beast::http::string_to_verb(verb), target, m_httpVersion}; + if (!isRfcToken(verb)) { + throw std::invalid_argument("verb must be a non-empty RFC token"); + } + validateRequestTarget(target); + + HttpRequest request; + request.version(m_httpVersion); + request.method_string(verb); + request.target(target); + request.body() = body; boost::beast::error_code hostAddressError; const boost::asio::ip::address hostAddress = @@ -398,11 +592,11 @@ HttpClientImpl::prepareRequest(const std::string &verb, ? "[" + m_host + "]:" + m_port : m_host + ":" + m_port; request.set(boost::beast::http::field::host, hostHeader); - request.body() = body; request.set( boost::beast::http::field::user_agent, BOOST_BEAST_VERSION_STRING); for (const auto &header : headers) { + validateCallerHeader(header.first, header.second); request.set(header.first, header.second); } diff --git a/modules/openapi-generator/src/test/java/org/openapitools/codegen/cppboostbeast/CppBoostBeastClientCodegenTest.java b/modules/openapi-generator/src/test/java/org/openapitools/codegen/cppboostbeast/CppBoostBeastClientCodegenTest.java index ba862be7b4f7..02cce2b34df0 100644 --- a/modules/openapi-generator/src/test/java/org/openapitools/codegen/cppboostbeast/CppBoostBeastClientCodegenTest.java +++ b/modules/openapi-generator/src/test/java/org/openapitools/codegen/cppboostbeast/CppBoostBeastClientCodegenTest.java @@ -67,7 +67,10 @@ public void generatesTypedJsonValuesForOpenApi31Schemas() throws IOException { "SSL_CTX_set_min_proto_version(", "TLS1_2_VERSION", "boost::asio::ssl::verify_peer", - "boost::asio::ssl::host_name_verification(m_host)"); + "boost::asio::ssl::host_name_verification(m_host)", + "request.method_string(verb)", + "target must use HTTP origin-form", + "header name is reserved by the transport"); } @Test diff --git a/samples/client/petstore/cpp-boost-beast/generated/api/HttpClientImpl.cpp b/samples/client/petstore/cpp-boost-beast/generated/api/HttpClientImpl.cpp index 6a3c8132b5ff..c63560027a47 100644 --- a/samples/client/petstore/cpp-boost-beast/generated/api/HttpClientImpl.cpp +++ b/samples/client/petstore/cpp-boost-beast/generated/api/HttpClientImpl.cpp @@ -26,6 +26,196 @@ namespace { using OperationCompletion = std::function; +bool isAsciiAlphaNumeric(const unsigned char character) { + return (character >= '0' && character <= '9') || + (character >= 'A' && character <= 'Z') || + (character >= 'a' && character <= 'z'); +} + +bool isRfcTokenCharacter(const unsigned char character) { + if (isAsciiAlphaNumeric(character)) { + return true; + } + + switch (character) { + case '!': + case '#': + case '$': + case '%': + case '&': + case '\'': + case '*': + case '+': + case '-': + case '.': + case '^': + case '_': + case '`': + case '|': + case '~': + return true; + default: + return false; + } +} + +bool isRfcToken(const std::string &token) { + if (token.empty()) { + return false; + } + + for (const unsigned char character : token) { + if (!isRfcTokenCharacter(character)) { + return false; + } + } + return true; +} + +bool containsControlCharacter(const std::string &value) { + for (const unsigned char character : value) { + if (character < 0x20 || character == 0x7f) { + return true; + } + } + return false; +} + +bool isHexDigit(const unsigned char character) { + return (character >= '0' && character <= '9') || + (character >= 'A' && character <= 'F') || + (character >= 'a' && character <= 'f'); +} + +bool isOriginFormTargetCharacter(const unsigned char character) { + if (isAsciiAlphaNumeric(character)) { + return true; + } + + switch (character) { + case '!': + case '$': + case '&': + case '\'': + case '(': + case ')': + case '*': + case '+': + case ',': + case '-': + case '.': + case '/': + case ':': + case ';': + case '=': + case '?': + case '@': + case '_': + case '~': + return true; + default: + return false; + } +} + +void validateHost(const std::string &host) { + if (host.empty()) { + throw std::invalid_argument("host must not be empty"); + } + if (containsControlCharacter(host) || + host.find_first_of(" /\\?#") != std::string::npos) { + throw std::invalid_argument("host contains an invalid character"); + } +} + +void validatePort(const std::string &port) { + if (port.empty()) { + throw std::invalid_argument("port must not be empty"); + } + if (containsControlCharacter(port) || + port.find_first_of(" /\\:?#") != std::string::npos) { + throw std::invalid_argument("port contains an invalid character"); + } +} + +void validateRequestTarget(const std::string &target) { + if (target.empty() || target.front() != '/') { + throw std::invalid_argument("target must use HTTP origin-form"); + } + + for (std::size_t index = 0; index < target.size(); ++index) { + const unsigned char character = + static_cast(target[index]); + if (character == '%') { + if (index + 2 >= target.size() || + !isHexDigit(static_cast(target[index + 1])) || + !isHexDigit(static_cast(target[index + 2]))) { + throw std::invalid_argument("target contains an invalid percent escape"); + } + index += 2; + } else if (!isOriginFormTargetCharacter(character)) { + throw std::invalid_argument("target contains an invalid character"); + } + } +} + +bool asciiCaseInsensitiveEqual(const std::string &headerName, + const char *reservedHeaderName) { + const std::size_t reservedHeaderNameLength = + std::char_traits::length(reservedHeaderName); + if (headerName.size() != reservedHeaderNameLength) { + return false; + } + + for (std::size_t index = 0; index < headerName.size(); ++index) { + const unsigned char headerCharacter = + static_cast(headerName[index]); + const unsigned char reservedCharacter = + static_cast(reservedHeaderName[index]); + const unsigned char lowerHeaderCharacter = + headerCharacter >= 'A' && headerCharacter <= 'Z' + ? static_cast(headerCharacter + ('a' - 'A')) + : headerCharacter; + const unsigned char lowerReservedCharacter = + reservedCharacter >= 'A' && reservedCharacter <= 'Z' + ? static_cast(reservedCharacter + ('a' - 'A')) + : reservedCharacter; + if (lowerHeaderCharacter != lowerReservedCharacter) { + return false; + } + } + return true; +} + +void validateCallerHeader(const std::string &headerName, + const std::string &headerValue) { + if (!isRfcToken(headerName)) { + throw std::invalid_argument("header name must be a non-empty RFC token"); + } + + const char *const reservedHeaderNames[] = { + "Host", + "Content-Length", + "Transfer-Encoding", + "Connection", + "Proxy-Connection", + "Keep-Alive", + "Upgrade", + "TE", + "Trailer", + "Expect", + "Proxy-Authorization"}; + for (const char *reservedHeaderName : reservedHeaderNames) { + if (asciiCaseInsensitiveEqual(headerName, reservedHeaderName)) { + throw std::invalid_argument("header name is reserved by the transport"); + } + } + + if (containsControlCharacter(headerValue)) { + throw std::invalid_argument("header value must not contain control characters"); + } +} + void throwOperationError(const boost::beast::error_code &operationError, const char *operationName) { if (operationError) { @@ -162,12 +352,8 @@ HttpClientImpl::HttpClientImpl( m_operationTimeout(operationTimeout), m_responseBodyLimit(responseBodyLimit) { - if (m_host.empty()) { - throw std::invalid_argument("host must not be empty"); - } - if (m_port.empty()) { - throw std::invalid_argument("port must not be empty"); - } + validateHost(m_host); + validatePort(m_port); if (m_httpVersion != 10 && m_httpVersion != 11) { throw std::invalid_argument("httpVersion must be 10 or 11"); } @@ -389,8 +575,16 @@ HttpClientImpl::prepareRequest(const std::string &verb, const std::string &target, const std::string &body, const std::map &headers) { - HttpRequest request{ - boost::beast::http::string_to_verb(verb), target, m_httpVersion}; + if (!isRfcToken(verb)) { + throw std::invalid_argument("verb must be a non-empty RFC token"); + } + validateRequestTarget(target); + + HttpRequest request; + request.version(m_httpVersion); + request.method_string(verb); + request.target(target); + request.body() = body; boost::beast::error_code hostAddressError; const boost::asio::ip::address hostAddress = @@ -399,11 +593,11 @@ HttpClientImpl::prepareRequest(const std::string &verb, ? "[" + m_host + "]:" + m_port : m_host + ":" + m_port; request.set(boost::beast::http::field::host, hostHeader); - request.body() = body; request.set( boost::beast::http::field::user_agent, BOOST_BEAST_VERSION_STRING); for (const auto &header : headers) { + validateCallerHeader(header.first, header.second); request.set(header.first, header.second); } diff --git a/samples/client/petstore/cpp-boost-beast/tests/api/http_client_test.cpp b/samples/client/petstore/cpp-boost-beast/tests/api/http_client_test.cpp index 4d6089f0ae44..04a7f04b1eb7 100644 --- a/samples/client/petstore/cpp-boost-beast/tests/api/http_client_test.cpp +++ b/samples/client/petstore/cpp-boost-beast/tests/api/http_client_test.cpp @@ -181,6 +181,14 @@ class RequestInspectingHttpClient final : public HttpClientImpl { const beast::string_view hostHeader = request[http::field::host]; return std::string(hostHeader.data(), hostHeader.size()); } + + HttpRequest prepareRequestForInspection( + const std::string &verb, + const std::string &target, + const std::string &body, + const std::map &headers) { + return prepareRequest(verb, target, body, headers); + } }; class SerializedLifecycleHttpClient final : public HttpClientImpl { @@ -609,10 +617,64 @@ BOOST_AUTO_TEST_CASE(prepare_request_formats_host_for_address_type) { BOOST_REQUIRE_EQUAL(ipv6Client.prepareHostHeader({}), "[2001:db8::10]:8082"); - const std::map overridingHeaders{ + const std::map reservedHeaders{ {"Host", "caller.example:9090"}}; - BOOST_REQUIRE_EQUAL(ipv6Client.prepareHostHeader(overridingHeaders), - "caller.example:9090"); + BOOST_CHECK_THROW(ipv6Client.prepareHostHeader(reservedHeaders), + std::invalid_argument); +} + +BOOST_AUTO_TEST_CASE(prepare_request_rejects_unsafe_wire_components) { + RequestInspectingHttpClient client("example.test", "https"); + + const auto customRequest = + client.prepareRequestForInspection( + "PURGE", + "/items%20one?tag=alpha&tag=beta", + "body", + {{"User-Agent", "test-client"}, {"X-Test", "value"}}); + const beast::string_view customMethod = customRequest.method_string(); + const beast::string_view customTarget = customRequest.target(); + const beast::string_view customUserAgent = + customRequest[http::field::user_agent]; + BOOST_REQUIRE_EQUAL( + std::string(customMethod.data(), customMethod.size()), "PURGE"); + BOOST_REQUIRE_EQUAL( + std::string(customTarget.data(), customTarget.size()), + "/items%20one?tag=alpha&tag=beta"); + BOOST_REQUIRE_EQUAL( + std::string(customUserAgent.data(), customUserAgent.size()), + "test-client"); + + BOOST_CHECK_THROW( + RequestInspectingHttpClient("bad\r\nhost", "443"), + std::invalid_argument); + BOOST_CHECK_THROW( + RequestInspectingHttpClient("example.test", "443\r\nX-Test: injected"), + std::invalid_argument); + BOOST_CHECK_THROW( + client.prepareRequestForInspection("GET\r\nX-Test: injected", "/", "", {}), + std::invalid_argument); + BOOST_CHECK_THROW( + client.prepareRequestForInspection("GET", "relative", "", {}), + std::invalid_argument); + BOOST_CHECK_THROW( + client.prepareRequestForInspection("GET", "/fragment#value", "", {}), + std::invalid_argument); + BOOST_CHECK_THROW( + client.prepareRequestForInspection("GET", "/invalid%2", "", {}), + std::invalid_argument); + BOOST_CHECK_THROW( + client.prepareRequestForInspection("GET", "/", "", {{"Bad Header", "value"}}), + std::invalid_argument); + BOOST_CHECK_THROW( + client.prepareRequestForInspection("GET", "/", "", {{"X-Test", "value\r\nInjected: true"}}), + std::invalid_argument); + BOOST_CHECK_THROW( + client.prepareRequestForInspection("GET", "/", "", {{"content-length", "0"}}), + std::invalid_argument); + BOOST_CHECK_THROW( + client.prepareRequestForInspection("GET", "/", "", {{"Host", "other.example"}}), + std::invalid_argument); } BOOST_AUTO_TEST_CASE(