Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 38 additions & 4 deletions wish/cpp/src/handshake.cc
Original file line number Diff line number Diff line change
Expand Up @@ -198,10 +198,14 @@ bool ValidateHeaders(const phr_header* headers, size_t num_headers) {

ClientHandshake::ClientHandshake(bufferevent* bev,
OnOpenCallback on_open,
OnErrorCallback on_error)
OnErrorCallback on_error,
size_t max_header_size,
int timeout_seconds)
: bev_(bev),
on_open_(std::move(on_open)),
on_error_(std::move(on_error)) {}
on_error_(std::move(on_error)),
max_header_size_(max_header_size),
timeout_seconds_(timeout_seconds) {}

ClientHandshake::~ClientHandshake() {
if (bev_) {
Expand All @@ -212,6 +216,11 @@ ClientHandshake::~ClientHandshake() {
void ClientHandshake::Start() {
bufferevent_setcb(bev_, ReadCb, nullptr, EventCb, this);

if (timeout_seconds_ > 0) {
struct timeval tv = {timeout_seconds_, 0};
bufferevent_set_timeouts(bev_, &tv, nullptr);
}

int enable_rv = bufferevent_enable(bev_, EV_READ | EV_WRITE);
if (enable_rv != 0) {
VLOG(1) << "bufferevent_enable() failed";
Expand Down Expand Up @@ -252,6 +261,14 @@ void ClientHandshake::HandleRead() {
return;
}

if (len > max_header_size_) {
VLOG(2) << "Client handshake header size exceeded limit: " << len;

InvokeError();

return;
}

const char* data = reinterpret_cast<const char*>(evbuffer_pullup(input, -1));
if (!data) {
ABSL_UNREACHABLE();
Expand Down Expand Up @@ -360,11 +377,15 @@ void ClientHandshake::InvokeError() {
ServerHandshake::ServerHandshake(bufferevent* bev,
OnOpenCallback on_open,
OnErrorCallback on_error,
CleanupCallback cleanup)
CleanupCallback cleanup,
size_t max_header_size,
int timeout_seconds)
: bev_(bev),
on_open_(std::move(on_open)),
on_error_(std::move(on_error)),
cleanup_(std::move(cleanup)) {}
cleanup_(std::move(cleanup)),
max_header_size_(max_header_size),
timeout_seconds_(timeout_seconds) {}

ServerHandshake::~ServerHandshake() {
if (bev_) {
Expand All @@ -375,6 +396,11 @@ ServerHandshake::~ServerHandshake() {
void ServerHandshake::Start() {
bufferevent_setcb(bev_, ReadCb, nullptr, EventCb, this);

if (timeout_seconds_ > 0) {
struct timeval tv = {timeout_seconds_, 0};
bufferevent_set_timeouts(bev_, &tv, nullptr);
}

int enable_rv = bufferevent_enable(bev_, EV_READ | EV_WRITE);
if (enable_rv != 0) {
VLOG(1) << "bufferevent_enable() failed";
Expand All @@ -401,6 +427,14 @@ void ServerHandshake::HandleRead() {
return;
}

if (len > max_header_size_) {
VLOG(2) << "Server handshake header size exceeded limit: " << len;

InvokeError();

return;
}

const char* data = reinterpret_cast<const char*>(evbuffer_pullup(input, -1));
if (!data) {
ABSL_UNREACHABLE();
Expand Down
16 changes: 14 additions & 2 deletions wish/cpp/src/handshake.h
Original file line number Diff line number Diff line change
Expand Up @@ -24,12 +24,17 @@
#include <memory>
#include <string>

constexpr size_t kDefaultMaxHeaderSize = 64 * 1024; // 64 KB
constexpr int kDefaultHandshakeTimeoutSeconds = 10; // 10 seconds

class ClientHandshake {
public:
using OnOpenCallback = std::function<void(bufferevent*)>;
using OnErrorCallback = std::function<void()>;

ClientHandshake(bufferevent* bev, OnOpenCallback on_open, OnErrorCallback on_error);
ClientHandshake(bufferevent* bev, OnOpenCallback on_open, OnErrorCallback on_error,
size_t max_header_size = kDefaultMaxHeaderSize,
int timeout_seconds = kDefaultHandshakeTimeoutSeconds);
~ClientHandshake();

void Start();
Expand All @@ -46,6 +51,8 @@ class ClientHandshake {

OnOpenCallback on_open_;
OnErrorCallback on_error_;
size_t max_header_size_;
int timeout_seconds_;
};

class ServerHandshake {
Expand All @@ -54,7 +61,10 @@ class ServerHandshake {
using OnErrorCallback = std::function<void()>;
using CleanupCallback = std::function<void(ServerHandshake*)>;

ServerHandshake(bufferevent* bev, OnOpenCallback on_open, OnErrorCallback on_error, CleanupCallback cleanup = nullptr);
ServerHandshake(bufferevent* bev, OnOpenCallback on_open, OnErrorCallback on_error,
CleanupCallback cleanup = nullptr,
size_t max_header_size = kDefaultMaxHeaderSize,
int timeout_seconds = kDefaultHandshakeTimeoutSeconds);
~ServerHandshake();

void Start();
Expand All @@ -72,6 +82,8 @@ class ServerHandshake {
OnOpenCallback on_open_;
OnErrorCallback on_error_;
CleanupCallback cleanup_;
size_t max_header_size_;
int timeout_seconds_;
};

#endif // WISH_CPP_SRC_HANDSHAKE_H_
63 changes: 63 additions & 0 deletions wish/cpp/src/handshake_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -651,3 +651,66 @@ TEST_F(HandshakeTest, ServerHandshakeRejectsHTTP10) {
EXPECT_TRUE(error_called);
bufferevent_free(pair[1]);
}

TEST_F(HandshakeTest, ServerHandshakeRejectsExceededHeaderSize) {
bufferevent* pair[2];
int rv = bufferevent_pair_new(base_, BEV_OPT_CLOSE_ON_FREE | BEV_OPT_DEFER_CALLBACKS, pair);
ASSERT_EQ(rv, 0);
bufferevent_enable(pair[1], EV_READ | EV_WRITE);

bool open_called = false;
bool error_called = false;
// Set tiny max header size limit of 100 bytes
auto server = std::make_unique<ServerHandshake>(
pair[0],
[&](bufferevent* bev) { open_called = true; bufferevent_free(bev); event_base_loopbreak(base_); },
[&]() { error_called = true; event_base_loopbreak(base_); },
nullptr,
100);
server->Start();

// Send incomplete header larger than 100 bytes
std::string huge_header = "POST / HTTP/1.1\r\nHost: localhost\r\nX-Custom-Header: ";
huge_header.append(200, 'A');
bufferevent_write(pair[1], huge_header.c_str(), huge_header.size());
event_base_dispatch(base_);

EXPECT_FALSE(open_called);
EXPECT_TRUE(error_called);
bufferevent_free(pair[1]);
}

TEST_F(HandshakeTest, ClientHandshakeRejectsExceededHeaderSize) {
bufferevent* pair[2];
int rv = bufferevent_pair_new(base_, BEV_OPT_CLOSE_ON_FREE | BEV_OPT_DEFER_CALLBACKS, pair);
ASSERT_EQ(rv, 0);
bufferevent_enable(pair[1], EV_READ | EV_WRITE);

bool open_called = false;
bool error_called = false;
// Set tiny max header size limit of 100 bytes
auto client = std::make_unique<ClientHandshake>(
pair[0],
[&](bufferevent* bev) { open_called = true; bufferevent_free(bev); event_base_loopbreak(base_); },
[&]() { error_called = true; event_base_loopbreak(base_); },
100);
client->Start();

// Drain request
int limit = 100;
while (evbuffer_get_length(bufferevent_get_input(pair[1])) == 0 && --limit > 0) {
event_base_loop(base_, EVLOOP_NONBLOCK);
}
std::string req = ReadAllData(pair[1]); (void)req;

// Send incomplete response header larger than 100 bytes
std::string huge_header = "HTTP/1.1 200 OK\r\nContent-Type: application/web-stream\r\nX-Custom: ";
huge_header.append(200, 'B');
bufferevent_write(pair[1], huge_header.c_str(), huge_header.size());
event_base_dispatch(base_);

EXPECT_FALSE(open_called);
EXPECT_TRUE(error_called);
bufferevent_free(pair[1]);
}