Skip to content
Merged
235 changes: 181 additions & 54 deletions src/WebSocketContext.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -182,10 +182,15 @@ void WebSocketContext::run() {
if (!_cfg.tls.disableHostnameValidation) {
X509_VERIFY_PARAM* param = SSL_get0_param(ssl);
if (param) {
int ret = X509_VERIFY_PARAM_set1_host(param, _cfg.host.c_str(), 0); // No port matching
// Verify IP addresses against IP SANs, hostnames against DNS SANs.
int ret = _cfg.is_ip_address
? X509_VERIFY_PARAM_set1_ip_asc(param, _cfg.host.c_str())
: X509_VERIFY_PARAM_set1_host(param, _cfg.host.c_str(), 0);
if (ret != 1) {
log_error("Failed to set hostname for verification");
sendError(ErrorCode::TLS_INIT_FAILED, "Failed hostname verification setup");
log_error("Failed to set %s for verification",
_cfg.is_ip_address ? "IP address" : "hostname");
sendError(ErrorCode::TLS_INIT_FAILED,
_cfg.is_ip_address ? "Failed IP address verification setup" : "Failed hostname verification setup");
SSL_free(ssl);
_tls.reset();
return;
Expand Down Expand Up @@ -346,7 +351,19 @@ void WebSocketContext::timeoutCallback(evutil_socket_t /*fd*/, short /*event*/,

void WebSocketContext::pingCallback(evutil_socket_t /*fd*/, short /*event*/, void *arg) {
auto* self = static_cast<WebSocketContext*>(arg);
// Heartbeat only runs after the WebSocket upgrade.
if (!self->upgraded.load(std::memory_order_acquire)) return;

// Disconnect after too many unanswered pings.
if (self->pings_outstanding >= MAX_MISSED_PONGS) {
log_error("ping timeout: %d unanswered ping(s)", self->pings_outstanding);
self->sendError(ErrorCode::PING_TIMEOUT, "Ping timeout (no pong)");
self->connection_state.store(ConnectionState::DISCONNECTING, std::memory_order_release);
self->requestLoopExit();
return;
}
self->sendPing();
self->pings_outstanding++;
}

void WebSocketContext::wakeupCallback(evutil_socket_t, short, void* arg) {
Expand Down Expand Up @@ -619,30 +636,122 @@ void WebSocketContext::handleRead(bufferevent* bev) {

if (!upgraded.load()) {

static constexpr size_t MAX_HANDSHAKE_BYTES = 64u * 1024u;

const size_t len = evbuffer_get_length(input);
if (len < 4) return;

std::vector<char> snap(len);
evbuffer_copyout(input, snap.data(), len);
const char* b = snap.data();

// Find end of headers: "\r\n\r\n" (length-bounded)
size_t headerBytes = 0;
for (size_t i = 0; i + 3 < len; ++i) {
if (b[i] == '\r' && b[i+1] == '\n' && b[i+2] == '\r' && b[i+3] == '\n') {
headerBytes = i + 4;
break;
const evbuffer_ptr headerEnd = evbuffer_search(input, "\r\n\r\n", 4, nullptr);

if (headerEnd.pos < 0) {
if (len > MAX_HANDSHAKE_BYTES) {
log_error("handshake response exceeded %zu bytes without header terminator",
static_cast<size_t>(MAX_HANDSHAKE_BYTES));
connection_state.store(ConnectionState::FAILED, std::memory_order_release);
sendError(ErrorCode::CONNECT_FAILED, "handshake response too large");
evbuffer_drain(input, len);
requestLoopExit();
}
return; // wait for more header bytes
}
if (headerBytes == 0) return;

const size_t headerBytes = static_cast<size_t>(headerEnd.pos) + 4;

if (headerBytes > MAX_HANDSHAKE_BYTES) {
log_error("handshake response exceeded %zu bytes",
static_cast<size_t>(MAX_HANDSHAKE_BYTES));

connection_state.store(ConnectionState::FAILED,
std::memory_order_release);

sendError(ErrorCode::CONNECT_FAILED,
"handshake response too large");

evbuffer_drain(input, len);
requestLoopExit();
return;
}

// Copy only the HTTP handshake headers, not any WebSocket data
// that may already follow in the same input buffer.
std::vector<char> snap(headerBytes);
evbuffer_copyout(input, snap.data(), headerBytes);

const char* b = snap.data();

std::string resp(b, headerBytes);

//log_debug("RESP: %s", resp.c_str());

if (resp.find("HTTP/1.1 101", 0) == std::string::npos ||
!containsHeader(resp, "Sec-WebSocket-Accept:"))
// RFC 6455: validate the HTTP 101 status and Sec-WebSocket-Accept value.
auto headerValue = [&resp](const std::string& expectedName) -> std::string {
size_t pos = 0;

while (pos < resp.size()) {
size_t lineEnd = resp.find("\r\n", pos);
if (lineEnd == std::string::npos) {
lineEnd = resp.size();
}

if (lineEnd == pos) {
break;
}

const size_t colon = resp.find(':', pos);

if (colon != std::string::npos && colon < lineEnd) {
std::string headerName = resp.substr(pos, colon - pos);

std::transform(
headerName.begin(),
headerName.end(),
headerName.begin(),
[](unsigned char c) {
return static_cast<char>(std::tolower(c));
});

if (headerName == expectedName) {
size_t valueStart = colon + 1;

while (valueStart < lineEnd &&
(resp[valueStart] == ' ' || resp[valueStart] == '\t')) {
++valueStart;
}

size_t valueEnd = lineEnd;

while (valueEnd > valueStart &&
(resp[valueEnd - 1] == ' ' ||
resp[valueEnd - 1] == '\t')) {
--valueEnd;
}

return resp.substr(valueStart, valueEnd - valueStart);
}
}

if (lineEnd == resp.size()) {
break;
}

pos = lineEnd + 2;
}

return {};
};

const size_t statusEnd = resp.find("\r\n");

const bool is101 = statusEnd != std::string::npos &&
resp.size() >= 12 &&
resp.compare(0, 12, "HTTP/1.1 101") == 0 &&
(statusEnd == 12 || resp[12] == ' ' || resp[12] == '\t');

const std::string acceptValue = headerValue("sec-websocket-accept");

if (!is101 || acceptValue.empty() || acceptValue != accept)
{
log_error("WebSocket upgrade failed");
log_error("WebSocket upgrade failed (status/accept mismatch)");
connection_state.store(ConnectionState::FAILED, std::memory_order_release);
sendError(ErrorCode::CONNECT_FAILED, "WebSocket upgrade failed");
evbuffer_drain(input, len);
Expand All @@ -653,14 +762,9 @@ void WebSocketContext::handleRead(bufferevent* bev) {
bool negotiated = false;

if (_cfg.compression_requested) {
std::string lowerResp = resp;
std::transform(lowerResp.begin(), lowerResp.end(), lowerResp.begin(), ::tolower);
const std::string key = "sec-websocket-extensions:";
size_t extHeaderPos = lowerResp.find(key);
if (extHeaderPos != std::string::npos) {
size_t lineEnd = resp.find("\r\n", extHeaderPos);
if (lineEnd == std::string::npos) lineEnd = resp.size();
std::string extLine = resp.substr(extHeaderPos, lineEnd - extHeaderPos);
const std::string extLine = headerValue("sec-websocket-extensions");

if (!extLine.empty()) {

if (containsHeader(extLine, "permessage-deflate")) {
negotiated = true;
Expand Down Expand Up @@ -795,10 +899,39 @@ void WebSocketContext::flushSendQueue() {
}
}

static bool hasInvalidHandshakeChars(const std::string& s) {
return s.find('\r') != std::string::npos ||
s.find('\n') != std::string::npos ||
s.find('\0') != std::string::npos;
}

void WebSocketContext::sendHandshakeRequest() {
if (!_bev) return;

log_debug("Sending WebSocket handshake request");

if (hasInvalidHandshakeChars(_cfg.uri) ||
hasInvalidHandshakeChars(_cfg.host)) {

log_error("Invalid CR/LF/NUL character in WebSocket URI or host");
sendError(ErrorCode::CONNECT_FAILED,
"Invalid WebSocket handshake configuration");
requestLoopExit();
return;
}

for (const auto& header : _cfg.headers.headers) {
if (hasInvalidHandshakeChars(header.first) ||
hasInvalidHandshakeChars(header.second)) {

log_error("Invalid CR/LF/NUL character in WebSocket header");
sendError(ErrorCode::CONNECT_FAILED,
"Invalid WebSocket handshake header");
requestLoopExit();
return;
}
}

auto out = bufferevent_get_output(_bev);

evbuffer_add_printf(out, "GET %s HTTP/1.1\r\n", _cfg.uri.c_str());
Expand All @@ -807,21 +940,29 @@ void WebSocketContext::sendHandshakeRequest() {
evbuffer_add_printf(out, "Connection:upgrade\r\n");
evbuffer_add_printf(out, "Sec-WebSocket-Key:%s\r\n", key.c_str());
evbuffer_add_printf(out, "Sec-WebSocket-Version:13\r\n");

if (_cfg.compression_requested) {
evbuffer_add_printf(out, "Sec-WebSocket-Extensions:permessage-deflate; client_no_context_takeover; server_no_context_takeover; client_max_window_bits=9\r\n");
evbuffer_add_printf(
out,
"Sec-WebSocket-Extensions:permessage-deflate; "
"client_no_context_takeover; "
"server_no_context_takeover; "
"client_max_window_bits=9\r\n");
}

evbuffer_add_printf(out, "Origin:http://%s:%d\r\n", _cfg.host.c_str(), _cfg.port);

if (!_cfg.headers.headers.empty()) {
for (const auto& header : _cfg.headers.headers) {
evbuffer_add_printf(out, "%s:%s\r\n", header.first.c_str(), header.second.c_str());
}
evbuffer_add_printf(out,
"Origin:http://%s:%d\r\n",
_cfg.host.c_str(),
_cfg.port);

for (const auto& header : _cfg.headers.headers) {
evbuffer_add_printf(out,
"%s:%s\r\n",
header.first.c_str(),
header.second.c_str());
}

evbuffer_add_printf(out, "\r\n");

}

void WebSocketContext::sendError(int error_code, const std::string& error_message) {
Expand Down Expand Up @@ -1048,35 +1189,20 @@ void WebSocketContext::send(evbuffer* buf, const void* raw_data, size_t raw_len,
// ---- Fast masking (single evbuffer_add) ----
uint8_t mask_key[4];

thread_local uint32_t s = 0;
if (s == 0) {
uint64_t t = static_cast<uint64_t>(time(nullptr));
uintptr_t a = reinterpret_cast<uintptr_t>(&s);
s = static_cast<uint32_t>((t ^ (t >> 32) ^ a) | 1u);
}

auto next_u32 = [&]() -> uint32_t {
s += 0x9E3779B9u;
uint32_t z = s;
z ^= z >> 16;
z *= 0x85EBCA6Bu;
z ^= z >> 13;
z *= 0xC2B2AE35u;
z ^= z >> 16;
return z;
};
thread_local std::random_device rd;

uint32_t mask32 = next_u32();
std::memcpy(mask_key, &mask32, 4);
const uint32_t mask32 = static_cast<uint32_t>(rd());
std::memcpy(mask_key, &mask32, sizeof(mask_key));

// Write mask key
evbuffer_add(out, mask_key, 4);
evbuffer_add(out, mask_key, sizeof(mask_key));

// Mask payload into one contiguous buffer, then add once
static thread_local std::vector<uint8_t> masked;
masked.resize(payload_len);

const uint8_t* src = payload_ptr;

for (size_t i = 0; i < payload_len; ++i) {
masked[i] = src[i] ^ mask_key[i & 3];
}
Expand Down Expand Up @@ -1171,6 +1297,7 @@ bool WebSocketContext::rxCompressionEnabled() const {
void WebSocketContext::onRxPong(std::vector<uint8_t>&& payload) {
log_debug("Received pong frame (%zu bytes)", payload.size());
(void)payload;
pings_outstanding = 0;
}

void WebSocketContext::onRxPing(std::vector<uint8_t>&& payload) {
Expand Down
4 changes: 4 additions & 0 deletions src/WebSocketContext.h
Original file line number Diff line number Diff line change
Expand Up @@ -219,6 +219,10 @@ class WebSocketContext: public std::enable_shared_from_this<WebSocketContext>,
struct event *ping_event = nullptr;
struct event *wakeup_event = nullptr;

// Ping liveness
int pings_outstanding = 0;
static constexpr int MAX_MISSED_PONGS = 2;

// Sender
struct event *send_event = nullptr;
std::atomic_bool send_flush_pending{false};
Expand Down
28 changes: 28 additions & 0 deletions src/WebSocketReceiver.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -279,6 +279,11 @@ bool WebSocketReceiver::rxInflate(const uint8_t* in, size_t in_len, std::vector<

if (produced) {
const size_t old = out.size();
if (old > MAX_MESSAGE_SIZE || produced > MAX_MESSAGE_SIZE - old) {
log_error("permessage-deflate output exceeds %zu bytes; aborting (possible decompression bomb)",
static_cast<size_t>(MAX_MESSAGE_SIZE));
return false;
}
out.resize(old + produced);
std::memcpy(out.data() + old, tmp, produced);
}
Expand Down Expand Up @@ -330,6 +335,12 @@ void WebSocketReceiver::onData(evbuffer* buf) {
return;
}

// RFC 7692: RSV1 is invalid on control and continuation frames.
if (rsv1 && ((opcode & 0x08) != 0 || opcode == 0x00)) {
_sinks.onRxProtocolError(1002, "RSV1 set on control or continuation frame");
return;
}

if ((opcode & 0x08) != 0 && !fin) {
_sinks.onRxProtocolError(1002, "Control frame fragmented");
return;
Expand Down Expand Up @@ -359,6 +370,11 @@ void WebSocketReceiver::onData(evbuffer* buf) {
return;
}

if ((opcode & 0x08) == 0 && payload_len > MAX_MESSAGE_SIZE) {
_sinks.onRxProtocolError(1009, "Frame payload too large");
return;
}

const size_t need = header_len + static_cast<size_t>(payload_len);
if (data_len < need) break; // wait for full frame

Expand Down Expand Up @@ -407,6 +423,18 @@ void WebSocketReceiver::handleContinuationFrame(const unsigned char* payload, si
return;
}

const size_t current_size = fragmented_message.size();

if (current_size > MAX_MESSAGE_SIZE ||
payload_len > static_cast<uint64_t>(MAX_MESSAGE_SIZE - current_size))
{
log_error("reassembled message exceeds %zu bytes; aborting",
static_cast<size_t>(MAX_MESSAGE_SIZE));

_sinks.onRxProtocolError(1009, "Message too large");
return;
}

fragmented_message.insert(fragmented_message.end(),
payload,
payload + payload_len);
Expand Down
2 changes: 2 additions & 0 deletions src/WebSocketReceiver.h
Original file line number Diff line number Diff line change
Expand Up @@ -79,4 +79,6 @@ class WebSocketReceiver {

Utf8Validator utf8Validator;
bool isValidUtf8(const char *str, size_t len);

static constexpr size_t MAX_MESSAGE_SIZE = 16u * 1024u * 1024u;
};
Loading