diff --git a/headroom/cli/proxy.py b/headroom/cli/proxy.py index e5693862d..5d0c8aeb5 100644 --- a/headroom/cli/proxy.py +++ b/headroom/cli/proxy.py @@ -373,6 +373,26 @@ def dashboard(port: int, no_open: bool) -> None: "Env: HEADROOM_RETRY_MAX_ATTEMPTS." ), ) +@click.option( + "--retry-base-delay-ms", + type=click.IntRange(min=0), + default=None, + envvar="HEADROOM_RETRY_BASE_DELAY_MS", + help=( + "Initial upstream retry delay in milliseconds (minimum: 0, default: 1000). " + "Env: HEADROOM_RETRY_BASE_DELAY_MS." + ), +) +@click.option( + "--retry-max-delay-ms", + type=click.IntRange(min=0), + default=None, + envvar="HEADROOM_RETRY_MAX_DELAY_MS", + help=( + "Maximum upstream retry delay in milliseconds (minimum: 0, default: 30000). " + "Env: HEADROOM_RETRY_MAX_DELAY_MS." + ), +) @click.option( "--request-timeout-seconds", type=int, @@ -876,6 +896,8 @@ def proxy( no_subscription_tracking: bool, subscription_poll_interval: int | None, retry_max_attempts: int | None, + retry_base_delay_ms: int | None, + retry_max_delay_ms: int | None, request_timeout_seconds: int | None, connect_timeout_seconds: int | None, anthropic_buffered_request_timeout_seconds: int | None, @@ -1129,6 +1151,8 @@ def proxy( subscription_poll_interval if subscription_poll_interval is not None else 300 ), retry_max_attempts=retry_max_attempts if retry_max_attempts is not None else 3, + retry_base_delay_ms=retry_base_delay_ms if retry_base_delay_ms is not None else 1000, + retry_max_delay_ms=retry_max_delay_ms if retry_max_delay_ms is not None else 30000, request_timeout_seconds=request_timeout_seconds if request_timeout_seconds is not None and request_timeout_seconds > 0 else 300, diff --git a/tests/test_cli_proxy_improvements.py b/tests/test_cli_proxy_improvements.py index c10263d26..d6787da13 100644 --- a/tests/test_cli_proxy_improvements.py +++ b/tests/test_cli_proxy_improvements.py @@ -193,6 +193,23 @@ class TestRetryMaxAttemptsValidation: assert result.exit_code != 0 +class TestRetryDelayValidation: + def test_retry_delays_are_forwarded(self, runner: CliRunner, mock_run_server: dict) -> None: + result = runner.invoke( + main, + ["proxy", "--retry-base-delay-ms", "250", "--retry-max-delay-ms", "5000"], + catch_exceptions=False, + ) + assert result.exit_code == 0, result.output + assert mock_run_server["config"].retry_base_delay_ms == 250 + assert mock_run_server["config"].retry_max_delay_ms == 5000 + + @pytest.mark.parametrize("option", ["--retry-base-delay-ms", "--retry-max-delay-ms"]) + def test_negative_delay_is_rejected(self, runner: CliRunner, option: str) -> None: + result = runner.invoke(main, ["proxy", option, "-1"]) + assert result.exit_code != 0 + + class TestConnectTimeoutSecondsValidation: """--connect-timeout-seconds should accept 1-300, reject outside that range.""" @@ -350,6 +367,20 @@ class TestNewEnvVarWiring: assert result.exit_code == 0, result.output assert mock_run_server["config"].retry_max_attempts == 5 + def test_headroom_retry_delays_from_env(self, runner: CliRunner, mock_run_server: dict) -> None: + result = runner.invoke( + main, + ["proxy"], + env={ + "HEADROOM_RETRY_BASE_DELAY_MS": "125", + "HEADROOM_RETRY_MAX_DELAY_MS": "8000", + }, + catch_exceptions=False, + ) + assert result.exit_code == 0, result.output + assert mock_run_server["config"].retry_base_delay_ms == 125 + assert mock_run_server["config"].retry_max_delay_ms == 8000 + def test_headroom_connect_timeout_from_env( self, runner: CliRunner, mock_run_server: dict ) -> None: