diff --git a/src/CMakeLists.txt b/src/CMakeLists.txt index 8e864dc..e0b0b87 100644 --- a/src/CMakeLists.txt +++ b/src/CMakeLists.txt @@ -2,6 +2,7 @@ add_library(internal ${CMAKE_CURRENT_SOURCE_DIR}/dummy.cpp ${CMAKE_CURRENT_SOURCE_DIR}/url.cpp ${CMAKE_CURRENT_SOURCE_DIR}/headers.cpp + ${CMAKE_CURRENT_SOURCE_DIR}/http/request_serializer.cpp ${CMAKE_CURRENT_SOURCE_DIR}/platform/native_socket.cpp ${CMAKE_CURRENT_SOURCE_DIR}/platform/network_runtime.cpp ${CMAKE_CURRENT_SOURCE_DIR}/platform/socket_mode.cpp diff --git a/src/http/request_serializer.cpp b/src/http/request_serializer.cpp new file mode 100644 index 0000000..954390b --- /dev/null +++ b/src/http/request_serializer.cpp @@ -0,0 +1,240 @@ +#include "http/request_serializer.hpp" + +#include +#include +#include +#include +#include + +namespace cpp_request::detail::http { +namespace { + +[[nodiscard]] constexpr std::string_view method_token(Method method) noexcept { + switch (method) { + case Method::Get: return "GET"; + case Method::Head: return "HEAD"; + case Method::Post: return "POST"; + case Method::Put: return "PUT"; + case Method::Patch: return "PATCH"; + case Method::Delete: return "DELETE"; + } + return {}; +} + +[[nodiscard]] constexpr char ascii_lower(char ch) noexcept { + return ch >= 'A' && ch <= 'Z' + ? static_cast(ch + ('a' - 'A')) + : ch; +} + +[[nodiscard]] bool ascii_iequals( + std::string_view lhs, + std::string_view rhs) noexcept { + if (lhs.size() != rhs.size()) { + return false; + } + + for (std::size_t index = 0; index < lhs.size(); ++index) { + if (ascii_lower(lhs[index]) != ascii_lower(rhs[index])) { + return false; + } + } + return true; +} + +[[nodiscard]] constexpr bool is_tchar(unsigned char ch) noexcept { + if ((ch >= '0' && ch <= '9') + || (ch >= 'A' && ch <= 'Z') + || (ch >= 'a' && ch <= 'z')) { + return true; + } + + switch (ch) { + case '!': + case '#': + case '$': + case '%': + case '&': + case '\'': + case '*': + case '+': + case '-': + case '.': + case '^': + case '_': + case '`': + case '|': + case '~': + return true; + default: + return false; + } +} + +[[nodiscard]] bool valid_header_name(std::string_view name) noexcept { + if (name.empty()) { + return false; + } + + for (const unsigned char ch : name) { + if (!is_tchar(ch)) { + return false; + } + } + return true; +} + +[[nodiscard]] bool valid_header_value(std::string_view value) noexcept { + for (const unsigned char ch : value) { + if (ch == '\t') { + continue; + } + if (ch < 0x20 || ch == 0x7f) { + return false; + } + } + return true; +} + +[[nodiscard]] bool parse_content_length( + std::string_view text, + std::size_t& value) noexcept { + if (text.empty()) { + return false; + } + + const char* const begin = text.data(); + const char* const end = text.data() + text.size(); + const auto parsed = std::from_chars(begin, end, value, 10); + return parsed.ec == std::errc{} && parsed.ptr == end; +} + +void append_decimal(std::string& output, std::size_t value) { + char buffer[32]{}; + const auto converted = std::to_chars( + buffer, + buffer + sizeof(buffer), + value, + 10); + output.append(buffer, converted.ptr); +} + +void append_port(std::string& output, std::uint16_t port) { + char buffer[8]{}; + const auto converted = std::to_chars( + buffer, + buffer + sizeof(buffer), + port, + 10); + output.append(buffer, converted.ptr); +} + +void append_generated_host(std::string& output, const Url& url) { + output.append("Host: "); + if (url.host_is_ipv6_literal()) { + output.push_back('['); + output.append(url.host().data(), url.host().size()); + output.push_back(']'); + } else { + output.append(url.host().data(), url.host().size()); + } + + if (url.has_explicit_port()) { + output.push_back(':'); + append_port(output, url.port()); + } + output.append("\r\n"); +} + +} // namespace + +Result serialize_request(const Request& request) { + auto parsed_url = Url::parse(request.url()); + if (!parsed_url) { + return parsed_url.error(); + } + + const std::string_view method = method_token(request.method()); + if (method.empty()) { + return Error{ErrorCode::Unknown}; + } + + std::size_t host_count = 0; + std::size_t content_length_count = 0; + std::size_t declared_content_length = 0; + + for (const auto& field : request.headers()) { + if (!valid_header_name(field.name) || !valid_header_value(field.value)) { + return Error{ErrorCode::InvalidHeader}; + } + + if (ascii_iequals(field.name, "Host")) { + ++host_count; + if (field.value.empty()) { + return Error{ErrorCode::InvalidHeader}; + } + } else if (ascii_iequals(field.name, "Content-Length")) { + ++content_length_count; + if (!parse_content_length(field.value, declared_content_length)) { + return Error{ErrorCode::InvalidContentLength}; + } + } else if (ascii_iequals(field.name, "Transfer-Encoding")) { + return Error{ErrorCode::ConflictingMessageFraming}; + } + } + + if (host_count > 1) { + return Error{ErrorCode::InvalidHeader}; + } + if (content_length_count > 1) { + return Error{ErrorCode::InvalidContentLength}; + } + if (content_length_count == 1 + && declared_content_length != request.body().size()) { + return Error{ErrorCode::InvalidContentLength}; + } + + SerializedRequest serialized; + serialized.url = std::move(parsed_url).value(); + serialized.body = request.body(); + + std::size_t reserve_size = method.size() + + 1 + + serialized.url.target().size() + + sizeof(" HTTP/1.1\r\n") - 1 + + request.headers().size() * 16 + + 64; + for (const auto& field : request.headers()) { + reserve_size += field.name.size() + field.value.size(); + } + serialized.head.reserve(reserve_size); + + serialized.head.append(method.data(), method.size()); + serialized.head.push_back(' '); + serialized.head.append( + serialized.url.target().data(), + serialized.url.target().size()); + serialized.head.append(" HTTP/1.1\r\n"); + + if (host_count == 0) { + append_generated_host(serialized.head, serialized.url); + } + + for (const auto& field : request.headers()) { + serialized.head.append(field.name); + serialized.head.append(": "); + serialized.head.append(field.value); + serialized.head.append("\r\n"); + } + + if (!request.body().empty() && content_length_count == 0) { + serialized.head.append("Content-Length: "); + append_decimal(serialized.head, request.body().size()); + serialized.head.append("\r\n"); + } + + serialized.head.append("\r\n"); + return serialized; +} + +} // namespace cpp_request::detail::http diff --git a/src/http/request_serializer.hpp b/src/http/request_serializer.hpp new file mode 100644 index 0000000..24eee57 --- /dev/null +++ b/src/http/request_serializer.hpp @@ -0,0 +1,20 @@ +#pragma once + +#include +#include + +#include +#include +#include + +namespace cpp_request::detail::http { + +struct SerializedRequest final { + Url url; + std::string head; + std::string_view body; +}; + +[[nodiscard]] Result serialize_request(const Request& request); + +} // namespace cpp_request::detail::http diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index e07818a..e73b3e4 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -5,6 +5,7 @@ add_executable(cpp_request_tests ${CMAKE_CURRENT_SOURCE_DIR}/url_test.cpp ${CMAKE_CURRENT_SOURCE_DIR}/headers_test.cpp ${CMAKE_CURRENT_SOURCE_DIR}/request_test.cpp + ${CMAKE_CURRENT_SOURCE_DIR}/request_serializer_test.cpp ${CMAKE_CURRENT_SOURCE_DIR}/native_socket_test.cpp ${CMAKE_CURRENT_SOURCE_DIR}/resolver_test.cpp ${CMAKE_CURRENT_SOURCE_DIR}/tcp_connection_test.cpp diff --git a/tests/request_serializer_test.cpp b/tests/request_serializer_test.cpp new file mode 100644 index 0000000..1a77658 --- /dev/null +++ b/tests/request_serializer_test.cpp @@ -0,0 +1,194 @@ +#include + +#include "http/request_serializer.hpp" + +#include +#include + +#include + +namespace { + +using cpp_request::ErrorCode; +using cpp_request::Method; +using cpp_request::Request; +using cpp_request::detail::http::serialize_request; + +TEST(RequestSerializerTest, SerializesBasicGetRequest) { + Request request{Method::Get, "http://example.com/items?limit=10"}; + + auto result = serialize_request(request); + + ASSERT_TRUE(result); + EXPECT_EQ( + result.value().head, + "GET /items?limit=10 HTTP/1.1\r\n" + "Host: example.com\r\n" + "\r\n"); + EXPECT_TRUE(result.value().body.empty()); +} + +TEST(RequestSerializerTest, SerializesAllSupportedMethods) { + struct Case { + Method method; + const char* token; + }; + const Case cases[] = { + {Method::Get, "GET"}, + {Method::Head, "HEAD"}, + {Method::Post, "POST"}, + {Method::Put, "PUT"}, + {Method::Patch, "PATCH"}, + {Method::Delete, "DELETE"}, + }; + + for (const auto& test_case : cases) { + Request request{test_case.method, "http://example.com/"}; + auto result = serialize_request(request); + ASSERT_TRUE(result); + EXPECT_EQ( + result.value().head.substr(0, std::string{test_case.token}.size()), + test_case.token); + } +} + +TEST(RequestSerializerTest, PreservesCustomHeadersAndDuplicates) { + Request request{Method::Get, "http://example.com/"}; + request.headers().add("X-Test", "one"); + request.headers().add("X-Test", "two"); + + auto result = serialize_request(request); + + ASSERT_TRUE(result); + EXPECT_NE(result.value().head.find("X-Test: one\r\n"), std::string::npos); + EXPECT_NE(result.value().head.find("X-Test: two\r\n"), std::string::npos); +} + +TEST(RequestSerializerTest, UsesCallerProvidedHostWithoutGeneratingAnother) { + Request request{Method::Get, "http://example.com/"}; + request.headers().add("Host", "virtual.example"); + + auto result = serialize_request(request); + + ASSERT_TRUE(result); + EXPECT_NE(result.value().head.find("Host: virtual.example\r\n"), std::string::npos); + EXPECT_EQ(result.value().head.find("Host: example.com\r\n"), std::string::npos); +} + +TEST(RequestSerializerTest, GeneratedHostIncludesExplicitPort) { + Request request{Method::Get, "http://example.com:8080/"}; + + auto result = serialize_request(request); + + ASSERT_TRUE(result); + EXPECT_NE(result.value().head.find("Host: example.com:8080\r\n"), std::string::npos); +} + +TEST(RequestSerializerTest, GeneratedHostFormatsIpv6Literal) { + Request request{Method::Get, "http://[::1]:8080/health"}; + + auto result = serialize_request(request); + + ASSERT_TRUE(result); + EXPECT_NE(result.value().head.find("Host: [::1]:8080\r\n"), std::string::npos); +} + +TEST(RequestSerializerTest, AddsContentLengthForNonEmptyBodyWithoutCopyingBody) { + std::string body = "hello"; + Request request{Method::Post, "http://example.com/upload"}; + request.set_body(body); + + auto result = serialize_request(request); + + ASSERT_TRUE(result); + EXPECT_NE(result.value().head.find("Content-Length: 5\r\n"), std::string::npos); + EXPECT_EQ(result.value().body, body); + EXPECT_EQ(result.value().body.data(), body.data()); +} + +TEST(RequestSerializerTest, AcceptsMatchingCallerContentLength) { + Request request{Method::Post, "http://example.com/upload"}; + request.set_body("abc"); + request.headers().add("Content-Length", "3"); + + auto result = serialize_request(request); + + ASSERT_TRUE(result); + EXPECT_NE(result.value().head.find("Content-Length: 3\r\n"), std::string::npos); +} + +TEST(RequestSerializerTest, RejectsMismatchedContentLength) { + Request request{Method::Post, "http://example.com/upload"}; + request.set_body("abc"); + request.headers().add("Content-Length", "2"); + + auto result = serialize_request(request); + + ASSERT_FALSE(result); + EXPECT_EQ(result.error().code, ErrorCode::InvalidContentLength); +} + +TEST(RequestSerializerTest, RejectsDuplicateContentLength) { + Request request{Method::Post, "http://example.com/upload"}; + request.set_body("abc"); + request.headers().add("Content-Length", "3"); + request.headers().add("content-length", "3"); + + auto result = serialize_request(request); + + ASSERT_FALSE(result); + EXPECT_EQ(result.error().code, ErrorCode::InvalidContentLength); +} + +TEST(RequestSerializerTest, RejectsDuplicateHost) { + Request request{Method::Get, "http://example.com/"}; + request.headers().add("Host", "example.com"); + request.headers().add("host", "example.com"); + + auto result = serialize_request(request); + + ASSERT_FALSE(result); + EXPECT_EQ(result.error().code, ErrorCode::InvalidHeader); +} + +TEST(RequestSerializerTest, RejectsHeaderNameWithInvalidTokenByte) { + Request request{Method::Get, "http://example.com/"}; + request.headers().add("Bad Header", "value"); + + auto result = serialize_request(request); + + ASSERT_FALSE(result); + EXPECT_EQ(result.error().code, ErrorCode::InvalidHeader); +} + +TEST(RequestSerializerTest, RejectsHeaderValueWithCrLfInjection) { + Request request{Method::Get, "http://example.com/"}; + request.headers().add("X-Test", "ok\r\nInjected: yes"); + + auto result = serialize_request(request); + + ASSERT_FALSE(result); + EXPECT_EQ(result.error().code, ErrorCode::InvalidHeader); +} + +TEST(RequestSerializerTest, RejectsTransferEncodingInV1Requests) { + Request request{Method::Post, "http://example.com/"}; + request.set_body("abc"); + request.headers().add("Transfer-Encoding", "chunked"); + + auto result = serialize_request(request); + + ASSERT_FALSE(result); + EXPECT_EQ(result.error().code, ErrorCode::ConflictingMessageFraming); +} + +TEST(RequestSerializerTest, PropagatesUrlParseFailure) { + Request request{Method::Get, "https://example.com/"}; + + auto result = serialize_request(request); + + ASSERT_FALSE(result); + EXPECT_EQ(result.error().code, ErrorCode::UnsupportedScheme); +} + +} // namespace