diff --git a/wish/cpp/src/h2_server.cc b/wish/cpp/src/h2_server.cc index ab2ace7..427df0a 100644 --- a/wish/cpp/src/h2_server.cc +++ b/wish/cpp/src/h2_server.cc @@ -96,9 +96,16 @@ void H2Server::AcceptConnCb(evconnlistener* listener, int /*socklen*/, void* ctx) { H2Server* server = static_cast(ctx); - event_base* base = evconnlistener_get_base(listener); + if (server->max_connections_ > 0 && server->active_sessions_count_ >= server->max_connections_) { + VLOG(1) << "H2Server: Max connections limit reached (" << server->max_connections_ << "). Rejecting connection."; + + evutil_closesocket(fd); + + return; + } + int one = 1; int set_rv = setsockopt(fd, IPPROTO_TCP, @@ -145,6 +152,7 @@ void H2Server::AcceptConnCb(evconnlistener* listener, return; } + server->active_sessions_count_++; // Send server connection preface (SETTINGS frame). nghttp2_settings_entry iv[] = { @@ -224,6 +232,9 @@ void H2Server::HandleSessionError(Session* sess) { if (!sess) { return; } + if (sess->server && sess->server->active_sessions_count_ > 0) { + sess->server->active_sessions_count_--; + } for (auto& [sid, info] : sess->incoming_streams) { if (info.web_stream) { info.web_stream->OnError(); diff --git a/wish/cpp/src/h2_server.h b/wish/cpp/src/h2_server.h index 3cae974..a2adeca 100644 --- a/wish/cpp/src/h2_server.h +++ b/wish/cpp/src/h2_server.h @@ -28,6 +28,8 @@ #include "nghttp2_web_stream.h" +constexpr size_t kDefaultMaxH2Connections = 10000; + // H2Server listens for plain (cleartext) HTTP/2 (h2c) connections and // exposes each incoming web-stream stream to the caller. class H2Server { @@ -43,6 +45,8 @@ class H2Server { bool Init(); int GetPort() const; void SetOnStream(StreamCallback cb); + void SetMaxConnections(size_t max_connections) { max_connections_ = max_connections; } + size_t active_connections() const { return active_sessions_count_; } int Run(); private: @@ -129,6 +133,9 @@ class H2Server { evconnlistener* listener_; StreamCallback on_stream_; + + size_t max_connections_ = kDefaultMaxH2Connections; + size_t active_sessions_count_ = 0; }; #endif // WISH_CPP_SRC_H2_SERVER_H_ diff --git a/wish/cpp/src/h2_server_test.cc b/wish/cpp/src/h2_server_test.cc index 2068d3c..da3d9b3 100644 --- a/wish/cpp/src/h2_server_test.cc +++ b/wish/cpp/src/h2_server_test.cc @@ -70,3 +70,31 @@ TEST_F(H2ServerTest, InvalidConnectionPrefaceCleanlyDeallocatesSession) { close(client_fd); event_base_loop(base_, EVLOOP_NONBLOCK); } + +TEST_F(H2ServerTest, MaxConnectionsLimitRejectsExtraConnections) { + H2Server server(base_, 0); + server.SetMaxConnections(1); + ASSERT_TRUE(server.Init()); + + sockaddr_in addr{}; + addr.sin_family = AF_INET; + addr.sin_port = htons(server.GetPort()); + inet_pton(AF_INET, "127.0.0.1", &addr.sin_addr); + + // Connection 1 + int fd1 = socket(AF_INET, SOCK_STREAM, 0); + ASSERT_GE(fd1, 0); + ASSERT_EQ(connect(fd1, reinterpret_cast(&addr), sizeof(addr)), 0); + event_base_loop(base_, EVLOOP_NONBLOCK); + EXPECT_EQ(server.active_connections(), 1u); + + // Connection 2 (Should be rejected) + int fd2 = socket(AF_INET, SOCK_STREAM, 0); + ASSERT_GE(fd2, 0); + ASSERT_EQ(connect(fd2, reinterpret_cast(&addr), sizeof(addr)), 0); + event_base_loop(base_, EVLOOP_NONBLOCK); + EXPECT_EQ(server.active_connections(), 1u); + + close(fd1); + close(fd2); +} diff --git a/wish/cpp/src/h2_tls_server.cc b/wish/cpp/src/h2_tls_server.cc index 26b3aea..d92cdd6 100644 --- a/wish/cpp/src/h2_tls_server.cc +++ b/wish/cpp/src/h2_tls_server.cc @@ -119,9 +119,16 @@ void H2TlsServer::AcceptConnCb(evconnlistener* listener, int /*socklen*/, void* ctx) { H2TlsServer* server = static_cast(ctx); - event_base* base = evconnlistener_get_base(listener); + if (server->max_connections_ > 0 && server->active_sessions_count_ >= server->max_connections_) { + VLOG(1) << "H2TlsServer: Max connections limit reached (" << server->max_connections_ << "). Rejecting connection."; + + evutil_closesocket(fd); + + return; + } + int one = 1; int set_rv = setsockopt(fd, IPPROTO_TCP, @@ -175,6 +182,7 @@ void H2TlsServer::AcceptConnCb(evconnlistener* listener, return; } + server->active_sessions_count_++; nghttp2_settings_entry iv[] = { {NGHTTP2_SETTINGS_MAX_CONCURRENT_STREAMS, 100}, @@ -252,6 +260,9 @@ void H2TlsServer::HandleSessionError(Session* sess) { if (!sess) { return; } + if (sess->server && sess->server->active_sessions_count_ > 0) { + sess->server->active_sessions_count_--; + } for (auto& [sid, info] : sess->incoming_streams) { if (info.web_stream) { info.web_stream->OnError(); diff --git a/wish/cpp/src/h2_tls_server.h b/wish/cpp/src/h2_tls_server.h index fb6f654..508c430 100644 --- a/wish/cpp/src/h2_tls_server.h +++ b/wish/cpp/src/h2_tls_server.h @@ -32,6 +32,8 @@ // H2TlsServer listens for TLS-encrypted HTTP/2 connections. // ALPN "h2" is advertised so standard HTTP/2 clients can connect. // mTLS is enforced (client certificates are required), matching TlsServer. +constexpr size_t kDefaultMaxH2TlsConnections = 10000; + class H2TlsServer { public: using StreamCallback = std::function; @@ -45,6 +47,8 @@ class H2TlsServer { bool Init(); void SetOnStream(StreamCallback cb); + void SetMaxConnections(size_t max_connections) { max_connections_ = max_connections; } + size_t active_connections() const { return active_sessions_count_; } int Run(); private: @@ -133,6 +137,9 @@ class H2TlsServer { TlsContext tls_ctx_; StreamCallback on_stream_; + + size_t max_connections_ = kDefaultMaxH2TlsConnections; + size_t active_sessions_count_ = 0; }; #endif // WISH_CPP_SRC_H2_TLS_SERVER_H_ diff --git a/wish/cpp/src/plain_server.cc b/wish/cpp/src/plain_server.cc index 14e106e..4ee33a0 100644 --- a/wish/cpp/src/plain_server.cc +++ b/wish/cpp/src/plain_server.cc @@ -90,6 +90,15 @@ void PlainServer::AcceptConnCb(evconnlistener* listener, event_base* base = evconnlistener_get_base(listener); PlainServer* server = static_cast(ctx); + if (server->max_connections_ > 0 && + (server->active_handshakes_.size() + server->active_streams_.size()) >= server->max_connections_) { + VLOG(1) << "Max connections limit reached (" << server->max_connections_ << "). Rejecting connection."; + + evutil_closesocket(fd); + + return; + } + int one = 1; int set_opt_rv = setsockopt(fd, IPPROTO_TCP, diff --git a/wish/cpp/src/plain_server.h b/wish/cpp/src/plain_server.h index 899d941..c1eeaaf 100644 --- a/wish/cpp/src/plain_server.h +++ b/wish/cpp/src/plain_server.h @@ -30,6 +30,8 @@ class ServerHandshake; class BufferEventWebStream; +constexpr size_t kDefaultMaxConnections = 10000; + class PlainServer { public: using StreamCallback = std::function; @@ -40,6 +42,8 @@ class PlainServer { bool Init(); void SetOnStream(StreamCallback cb); + void SetMaxConnections(size_t max_connections) { max_connections_ = max_connections; } + size_t active_connections() const { return active_handshakes_.size() + active_streams_.size(); } int Run(); private: @@ -62,6 +66,8 @@ class PlainServer { StreamCallback on_stream_; + size_t max_connections_ = kDefaultMaxConnections; + std::vector> active_handshakes_; std::vector> active_streams_; }; diff --git a/wish/cpp/src/tls_server.cc b/wish/cpp/src/tls_server.cc index c520bc3..bbd88c3 100644 --- a/wish/cpp/src/tls_server.cc +++ b/wish/cpp/src/tls_server.cc @@ -110,6 +110,15 @@ void TlsServer::AcceptConnCb(evconnlistener* listener, event_base* base = evconnlistener_get_base(listener); TlsServer* server = static_cast(ctx); + if (server->max_connections_ > 0 && + (server->active_handshakes_.size() + server->active_streams_.size()) >= server->max_connections_) { + VLOG(1) << "Max TLS connections limit reached (" << server->max_connections_ << "). Rejecting connection."; + + evutil_closesocket(fd); + + return; + } + int one = 1; int set_opt_rv = setsockopt(fd, IPPROTO_TCP, diff --git a/wish/cpp/src/tls_server.h b/wish/cpp/src/tls_server.h index 62786f7..dda5f3c 100644 --- a/wish/cpp/src/tls_server.h +++ b/wish/cpp/src/tls_server.h @@ -32,6 +32,8 @@ class ServerHandshake; class BufferEventWebStream; +constexpr size_t kDefaultMaxTlsConnections = 10000; + class TlsServer { public: using StreamCallback = std::function; @@ -45,6 +47,8 @@ class TlsServer { bool Init(); void SetOnStream(StreamCallback cb); + void SetMaxConnections(size_t max_connections) { max_connections_ = max_connections; } + size_t active_connections() const { return active_handshakes_.size() + active_streams_.size(); } int Run(); private: @@ -73,6 +77,8 @@ class TlsServer { StreamCallback on_stream_; + size_t max_connections_ = kDefaultMaxTlsConnections; + std::vector> active_handshakes_; std::vector> active_streams_; };