From 0dd1a63269fa82d46cfc017b5218f7a305abaec5 Mon Sep 17 00:00:00 2001 From: Googler Date: Tue, 8 Sep 2026 16:04:27 -0700 Subject: [PATCH] Implement exponential backoff circuit breaker for failing producer peers in tpu_sync. When a prefill node experiences failures (e.g. frozen process, unresponsive socket, or crash), incoming requests targeting that producer can consume worker threads in push_pool_, potentially starving healthy producers. * Add a 2-state exponential backoff Circuit Breaker (CLOSED <-> OPEN) per peer endpoint: - Trips from CLOSED to OPEN after 2 consecutive connection/handshake failures. - Initial ban duration is 2 minutes, doubling on repeated failures up to 16 minutes (2m -> 4m -> 8m -> 16m). - Automatically transitions back to CLOSED upon expiry of the ban deadline. - Fast-fails incoming reads for banned peers in <1us without occupying threads in push_pool_. - Straggler protection: in-flight requests that fail while the circuit is already OPEN do not trigger additional backoff escalation. * Guard configuration behind opt-in environment variable (disabled by default): - Environment variable: `TPU_RAIDEN_ENABLE_CIRCUIT_BREAKER` ("1" / "true"). - Programmatic setter/getter: `set_circuit_breaker_enabled(bool)` / `circuit_breaker_enabled()`. * Add unit tests verifying fast-failover, non-starvation of healthy peers when enabled, disabled-by-default behavior, and straggler protection. PiperOrigin-RevId: 978163735 --- .../core/kv_cache_manager_with_transfer.cc | 80 +++- .../core/kv_cache_manager_with_transfer.h | 51 +++ ...ache_manager_with_transfer_control_test.cc | 405 +++++++++++++++++- 3 files changed, 532 insertions(+), 4 deletions(-) diff --git a/tpu_sync/core/kv_cache_manager_with_transfer.cc b/tpu_sync/core/kv_cache_manager_with_transfer.cc index 77aba54bc..09cd0390e 100644 --- a/tpu_sync/core/kv_cache_manager_with_transfer.cc +++ b/tpu_sync/core/kv_cache_manager_with_transfer.cc @@ -16,6 +16,7 @@ #include "tpu_sync/core/kv_cache_manager_with_transfer.h" #include +#include #include #include #include @@ -34,6 +35,7 @@ #include #include #include +#include #include #include #include @@ -278,7 +280,8 @@ int ConnectTcp(const std::string& endpoint, double timeout_s) { std::string err_str = std::strerror(errno); close(fd); freeaddrinfo(res); - throw std::runtime_error("connect(" + endpoint + ") failed: " + err_str); + throw std::runtime_error( + absl::StrCat("connect(", endpoint, ") failed: ", err_str)); } freeaddrinfo(res); return fd; @@ -522,6 +525,7 @@ KVCacheManagerWithTransfer::KVCacheManagerWithTransfer( timeout_s_(timeout_s), unsafe_skip_buffer_lock_(unsafe_skip_buffer_lock), metrics_collector_(std::move(metrics_collector)) { + circuit_breaker_enabled_ = IsCircuitBreakerEnvEnabled(); if (local_control_port_ >= 0) { if (max_blocks_ <= 0) { throw std::invalid_argument("max_blocks must be positive"); @@ -577,6 +581,7 @@ KVCacheManagerWithTransfer::KVCacheManagerWithTransfer( timeout_s_(timeout_s), unsafe_skip_buffer_lock_(unsafe_skip_buffer_lock), metrics_collector_(std::move(metrics_collector)) { + circuit_breaker_enabled_ = IsCircuitBreakerEnvEnabled(); if (num_layers() == 0 || num_shards() == 0) { return; } @@ -618,7 +623,9 @@ KVCacheManagerWithTransfer::KVCacheManagerWithTransfer( num_layers, num_shards, std::vector(num_layers, slice_byte_size), local_port, host_blocks_to_allocate, parallelism, node_id, local_control_port, - max_blocks, num_slots, timeout_s, std::move(metrics_collector)) {} + max_blocks, num_slots, timeout_s, std::move(metrics_collector)) { + circuit_breaker_enabled_ = IsCircuitBreakerEnvEnabled(); +} KVCacheManagerWithTransfer::KVCacheManagerWithTransfer( size_t num_layers, size_t num_shards, std::vector slice_byte_sizes, @@ -638,6 +645,7 @@ KVCacheManagerWithTransfer::KVCacheManagerWithTransfer( timeout_s_(timeout_s), unsafe_skip_buffer_lock_(false), metrics_collector_(std::move(metrics_collector)) { + circuit_breaker_enabled_ = IsCircuitBreakerEnvEnabled(); if (local_control_port_ >= 0) { if (max_blocks_ <= 0) { throw std::invalid_argument("max_blocks must be positive"); @@ -1802,12 +1810,76 @@ void KVCacheManagerWithTransfer::StartRead( local_block_ids, parallelism, local_host_block_ids); } +bool KVCacheManagerWithTransfer::IsCircuitBreakerEnvEnabled() { + const char* env = std::getenv("TPU_RAIDEN_ENABLE_CIRCUIT_BREAKER"); + if (env != nullptr) { + return absl::EqualsIgnoreCase(env, "1") || + absl::EqualsIgnoreCase(env, "true"); + } + return false; +} + +bool KVCacheManagerWithTransfer::IsPeerBanned( + absl::string_view remote_endpoint) { + absl::MutexLock lock(circuit_breaker_mu_); + if (!circuit_breaker_enabled_) return false; + auto it = circuit_breakers_.find(remote_endpoint); + if (it != circuit_breakers_.end() && absl::Now() < it->second.banned_until) { + return true; + } + return false; +} + +void KVCacheManagerWithTransfer::RecordCircuitBreakerSuccess( + absl::string_view remote_endpoint) { + absl::MutexLock lock(circuit_breaker_mu_); + if (!circuit_breaker_enabled_) return; + auto it = circuit_breakers_.find(remote_endpoint); + if (it != circuit_breakers_.end()) { + it->second.consecutive_failures = 0; + it->second.current_backoff = circuit_breaker_initial_backoff_; + it->second.banned_until = absl::InfinitePast(); + } +} + +void KVCacheManagerWithTransfer::RecordCircuitBreakerFailure( + absl::string_view remote_endpoint) { + absl::MutexLock lock(circuit_breaker_mu_); + if (!circuit_breaker_enabled_) return; + auto& cb = circuit_breakers_[remote_endpoint]; + if (cb.current_backoff == absl::ZeroDuration()) { + cb.current_backoff = circuit_breaker_initial_backoff_; + } + // Straggler Protection: If already OPEN (banned), ignore stragglers! + if (absl::Now() < cb.banned_until) { + return; + } + if (++cb.consecutive_failures >= 2) { + cb.banned_until = absl::Now() + cb.current_backoff; + LOG(WARNING) << "Peer " << remote_endpoint + << " failed 2 consecutive times. " + << "Banning for " << cb.current_backoff << "!"; + // Exponential backoff: double up to 16 minutes for the next ban + cb.current_backoff = std::min(cb.current_backoff * 2, absl::Minutes(16)); + cb.consecutive_failures = 0; + } +} + void KVCacheManagerWithTransfer::StartRead( const std::string& req_id, uint64_t uuid, const std::string& remote_endpoint, const std::vector& remote_block_ids, const std::vector& local_block_ids, int parallelism, std::optional> local_host_block_ids) { + if (IsPeerBanned(remote_endpoint)) { + LOG(WARNING) + << "StartRead: peer " << remote_endpoint + << " is currently BANNED by circuit breaker. Fast-failing req_id=" + << req_id; + absl::MutexLock lock(mu_); + failed_recving_.insert(req_id); + return; + } LOG(INFO) << "StartRead (initiate): req_id=" << req_id << ", uuid=" << uuid << ", numa=" << assigned_numa_node().value_or(-1); VLOG(1) << "KVCacheManagerWithTransfer::StartRead (Hybrid Bridge) called. " @@ -1949,6 +2021,7 @@ void KVCacheManagerWithTransfer::StartRead( throw std::runtime_error( "Remote producer rejected Hybrid Bridge read request"); } + RecordCircuitBreakerSuccess(remote_endpoint); VLOG(1) << "StartRead (Hybrid Bridge) successfully registered pull " "request with Producer. req_id: " << req_id; @@ -1956,6 +2029,7 @@ void KVCacheManagerWithTransfer::StartRead( LOG(ERROR) << "Raiden consumer error during Hybrid Bridge StartRead connect: " << e.what(); + RecordCircuitBreakerFailure(remote_endpoint); absl::MutexLock lock(mu_); failed_recving_.insert(req_id); auto it = active_recv_entries_.find(uuid); @@ -3348,7 +3422,7 @@ absl::Status KVCacheManagerWithTransfer::OnLayerReceived(size_t layer_idx, } auto future = future_or.value(); - future.OnReady([this, uuid, layer_idx, recv_slot, req_id, + future.OnReady([this, uuid, layer_idx, req_id, metrics_collector = metrics_collector_](auto status_or) { bool unregister_plan = false; uint64_t plan_generation = 0; diff --git a/tpu_sync/core/kv_cache_manager_with_transfer.h b/tpu_sync/core/kv_cache_manager_with_transfer.h index 96d3e09d9..76bb724e5 100644 --- a/tpu_sync/core/kv_cache_manager_with_transfer.h +++ b/tpu_sync/core/kv_cache_manager_with_transfer.h @@ -38,7 +38,9 @@ #include "absl/container/flat_hash_map.h" #include "absl/status/status.h" #include "absl/status/statusor.h" +#include "absl/strings/string_view.h" #include "absl/synchronization/mutex.h" +#include "absl/time/time.h" #include "absl/types/span.h" #include "xla/pjrt/pjrt_client.h" #include "tpu_sync/common/trace.h" @@ -265,6 +267,55 @@ class KVCacheManagerWithTransfer : public kv_cache::KVCacheManagerBase { virtual int local_control_port() const { return local_control_port_; } virtual int64_t node_id() const { return node_id_; } + struct CircuitBreakerState { + int consecutive_failures = 0; + absl::Duration current_backoff = absl::ZeroDuration(); + absl::Time banned_until = absl::InfinitePast(); + }; + + mutable absl::Mutex circuit_breaker_mu_; + absl::flat_hash_map circuit_breakers_ + ABSL_GUARDED_BY(circuit_breaker_mu_); + absl::Duration circuit_breaker_initial_backoff_ + ABSL_GUARDED_BY(circuit_breaker_mu_) = absl::Minutes(2); + bool circuit_breaker_enabled_ ABSL_GUARDED_BY(circuit_breaker_mu_) = false; + + static bool IsCircuitBreakerEnvEnabled(); + + void RecordCircuitBreakerSuccess(absl::string_view remote_endpoint); + void RecordCircuitBreakerFailure(absl::string_view remote_endpoint); + bool IsPeerBanned(absl::string_view remote_endpoint); + + void set_circuit_breaker_enabled(bool enabled) { + absl::MutexLock lock(circuit_breaker_mu_); + circuit_breaker_enabled_ = enabled; + } + bool circuit_breaker_enabled() const { + absl::MutexLock lock(circuit_breaker_mu_); + return circuit_breaker_enabled_; + } + + void set_circuit_breaker_initial_backoff(absl::Duration d) { + absl::MutexLock lock(circuit_breaker_mu_); + circuit_breaker_initial_backoff_ = d; + } + absl::Duration GetPeerCurrentBackoff( + absl::string_view remote_endpoint) const { + absl::MutexLock lock(circuit_breaker_mu_); + auto it = circuit_breakers_.find(remote_endpoint); + if (it != circuit_breakers_.end() && + it->second.current_backoff != absl::ZeroDuration()) { + return it->second.current_backoff; + } + return circuit_breaker_initial_backoff_; + } + absl::Time GetPeerBannedUntil(absl::string_view remote_endpoint) { + absl::MutexLock lock(circuit_breaker_mu_); + auto it = circuit_breakers_.find(remote_endpoint); + if (it != circuit_breakers_.end()) return it->second.banned_until; + return absl::InfinitePast(); + } + protected: std::vector BuildEndpoints(int64_t port) const; diff --git a/tpu_sync/core/kv_cache_manager_with_transfer_control_test.cc b/tpu_sync/core/kv_cache_manager_with_transfer_control_test.cc index 505aefb42..403e08b1b 100644 --- a/tpu_sync/core/kv_cache_manager_with_transfer_control_test.cc +++ b/tpu_sync/core/kv_cache_manager_with_transfer_control_test.cc @@ -22,7 +22,6 @@ #include #include -#include #include #include #include @@ -35,8 +34,10 @@ #include #include +#include "absl/base/thread_annotations.h" #include "absl/log/log.h" #include "absl/strings/str_cat.h" +#include "absl/synchronization/mutex.h" #include "absl/time/clock.h" #include "absl/time/time.h" #include "tpu_sync/core/kv_cache_manager_with_transfer.h" @@ -281,5 +282,407 @@ TEST(ControlHandshakeTest, ConsumerGivesUpOnProducerThatNeverAnswers) { } } +class DynamicFailingProducer { + public: + enum class State { + HEALTHY, + SILENT_FREEZE, + CRASH_WITH_RST, + }; + + DynamicFailingProducer() { + listen_fd_ = socket(AF_INET, SOCK_STREAM, 0); + EXPECT_GE(listen_fd_, 0); + int opt = 1; + setsockopt(listen_fd_, SOL_SOCKET, SO_REUSEADDR, &opt, sizeof(opt)); + sockaddr_in addr{}; + addr.sin_family = AF_INET; + addr.sin_addr.s_addr = htonl(INADDR_LOOPBACK); + addr.sin_port = 0; + EXPECT_EQ( + bind(listen_fd_, reinterpret_cast(&addr), sizeof(addr)), 0); + socklen_t len = sizeof(addr); + EXPECT_EQ(getsockname(listen_fd_, reinterpret_cast(&addr), &len), + 0); + port_ = ntohs(addr.sin_port); + EXPECT_EQ(listen(listen_fd_, 64), 0); + + listener_thread_ = std::thread([this] { AcceptLoop(); }); + } + + ~DynamicFailingProducer() { Stop(); } + + void SetState(State state) { + state_.store(state, std::memory_order_relaxed); + if (state == State::CRASH_WITH_RST) { + absl::MutexLock lock(mu_); + for (int client : clients_) { + linger sl{.l_onoff = 1, .l_linger = 0}; + setsockopt(client, SOL_SOCKET, SO_LINGER, &sl, sizeof(sl)); + close(client); + } + clients_.clear(); + if (listen_fd_ >= 0) { + shutdown(listen_fd_, SHUT_RDWR); + close(listen_fd_); + listen_fd_ = -1; + } + } + } + + void Stop() { + if (listen_fd_ >= 0) { + shutdown(listen_fd_, SHUT_RDWR); + close(listen_fd_); + listen_fd_ = -1; + } + if (listener_thread_.joinable()) { + listener_thread_.join(); + } + absl::MutexLock lock(mu_); + for (int client : clients_) { + close(client); + } + clients_.clear(); + for (auto& t : worker_threads_) { + if (t.joinable()) t.join(); + } + worker_threads_.clear(); + } + + std::string endpoint() const { return absl::StrCat("127.0.0.1:", port_); } + int accepted_count() const { return accepted_.load(); } + int served_count() const { return served_.load(); } + + private: + void AcceptLoop() { + while (true) { + int client = accept(listen_fd_, nullptr, nullptr); + if (client < 0) return; + ++accepted_; + absl::MutexLock lock(mu_); + clients_.push_back(client); + worker_threads_.emplace_back([this, client] { HandleClient(client); }); + } + } + + void HandleClient(int client_fd) { + State current = state_.load(std::memory_order_relaxed); + if (current == State::CRASH_WITH_RST) { + linger sl{.l_onoff = 1, .l_linger = 0}; + setsockopt(client_fd, SOL_SOCKET, SO_LINGER, &sl, sizeof(sl)); + close(client_fd); + return; + } + if (current == State::SILENT_FREEZE) { + // Hold connection open without writing anything back + return; + } + + // State::HEALTHY: Read request header and return valid + // ControlResponseHeader + TestManager::ControlRequestHeader req; + if (!ReadAll(client_fd, &req, sizeof(req))) return; + std::vector dummy(req.num_blocks * 2); + if (!ReadAll(client_fd, dummy.data(), dummy.size() * sizeof(int64_t))) + return; + + TestManager::ControlResponseHeader resp; + resp.magic = TestManager::kResponseMagic; + resp.status = 0; + resp.message_len = 0; + WriteAll(client_fd, &resp, sizeof(resp)); + ++served_; + } + + int listen_fd_ = -1; + int port_ = 0; + std::atomic state_{State::HEALTHY}; + std::atomic accepted_{0}; + std::atomic served_{0}; + absl::Mutex mu_; + std::vector clients_ ABSL_GUARDED_BY(mu_); + std::vector worker_threads_ ABSL_GUARDED_BY(mu_); + std::thread listener_thread_; +}; + +TEST(ControlHandshakeTest, + CircuitBreakerBansFailingPeerWithExponentialBackoff) { + DynamicFailingProducer p0; + DynamicFailingProducer p1; + TestManager consumer; + consumer.set_circuit_breaker_enabled(true); + consumer.set_circuit_breaker_initial_backoff(absl::Seconds(2)); + + // Freeze p0. + p0.SetState(DynamicFailingProducer::State::SILENT_FREEZE); + + // Send 2 requests to p0. + consumer.StartRead("req_p0_0", /*uuid=*/100, p0.endpoint(), + /*remote_block_ids=*/{0}, /*local_block_ids=*/{0}); + consumer.StartRead("req_p0_1", /*uuid=*/101, p0.endpoint(), + /*remote_block_ids=*/{0}, /*local_block_ids=*/{0}); + + // Complete them (they fail after timeout). + const absl::Time t_fail = absl::Now(); + int fail_count = 0; + while (SecondsSince(t_fail) < 5.0) { + auto [done_sending, done_recving, failed_recving] = + consumer.CompleteReadRaw(); + fail_count += failed_recving.size(); + if (fail_count >= 2) break; + std::this_thread::sleep_for(std::chrono::milliseconds(10)); + } + ASSERT_GE(fail_count, 2); + + // Assert EXPECT_TRUE(consumer.IsPeerBanned(p0.endpoint())); + EXPECT_TRUE(consumer.IsPeerBanned(p0.endpoint())); + const absl::Time banned_until = consumer.GetPeerBannedUntil(p0.endpoint()); + EXPECT_GT(banned_until, absl::Now()); + + // Send a 3rd request to p0. Call consumer.CompleteReadRaw(). Assert that it + // is in failed_recving immediately (<10ms)! + const absl::Time t_fast_fail = absl::Now(); + consumer.StartRead("req_p0_2", /*uuid=*/102, p0.endpoint(), + /*remote_block_ids=*/{0}, /*local_block_ids=*/{0}); + auto [s_3, r_3, f_3] = consumer.CompleteReadRaw(); + const double fast_fail_ms = + absl::ToDoubleMilliseconds(absl::Now() - t_fast_fail); + EXPECT_THAT(f_3, Contains("req_p0_2")); + EXPECT_LT(fast_fail_ms, 10.0); + + // Send a request to healthy p1. Assert that it completes successfully in + // <20ms! + const absl::Time t_healthy = absl::Now(); + consumer.StartRead("req_p1_0", /*uuid=*/200, p1.endpoint(), + /*remote_block_ids=*/{0}, /*local_block_ids=*/{0}); + bool p1_done = false; + while (SecondsSince(t_healthy) < 5.0) { + auto [s_p1, r_p1, f_p1] = consumer.CompleteReadRaw(); + for (const auto& id : r_p1) { + if (id == "req_p1_0") p1_done = true; + } + if (p1_done) break; + std::this_thread::sleep_for(std::chrono::milliseconds(2)); + } + const double healthy_ms = absl::ToDoubleMilliseconds(absl::Now() - t_healthy); + EXPECT_TRUE(p1_done); + EXPECT_LT(healthy_ms, 20.0); + + // Verify straggler protection: call + // consumer.RecordCircuitBreakerFailure(p0.endpoint()) while banned, assert + // consumer.GetPeerBannedUntil(p0.endpoint()) was not modified. + consumer.RecordCircuitBreakerFailure(p0.endpoint()); + EXPECT_EQ(consumer.GetPeerBannedUntil(p0.endpoint()), banned_until); +} + +TEST(ControlHandshakeTest, CircuitBreakerDisabledByDefaultDoesNotBan) { + DynamicFailingProducer p0; + TestManager consumer; + EXPECT_FALSE(consumer.circuit_breaker_enabled()); + + p0.SetState(DynamicFailingProducer::State::SILENT_FREEZE); + consumer.StartRead("req_dis_0", /*uuid=*/300, p0.endpoint(), + /*remote_block_ids=*/{0}, /*local_block_ids=*/{0}); + consumer.StartRead("req_dis_1", /*uuid=*/301, p0.endpoint(), + /*remote_block_ids=*/{0}, /*local_block_ids=*/{0}); + + const absl::Time t_fail = absl::Now(); + int fail_count = 0; + while (SecondsSince(t_fail) < 5.0) { + auto [done_sending, done_recving, failed_recving] = + consumer.CompleteReadRaw(); + fail_count += failed_recving.size(); + if (fail_count >= 2) break; + std::this_thread::sleep_for(std::chrono::milliseconds(10)); + } + ASSERT_GE(fail_count, 2); + + // Because the circuit breaker is disabled by default, the peer is not banned. + EXPECT_FALSE(consumer.IsPeerBanned(p0.endpoint())); + EXPECT_EQ(consumer.GetPeerBannedUntil(p0.endpoint()), absl::InfinitePast()); +} + +void RunFailingProducerWithRst(bool circuit_breaker_enabled, + bool expect_peer_banned, + bool expect_starvation) { + DynamicFailingProducer p0; + DynamicFailingProducer p1; + TestManager consumer; + consumer.set_circuit_breaker_enabled(circuit_breaker_enabled); + + p0.SetState(DynamicFailingProducer::State::CRASH_WITH_RST); + + for (int i = 0; i < kPoolSize; ++i) { + consumer.StartRead(absl::StrCat("req_crash_", i), /*uuid=*/100 + i, + p0.endpoint(), /*remote_block_ids=*/{0}, + /*local_block_ids=*/{0}); + } + + const absl::Time t_healthy_start = absl::Now(); + consumer.StartRead("req_healthy", /*uuid=*/999, p1.endpoint(), + /*remote_block_ids=*/{0}, /*local_block_ids=*/{0}); + + while (p1.served_count() == 0 && SecondsSince(t_healthy_start) < 5.0) { + std::this_thread::sleep_for(std::chrono::milliseconds(5)); + } + const double latency = SecondsSince(t_healthy_start); + + EXPECT_EQ(p1.served_count(), 1); + if (expect_starvation) { + EXPECT_GE(latency, 0.4); + } else { + // Sockets fail immediately with TCP RST, so worker threads are freed + // in <1ms without starvation. + EXPECT_LT(latency, 0.2); + } + EXPECT_EQ(consumer.IsPeerBanned(p0.endpoint()), expect_peer_banned); +} + +TEST(ControlHandshakeTest, FailingProducerWithRst) { + RunFailingProducerWithRst(/*circuit_breaker_enabled=*/false, + /*expect_peer_banned=*/false, + /*expect_starvation=*/false); + RunFailingProducerWithRst(/*circuit_breaker_enabled=*/true, + /*expect_peer_banned=*/true, + /*expect_starvation=*/false); +} + +void RunFailingProducerWithoutRst(bool circuit_breaker_enabled, + bool expect_peer_banned, + bool expect_starvation) { + DynamicFailingProducer p0; + DynamicFailingProducer p1; + TestManager consumer; + consumer.set_circuit_breaker_enabled(circuit_breaker_enabled); + + p0.SetState(DynamicFailingProducer::State::SILENT_FREEZE); + + // Dispatch 2 reads to p0 and wait for them to fail. + for (int i = 0; i < 2; ++i) { + consumer.StartRead(absl::StrCat("req_silent_init_", i), /*uuid=*/100 + i, + p0.endpoint(), /*remote_block_ids=*/{0}, + /*local_block_ids=*/{0}); + } + + const absl::Time t_fail_start = absl::Now(); + int fail_count = 0; + while (SecondsSince(t_fail_start) < 5.0) { + auto [done_sending, done_recving, failed_recving] = + consumer.CompleteReadRaw(); + fail_count += failed_recving.size(); + if (fail_count >= 2) break; + std::this_thread::sleep_for(std::chrono::milliseconds(10)); + } + ASSERT_GE(fail_count, 2); + + EXPECT_EQ(consumer.IsPeerBanned(p0.endpoint()), expect_peer_banned); + + // Dispatch 4 reads to p0. + // When banned, these fast-fail in <1us without occupying threads in + // push_pool_. When unbanned, these saturate all 4 worker threads in + // push_pool_. + for (int i = 0; i < kPoolSize; ++i) { + consumer.StartRead(absl::StrCat("req_silent_batch_", i), /*uuid=*/200 + i, + p0.endpoint(), /*remote_block_ids=*/{0}, + /*local_block_ids=*/{0}); + } + + if (expect_starvation) { + // When starvation is expected (peer not banned), wait until p0 has accepted + // the 4 connections so all push_pool_ threads are occupied. + while (p0.accepted_count() < 2 + kPoolSize) { + std::this_thread::sleep_for(std::chrono::milliseconds(5)); + } + } + + // Dispatch read to healthy p1. + const absl::Time t_healthy_start = absl::Now(); + consumer.StartRead("req_healthy", /*uuid=*/999, p1.endpoint(), + /*remote_block_ids=*/{0}, /*local_block_ids=*/{0}); + + while (p1.served_count() == 0 && SecondsSince(t_healthy_start) < 5.0) { + std::this_thread::sleep_for(std::chrono::milliseconds(5)); + } + const double latency = SecondsSince(t_healthy_start); + + EXPECT_EQ(p1.served_count(), 1); + if (expect_starvation) { + // Healthy producer is starved in queue for >= 0.4s until a silent request + // times out (kTimeoutS = 0.5s). + EXPECT_GE(latency, 0.4); + } else { + // Healthy producer is serviced immediately (<0.1s) because circuit breaker + // fast-failed requests to p0 without holding worker threads. + EXPECT_LT(latency, 0.1); + } +} + +TEST(ControlHandshakeTest, FailingProducerWithoutRst) { + RunFailingProducerWithoutRst(/*circuit_breaker_enabled=*/false, + /*expect_peer_banned=*/false, + /*expect_starvation=*/true); + RunFailingProducerWithoutRst(/*circuit_breaker_enabled=*/true, + /*expect_peer_banned=*/true, + /*expect_starvation=*/false); +} + +TEST(ControlHandshakeTest, + ConcurrentWorkerFailuresInOpenCircuitIgnoredWhileClosedCircuitCompounds) { + TestManager consumer; + consumer.set_circuit_breaker_enabled(true); + consumer.set_circuit_breaker_initial_backoff(absl::Milliseconds(200)); + const std::string peer = "127.0.0.1:54321"; + + // 1. In CLOSED circuit: launch 8 concurrent worker threads recording failure. + // The first 2 trip the circuit to OPEN (backoff becomes 400ms, banned for + // 200ms). The remaining 6 run while circuit is already OPEN, so their + // failures are ignored. + { + std::vector workers; + for (int i = 0; i < 8; ++i) { + workers.emplace_back( + [&consumer, &peer]() { consumer.RecordCircuitBreakerFailure(peer); }); + } + for (auto& t : workers) t.join(); + } + + EXPECT_TRUE(consumer.IsPeerBanned(peer)); + // In open circuit, worker thread failures are ignored and backoff remains + // unchanged (400ms for next trip) + EXPECT_EQ(consumer.GetPeerCurrentBackoff(peer), absl::Milliseconds(400)); + const absl::Time banned_until_1 = consumer.GetPeerBannedUntil(peer); + + // Any additional failures while OPEN do not extend banned_until or double + // backoff + consumer.RecordCircuitBreakerFailure(peer); + EXPECT_EQ(consumer.GetPeerBannedUntil(peer), banned_until_1); + EXPECT_EQ(consumer.GetPeerCurrentBackoff(peer), absl::Milliseconds(400)); + + // 2. Wait for the 200ms ban to expire (circuit returns to CLOSED) + while (absl::Now() < banned_until_1) { + std::this_thread::sleep_for(std::chrono::milliseconds(10)); + } + EXPECT_FALSE(consumer.IsPeerBanned(peer)); + + // 3. In CLOSED circuit: things compound! + // 2 consecutive failures now trip with the compounded 400ms backoff, + // and next backoff compounds to 800ms. + consumer.RecordCircuitBreakerFailure(peer); + EXPECT_FALSE(consumer.IsPeerBanned(peer)); // 1 failure: not banned yet + consumer.RecordCircuitBreakerFailure(peer); // 2nd failure: trips! + EXPECT_TRUE(consumer.IsPeerBanned(peer)); + EXPECT_EQ(consumer.GetPeerCurrentBackoff(peer), absl::Milliseconds(800)); + + // 4. On success: resets backoff to initial 200ms + const absl::Time banned_until_2 = consumer.GetPeerBannedUntil(peer); + while (absl::Now() < banned_until_2) { + std::this_thread::sleep_for(std::chrono::milliseconds(10)); + } + EXPECT_FALSE(consumer.IsPeerBanned(peer)); + consumer.RecordCircuitBreakerSuccess(peer); + EXPECT_EQ(consumer.GetPeerCurrentBackoff(peer), absl::Milliseconds(200)); +} + } // namespace } // namespace tpu_raiden