From 52c442e5f043b7cf7eef462d803274d8042db052 Mon Sep 17 00:00:00 2001 From: hemanth-openai Date: Sat, 7 Feb 2026 01:09:19 -0800 Subject: [PATCH] Fix ramp halt timing and tests --- src/gabriel/utils/openai_utils.py | 56 +++++++++++++++++++++++----- tests/test_parallelization_tuning.py | 38 ++++++++++++++++++- 2 files changed, 83 insertions(+), 11 deletions(-) diff --git a/src/gabriel/utils/openai_utils.py b/src/gabriel/utils/openai_utils.py index fdb7c3b..25177fe 100644 --- a/src/gabriel/utils/openai_utils.py +++ b/src/gabriel/utils/openai_utils.py @@ -2567,6 +2567,7 @@ async def worker() -> None: while True: if stop_event.is_set(): break + _maybe_emit_ramp_complete() current_cap = _current_parallel_cap() if active_workers < current_cap: break @@ -2977,7 +2978,7 @@ async def get_all_responses( # the API or tool backends. n_parallels: int = 650, ramp_up_seconds: float = 20.0, - ramp_up_start_fraction: float = 0.25, + ramp_up_start_fraction: float = 0.15, max_retries: int = 3, timeout_factor: float = 2.5, max_timeout: Optional[float] = None, @@ -4348,6 +4349,7 @@ def _effective_parallel_ceiling() -> int: last_wait_adjust = 0.0 timeout_notes: Deque[str] = deque(maxlen=3) timeout_errors_since_last_status = 0 + rate_limit_errors_since_last_status = 0 connection_errors_since_last_status = 0 first_timeout_logged = False first_rate_limit_logged = False @@ -4379,6 +4381,7 @@ def _effective_parallel_ceiling() -> int: ramp_up_end_time = ramp_up_start_time + ramp_up_seconds ramp_up_start_cap = max(1, int(math.ceil(concurrency_cap * ramp_up_start_fraction))) ramp_up_halted = False + ramp_up_complete_emitted = False def _maybe_reduce_output_headroom(*, allow_reduce: bool) -> bool: """Lower output headroom after the first token estimate refresh.""" @@ -4474,9 +4477,12 @@ def _log_timeout_once(message: str, note: str, *, dedup_key: Hashable) -> None: def _log_rate_limit_once(detail: Optional[str] = None) -> None: nonlocal first_rate_limit_logged if not first_rate_limit_logged: - logger.warning( + msg = ( "Encountered first rate limit error. Future rate limit errors will be silenced and tracked in periodic updates." ) + logger.warning(msg) + if message_verbose: + print(msg) first_rate_limit_logged = True else: if detail: @@ -4487,9 +4493,12 @@ def _log_rate_limit_once(detail: Optional[str] = None) -> None: def _log_connection_once(detail: Optional[str] = None) -> None: nonlocal first_connection_logged if not first_connection_logged: - logger.warning( + msg = ( "Encountered first connection error. Future connection errors will be silenced and tracked in periodic updates." ) + logger.warning(msg) + if message_verbose: + print(msg) first_connection_logged = True else: if detail: @@ -4501,6 +4510,8 @@ def _halt_ramp_up(reason: str) -> None: nonlocal ramp_up_halted, concurrency_cap if not ramp_up_enabled or ramp_up_halted: return + if time.time() >= ramp_up_end_time: + return ramp_cap = _current_ramp_cap() ramp_up_halted = True if concurrency_cap > ramp_cap: @@ -4553,6 +4564,21 @@ def _trigger_timeout_burst(now: Optional[float] = None) -> None: if drained and message_verbose: print(f"[timeouts] Drained {drained} queued prompts before restart.") + def _maybe_emit_ramp_complete(now: Optional[float] = None) -> None: + nonlocal ramp_up_complete_emitted + if not ramp_up_enabled or ramp_up_halted or ramp_up_complete_emitted: + return + now = time.time() if now is None else now + if now < ramp_up_end_time: + return + ramp_up_complete_emitted = True + msg = ( + f"[parallelization] Ramp-up complete at {concurrency_cap} parallel threads." + ) + logger.info(msg) + if message_verbose: + print(msg) + def _cost_progress_snapshot() -> Optional[Tuple[float, bool]]: """Return total cost so far and whether the value is sampled.""" @@ -4643,9 +4669,9 @@ def emit_parallelization_status( timeout_text = "" connection_text = "" total_completed = processed + denom = max(total_completed, 1) effective_cap = _current_parallel_cap() if status.num_timeout_errors or total_completed: - denom = max(total_completed, 1) timeout_text = f"timeouts={status.num_timeout_errors}/{denom}" if status_report_interval is not None and timeout_errors_since_last_status: timeout_text += f" (+{timeout_errors_since_last_status} since last)" @@ -4664,17 +4690,21 @@ def emit_parallelization_status( total_cost, sampled = cost_snapshot cost_text = f"cost_so_far={'~' if sampled else ''}${total_cost:.2f}" prefix = f"[{label}] " if label else "" - if status.num_connection_errors or connection_errors_since_last_status: - connection_text = f"connection_errors={status.num_connection_errors}" - if status_report_interval is not None and connection_errors_since_last_status: - connection_text += f" (+{connection_errors_since_last_status} since last)" + connection_text = f"connection_errors={status.num_connection_errors}/{denom}" + if status_report_interval is not None and connection_errors_since_last_status: + connection_text += f" (+{connection_errors_since_last_status} since last)" + rate_limit_text = f"rate_limit_errors={status.num_rate_limit_errors}/{denom}" + if status_report_interval is not None and rate_limit_errors_since_last_status: + rate_limit_text += ( + f" (+{rate_limit_errors_since_last_status} since last)" + ) status_bits: List[str] = [ f"cap={effective_cap}", f"active={active_workers}", f"inflight={len(inflight)}", f"queue={queue.qsize()}", f"processed={processed}/{status.num_tasks_started}", - f"rate_limit_errors={status.num_rate_limit_errors}", + rate_limit_text, ] if ramp_up_enabled and not ramp_up_halted and effective_cap < concurrency_cap: status_bits.insert(1, f"ramp_target={concurrency_cap}") @@ -4696,7 +4726,7 @@ def emit_parallelization_status( if ramp_up_enabled and concurrency_cap > 1: ramp_msg = ( f"[parallelization] Ramping up from {ramp_up_start_cap} to {concurrency_cap} " - f"over {int(ramp_up_seconds)}s." + f"parallel threads over {int(ramp_up_seconds)}s." ) logger.info(ramp_msg) if message_verbose: @@ -4918,6 +4948,7 @@ def maybe_adjust_concurrency() -> None: f"after {recent_errors} rate-limit errors in the last {int(round(rate_limit_window))}s." ) logger.warning(reason) + _halt_ramp_up("rate limit recovery") emit_parallelization_status(reason, force=True) else: concurrency_cap = new_cap @@ -4990,6 +5021,7 @@ def maybe_adjust_for_connection_errors() -> None: "or reduce `n_parallels`." ) logger.warning(reason) + _halt_ramp_up("connection errors") emit_parallelization_status(reason, force=True) connection_errors_since_adjust = 0 connection_error_times.clear() @@ -5023,6 +5055,7 @@ async def _maybe_retry(backoff: float) -> bool: async def _handle_rate_limit_error(error_text: str) -> None: nonlocal cooldown_until, rate_limit_errors_since_adjust, successes_since_adjust, processed + nonlocal rate_limit_errors_since_last_status inflight.pop(ident, None) status.num_rate_limit_errors += 1 status.time_of_last_rate_limit_error = time.time() @@ -5031,6 +5064,7 @@ async def _handle_rate_limit_error(error_text: str) -> None: error_logs[ident].append(error_text) rate_limit_error_times.append(time.time()) rate_limit_errors_since_adjust += 1 + rate_limit_errors_since_last_status += 1 successes_since_adjust = 0 _halt_ramp_up("rate limit error") if _is_quota_error_message(error_text): @@ -5842,6 +5876,7 @@ async def status_reporter() -> None: return try: nonlocal timeout_errors_since_last_status, connection_errors_since_last_status + nonlocal rate_limit_errors_since_last_status while not stop_event.is_set(): await asyncio.sleep(status_report_interval) if stop_event.is_set() or processed >= status.num_tasks_started: @@ -5850,6 +5885,7 @@ async def status_reporter() -> None: "Periodic status update", force=True, label=None ) timeout_errors_since_last_status = 0 + rate_limit_errors_since_last_status = 0 connection_errors_since_last_status = 0 except asyncio.CancelledError: pass diff --git a/tests/test_parallelization_tuning.py b/tests/test_parallelization_tuning.py index f599819..c6b6657 100644 --- a/tests/test_parallelization_tuning.py +++ b/tests/test_parallelization_tuning.py @@ -130,7 +130,7 @@ async def responder(prompt: str, **_: object): assert active["peak"] > 2 -def test_ramp_up_halts_on_rate_limit(tmp_path): +def test_ramp_up_halts_on_rate_limit(tmp_path, capsys): active = {"current": 0, "peak": 0} lock = asyncio.Lock() first_error = {"raised": False} @@ -168,4 +168,40 @@ async def responder(prompt: str, **_: object): ) ) + output = capsys.readouterr().out + assert "Halting ramp-up" in output assert active["peak"] <= 3 + + +def test_ramp_up_does_not_halt_after_window(tmp_path, capsys): + first_error = {"raised": False} + + async def responder(prompt: str, **_: object): + if not first_error["raised"]: + first_error["raised"] = True + await asyncio.sleep(0.06) + raise openai_utils.RateLimitError("rate limit") + await asyncio.sleep(0.01) + return [f"ok-{prompt}"], 0.01, [] + + asyncio.run( + openai_utils.get_all_responses( + prompts=[f"p{i}" for i in range(4)], + identifiers=[f"p{i}" for i in range(4)], + response_fn=responder, + use_dummy=False, + save_path=str(tmp_path / "responses.csv"), + reset_files=True, + dynamic_timeout=False, + max_retries=1, + n_parallels=4, + ramp_up_seconds=0.05, + ramp_up_start_fraction=0.25, + status_report_interval=None, + global_cooldown=0, + logging_level="error", + ) + ) + + output = capsys.readouterr().out + assert "Halting ramp-up" not in output