diff --git a/cecli/helpers/llms/litellm_compat.py b/cecli/helpers/llms/litellm_compat.py index 91042c07387..ebe76800c31 100644 --- a/cecli/helpers/llms/litellm_compat.py +++ b/cecli/helpers/llms/litellm_compat.py @@ -395,6 +395,8 @@ def __init__( ) -> None: super().__init__(message or "") self.status_code = status_code + for k, v in kwargs.items(): + setattr(self, k, v) class APIConnectionError(_FacadeException): @@ -508,21 +510,23 @@ def _translate_http_error(err: httpx.HTTPStatusError) -> _FacadeException: if status == 400 and any( token in body for token in ("context", "context_length", "maximum context") ): - return ContextWindowExceededError(message=message, status_code=status) + return ContextWindowExceededError( + message=message, status_code=status, response=err.response + ) if status in (401, 403): - return AuthenticationError(message=message, status_code=status) + return AuthenticationError(message=message, status_code=status, response=err.response) if status == 404: - return NotFoundError(message=message, status_code=status) + return NotFoundError(message=message, status_code=status, response=err.response) if status == 429: - return RateLimitError(message=message, status_code=status) + return RateLimitError(message=message, status_code=status, response=err.response) if status >= 500: - return InternalServerError(message=message, status_code=status) + return InternalServerError(message=message, status_code=status, response=err.response) - return APIError(message=message, status_code=status) + return APIError(message=message, status_code=status, response=err.response) # --------------------------------------------------------------------------- diff --git a/cecli/models.py b/cecli/models.py index eb48c352c69..e0e8eb44c39 100644 --- a/cecli/models.py +++ b/cecli/models.py @@ -12,12 +12,15 @@ import yaml -from cecli import __version__ +from cecli import __version__, utils from cecli.decoding import safe_open from cecli.dump import dump from cecli.exceptions import LiteLLMExceptions from cecli.helpers import coroutines, nested -from cecli.helpers.file_searcher import generate_search_path_list, handle_core_files +from cecli.helpers.file_searcher import ( + generate_search_path_list, + handle_core_files, +) from cecli.helpers.model_config import get_default_config from cecli.helpers.model_config.utils import get_entry_from_raw from cecli.helpers.model_providers import ModelProviderManager @@ -289,7 +292,9 @@ def fetch_openrouter_model_info(self, model): import re if re.search( - f"The model\\s*.*{re.escape(url_part)}.* is not available", html, re.IGNORECASE + f"The model\\s*.*{re.escape(url_part)}.* is not available", + html, + re.IGNORECASE, ): print(f"\x1b[91mError: Model '{url_part}' is not available\x1b[0m") return {} @@ -497,7 +502,8 @@ def __init__( "editor_model", nested.getter(from_model, "editor_model", None) ) editor_edit_format = kwargs.get( - "editor_edit_format", nested.getter(from_model, "editor_edit_format", None) + "editor_edit_format", + nested.getter(from_model, "editor_edit_format", None), ) else: agent_model = kwargs.get("agent_model", None) @@ -538,7 +544,8 @@ def __init__( self.editor_model = None self.agent_model = None self.extra_model_settings = next( - (ms for ms in MODEL_SETTINGS if ms.name == "cecli/extra_params"), None + (ms for ms in MODEL_SETTINGS if ms.name == "cecli/extra_params"), + None, ) self.info = self.get_model_info(model) self.model_config_defaults = get_default_config( @@ -654,18 +661,24 @@ def _apply_structured_kwargs(self, config, model_name): for key, value in config.items(): if key in ("agent", "model_settings", "model-settings"): if not isinstance(value, dict): - raise ValueError(f"override_kwargs '{key}' must be a dict, got {type(value)}") + raise ValueError( + f"override_kwargs '{key}' must be a dict, got" f" {type(value)}" + ) for setting_key, setting_value in value.items(): if setting_key not in valid_model_settings_fields: raise ValueError( - f"Invalid model_settings key '{setting_key}'. " - f"Must be one of: {sorted(valid_model_settings_fields)}" + f"Invalid model_settings key '{setting_key}'. Must" + f" be one of: {sorted(valid_model_settings_fields)}" ) setattr(self, setting_key, setting_value) - elif has_structured_keys and key in ("api", "api_settings", "api-settings"): + elif has_structured_keys and key in ( + "api", + "api_settings", + "api-settings", + ): # api_settings: merge each sub-key into extra_params if not isinstance(value, dict): raise ValueError(f"override_kwargs '{key}' must be a dict, got {type(value)}") @@ -687,11 +700,18 @@ def _apply_structured_kwargs(self, config, model_name): if isinstance(api_value, dict) and isinstance( self.extra_params.get(api_key), dict ): - self.extra_params[api_key] = {**self.extra_params[api_key], **api_value} + self.extra_params[api_key] = { + **self.extra_params[api_key], + **api_value, + } else: self.extra_params[api_key] = api_value - elif has_structured_keys and key in ("llm", "llm_settings", "llm-settings"): + elif has_structured_keys and key in ( + "llm", + "llm_settings", + "llm-settings", + ): # llm_settings: merge into self.info if not isinstance(value, dict): raise ValueError(f"override_kwargs '{key}' must be a dict, got {type(value)}") @@ -1213,7 +1233,10 @@ def set_thinking_tokens(self, value): elif "reasoning" in extra_body: del extra_body["reasoning"] elif num_tokens > 0: - extra_body["thinking"] = {"type": "enabled", "budget_tokens": num_tokens} + extra_body["thinking"] = { + "type": "enabled", + "budget_tokens": num_tokens, + } # extra_body is authoritative; drop any legacy top-level copy. self.extra_params.pop("thinking", None) else: @@ -1322,7 +1345,8 @@ async def send_completion( if effective_tools: sorted_tools = sorted( - effective_tools, key=lambda x: x.get("function", {}).get("name", "Invalid Name") + effective_tools, + key=lambda x: x.get("function", {}).get("name", "Invalid Name"), ) try: @@ -1352,7 +1376,10 @@ async def send_completion( if "name" in function: tool_name = function.get("name") if tool_name: - kwargs["tool_choice"] = {"type": "function", "function": {"name": tool_name}} + kwargs["tool_choice"] = { + "type": "function", + "function": {"name": tool_name}, + } if self.extra_params: kwargs.update(self.extra_params) @@ -1456,10 +1483,15 @@ async def send_completion( if ex_info.name == "ServiceUnavailableError": should_retry = should_retry or self.retry_on_unavailable - if should_retry: + custom_retry_delay = self._extract_gemini_retry_delay(err) + if custom_retry_delay is not None: + retry_delay = custom_retry_delay + should_retry = True + elif should_retry: retry_delay *= self.retry_backoff_factor - if retry_delay > self.retry_timeout: - should_retry = False + + if retry_delay > self.retry_timeout: + should_retry = False # Check for non-retryable RateLimitError within ServiceUnavailableError if ( @@ -1550,10 +1582,16 @@ async def simple_send_with_retries( if ex_info.description: print(ex_info.description) should_retry = ex_info.retry - if should_retry: + custom_retry_delay = self._extract_gemini_retry_delay(err) + if custom_retry_delay is not None: + retry_delay = custom_retry_delay + should_retry = True + elif should_retry: retry_delay *= 2 - if retry_delay > RETRY_TIMEOUT: - should_retry = False + + if retry_delay > RETRY_TIMEOUT: + should_retry = False + if not should_retry: return None print(f"Retrying in {retry_delay:.1f} seconds...") @@ -1574,7 +1612,7 @@ def model_error_response(self): finish_reason="stop", index=0, message=litellm.Message( - content="Model API Response Error. Please retry the previous request" + content=("Model API Response Error. Please retry the previous request") ), ) ], @@ -1584,6 +1622,104 @@ def model_error_response(self): async def model_error_response_stream(self): yield self.model_error_response() + def _extract_gemini_retry_delay(self, err): + """ + Extract suggested retry delay (in seconds) from a 429 rate-limit error. + + 1. Gemini payload body: Google RPC APIs return `google.rpc.RetryInfo` with + a `retryDelay` field (e.g. "15.2s") inside the `error.details` array. + 2. HTTP headers fallback: Standard providers (OpenAI, Anthropic, Groq, etc.) + supply `retry-after` (seconds) or `retry-after-ms` (milliseconds) headers. + """ + status_code = nested.getter(err, ["status_code", "response.status_code"], None) + if isinstance(status_code, (int, str, float)) and str(status_code) != "429": + return None + + # 1. Check Gemini response payload for google.rpc.RetryInfo retryDelay + candidates = [] + response = nested.getter(err, "response") + if callable(getattr(response, "json", None)): + try: + data = response.json() + if isinstance(data, dict): + candidates.append(data) + except Exception: + pass + + sources = [ + response, + nested.getter(err, "message"), + nested.getter(err, "body"), + nested.getter(err, "error"), + nested.getter(err, "raw_response"), + str(err) if isinstance(err, Exception) else None, + err, + ] + + for src in sources: + if not src: + continue + if isinstance(src, dict): + candidates.append(src) + elif isinstance(src, (str, bytes)): + text_str = src.decode("utf-8", errors="ignore") if isinstance(src, bytes) else src + for chunk in utils.split_concatenated_json(text_str): + try: + parsed = json.loads(chunk) + if isinstance(parsed, dict): + candidates.append(parsed) + except Exception: + pass + else: + text = nested.getter(src, ["text", "content"]) + if isinstance(text, (str, bytes)): + text_str = ( + text.decode("utf-8", errors="ignore") if isinstance(text, bytes) else text + ) + for chunk in utils.split_concatenated_json(text_str): + try: + parsed = json.loads(chunk) + if isinstance(parsed, dict): + candidates.append(parsed) + except Exception: + pass + + for candidate in candidates: + if not isinstance(candidate, dict): + continue + + details = nested.getter(candidate, "error.details") + if isinstance(details, list): + for item in details: + delay_val = nested.getter(item, "retryDelay") + if delay_val is not None: + delay_str = str(delay_val).strip() + if delay_str.endswith("s"): + delay_str = delay_str[:-1] + try: + return float(delay_str) + except (ValueError, TypeError): + pass + + # 2. Check HTTP headers fallback (retry-after, retry-after-ms) + headers = nested.getter(err, ["response.headers", "headers"], None) + if headers is not None: + retry_after = nested.getter(headers, ["retry-after"], None) + if retry_after is not None: + try: + return float(str(retry_after).strip()) + except (ValueError, TypeError): + pass + + retry_after_ms = nested.getter(headers, ["retry-after-ms"], None) + if retry_after_ms is not None: + try: + return float(str(retry_after_ms).strip()) / 1000.0 + except (ValueError, TypeError): + pass + + return None + def _log_messages(self, messages, name="message"): """ Log conversation messages to a JSON file. @@ -1591,7 +1727,11 @@ def _log_messages(self, messages, name="message"): os.makedirs(".cecli/logs/messages", exist_ok=True) with safe_open(f".cecli/logs/messages/{name}-{time.time()}.log", "w") as f: json.dump( - messages, f, indent=4, ensure_ascii=False, default=lambda o: "" + messages, + f, + indent=4, + ensure_ascii=False, + default=lambda o: "", ) def _log_request(self, model_call_dict): @@ -1691,8 +1831,8 @@ async def sanity_check_model(io, model): io.tool_output(f"- {key}: {status}") if platform.system() == "Windows": io.tool_output( - "Note: You may need to restart your terminal or command prompt for `setx` to take" - " effect." + "Note: You may need to restart your terminal or command prompt" + " for `setx` to take effect." ) elif not model.keys_in_environment: show = True @@ -1721,7 +1861,10 @@ async def check_for_dependencies(io, model_name): """ if model_name.startswith("bedrock/"): await check_pip_install_extra( - io, "boto3", "AWS Bedrock models require the boto3 package.", ["boto3"] + io, + "boto3", + "AWS Bedrock models require the boto3 package.", + ["boto3"], ) elif model_name.startswith("vertex_ai/"): await check_pip_install_extra( diff --git a/tests/unit/test_retry_backoff.py b/tests/unit/test_retry_backoff.py new file mode 100644 index 00000000000..1d25525de20 --- /dev/null +++ b/tests/unit/test_retry_backoff.py @@ -0,0 +1,536 @@ +import asyncio +import json +from unittest.mock import MagicMock, patch + +import pytest + +from cecli.llm import litellm +from cecli.models import Model + + +def test_extract_gemini_retry_delay_valid(): + model = Model("gemini/gemini-2.5-flash") + + payload = { + "error": { + "code": 429, + "message": "Quota exceeded for metric ... Please retry in 15.2s.", + "status": "RESOURCE_EXHAUSTED", + "details": [ + { + "@type": "://googleapis.com", + "violations": [ + { + "subject": "client_id:your_api_key_or_project", + "description": "Rate limit exceeded.", + } + ], + }, + {"@type": "://googleapis.com", "retryDelay": "15.2s"}, + ], + } + } + + # Exception with response object + err = Exception("Rate limit error") + err.status_code = 429 + mock_resp = MagicMock() + mock_resp.json.return_value = payload + err.response = mock_resp + + delay = model._extract_gemini_retry_delay(err) + assert delay == 15.2 + + +def test_extract_gemini_retry_delay_json_in_message(): + model = Model("gemini/gemini-2.5-flash") + + payload = { + "error": { + "code": 429, + "message": "Quota exceeded", + "details": [{"retryDelay": "8.5s"}], + } + } + + err = Exception(f"APIError: 429 {json.dumps(payload)}") + delay = model._extract_gemini_retry_delay(err) + assert delay == 8.5 + + +def test_extract_gemini_retry_delay_non_429(): + model = Model("gemini/gemini-2.5-flash") + + payload = { + "error": { + "code": 500, + "message": "Internal error", + "details": [{"retryDelay": "15.2s"}], + } + } + + err = Exception("Internal Error") + err.status_code = 500 + mock_resp = MagicMock() + mock_resp.json.return_value = payload + err.response = mock_resp + + delay = model._extract_gemini_retry_delay(err) + assert delay is None + + +def test_extract_gemini_retry_delay_missing_details(): + model = Model("gemini/gemini-2.5-flash") + + payload = { + "error": { + "code": 429, + "message": "Quota exceeded", + } + } + + err = Exception("Quota exceeded") + err.status_code = 429 + mock_resp = MagicMock() + mock_resp.json.return_value = payload + err.response = mock_resp + + delay = model._extract_gemini_retry_delay(err) + assert delay is None + + +def test_extract_gemini_retry_delay_headers_fallback(): + model = Model("gemini/gemini-2.5-flash") + + # Standard retry-after header (seconds) + err1 = Exception("Rate limit") + err1.status_code = 429 + mock_resp1 = MagicMock() + mock_resp1.json.return_value = {} + mock_resp1.headers = {"retry-after": "6.5"} + err1.response = mock_resp1 + assert model._extract_gemini_retry_delay(err1) == 6.5 + + # Standard retry-after-ms header (milliseconds) + err2 = Exception("Rate limit") + err2.status_code = 429 + mock_resp2 = MagicMock() + mock_resp2.json.return_value = {} + mock_resp2.headers = {"retry-after-ms": "2500"} + err2.response = mock_resp2 + assert model._extract_gemini_retry_delay(err2) == 2.5 + + +def test_extract_retry_delay_direct_err_headers(): + model = Model("openai/gpt-4o") + + # Direct headers dict on err + err = Exception("Rate limit") + err.status_code = 429 + err.headers = {"retry-after": "3.5"} + assert model._extract_gemini_retry_delay(err) == 3.5 + + err_ms = Exception("Rate limit") + err_ms.status_code = 429 + err_ms.headers = {"retry-after-ms": "4500"} + assert model._extract_gemini_retry_delay(err_ms) == 4.5 + + +def test_extract_retry_delay_malformed_headers(): + model = Model("openai/gpt-4o") + + # HTTP-date header value (non-numeric string) + err = Exception("Rate limit") + err.status_code = 429 + mock_resp = MagicMock() + mock_resp.json.return_value = {} + mock_resp.headers = {"retry-after": "Wed, 21 Oct 2026 07:28:00 GMT"} + err.response = mock_resp + assert model._extract_gemini_retry_delay(err) is None + + # Garbage string header value + mock_resp.headers = {"retry-after": "invalid"} + assert model._extract_gemini_retry_delay(err) is None + + +def test_extract_gemini_retry_delay_bytes_payload(): + model = Model("gemini/gemini-2.5-flash") + + payload_bytes = json.dumps( + { + "error": { + "code": 429, + "message": "Resource exhausted", + "details": [{"retryDelay": "4.0s"}], + } + } + ).encode("utf-8") + + err = Exception("Rate limit") + err.status_code = 429 + mock_resp = MagicMock() + mock_resp.json.side_effect = Exception("Not parsed") + mock_resp.text = payload_bytes + err.response = mock_resp + + assert model._extract_gemini_retry_delay(err) == 4.0 + + +def test_retry_fallback_to_unilateral_backoff_when_no_retry_delay(): + async def run_test(): + model = Model("gemini/gemini-2.5-flash") + model.caches_by_default = False + + rate_limit_err = litellm.RateLimitError( + message="Rate limit exceeded", + response=MagicMock( + json=lambda: {"error": {"code": 429, "message": "Rate limit exceeded"}} + ), + model="gemini/gemini-2.5-flash", + llm_provider="gemini", + ) + + call_count = 0 + slept_delays = [] + + async def mock_acompletion(*args, **kwargs): + nonlocal call_count + call_count += 1 + if call_count == 1: + raise rate_limit_err + return MagicMock(choices=[MagicMock(message=MagicMock(content="ok"))]) + + async def mock_sleep(delay): + slept_delays.append(delay) + + with ( + patch("cecli.llm.litellm.acompletion", side_effect=mock_acompletion), + patch("asyncio.sleep", side_effect=mock_sleep), + ): + await model.send_completion( + messages=[{"role": "user", "content": "hi"}], functions=None, stream=False + ) + + # Initial retry_delay is 0.125, multiplied by retry_backoff_factor (1.5) = 0.1875 + assert len(slept_delays) == 1 + assert pytest.approx(slept_delays[0]) == 0.125 * 1.5 + + asyncio.run(run_test()) + + +def test_send_completion_retry_after_header_integration(): + async def run_test(): + model = Model("openai/gpt-4o") + model.caches_by_default = False + + mock_resp = MagicMock() + mock_resp.json.return_value = {} + mock_resp.headers = {"retry-after": "2.5"} + + rate_limit_err = litellm.RateLimitError( + message="Rate limit exceeded", + response=mock_resp, + model="openai/gpt-4o", + llm_provider="openai", + ) + rate_limit_err.status_code = 429 + + call_count = 0 + slept_delays = [] + + async def mock_acompletion(*args, **kwargs): + nonlocal call_count + call_count += 1 + if call_count == 1: + raise rate_limit_err + return MagicMock(choices=[MagicMock(message=MagicMock(content="success"))]) + + async def mock_sleep(delay): + slept_delays.append(delay) + + with ( + patch("cecli.llm.litellm.acompletion", side_effect=mock_acompletion), + patch("asyncio.sleep", side_effect=mock_sleep), + ): + _hash, resp = await model.send_completion( + messages=[{"role": "user", "content": "hi"}], functions=None, stream=False + ) + + assert len(slept_delays) == 1 + assert slept_delays[0] == 2.5 + assert resp.choices[0].message.content == "success" + + asyncio.run(run_test()) + + +def test_send_completion_retry_after_ms_header_integration(): + async def run_test(): + model = Model("anthropic/claude-3-5-sonnet") + model.caches_by_default = False + + mock_resp = MagicMock() + mock_resp.json.return_value = {} + mock_resp.headers = {"retry-after-ms": "1500"} + + rate_limit_err = litellm.RateLimitError( + message="Rate limit exceeded", + response=mock_resp, + model="anthropic/claude-3-5-sonnet", + llm_provider="anthropic", + ) + rate_limit_err.status_code = 429 + + call_count = 0 + slept_delays = [] + + async def mock_acompletion(*args, **kwargs): + nonlocal call_count + call_count += 1 + if call_count == 1: + raise rate_limit_err + return MagicMock(choices=[MagicMock(message=MagicMock(content="success"))]) + + async def mock_sleep(delay): + slept_delays.append(delay) + + with ( + patch("cecli.llm.litellm.acompletion", side_effect=mock_acompletion), + patch("asyncio.sleep", side_effect=mock_sleep), + ): + _hash, resp = await model.send_completion( + messages=[{"role": "user", "content": "hi"}], functions=None, stream=False + ) + + assert len(slept_delays) == 1 + assert slept_delays[0] == 1.5 + assert resp.choices[0].message.content == "success" + + asyncio.run(run_test()) + + +def test_send_completion_gemini_payload_retry_delay_integration(): + async def run_test(): + model = Model("gemini/gemini-2.5-flash") + model.caches_by_default = False + + mock_resp = MagicMock() + mock_resp.json.return_value = { + "error": { + "code": 429, + "message": "Resource exhausted", + "details": [{"retryDelay": "3.5s"}], + } + } + + rate_limit_err = litellm.RateLimitError( + message="Resource exhausted", + response=mock_resp, + model="gemini/gemini-2.5-flash", + llm_provider="gemini", + ) + rate_limit_err.status_code = 429 + + call_count = 0 + slept_delays = [] + + async def mock_acompletion(*args, **kwargs): + nonlocal call_count + call_count += 1 + if call_count == 1: + raise rate_limit_err + return MagicMock(choices=[MagicMock(message=MagicMock(content="success"))]) + + async def mock_sleep(delay): + slept_delays.append(delay) + + with ( + patch("cecli.llm.litellm.acompletion", side_effect=mock_acompletion), + patch("asyncio.sleep", side_effect=mock_sleep), + ): + _hash, resp = await model.send_completion( + messages=[{"role": "user", "content": "hi"}], functions=None, stream=False + ) + + assert len(slept_delays) == 1 + assert slept_delays[0] == 3.5 + assert resp.choices[0].message.content == "success" + + asyncio.run(run_test()) + + +def test_send_completion_exceeds_retry_timeout(): + async def run_test(): + model = Model("openai/gpt-4o") + model.caches_by_default = False + model.retry_timeout = 10 + + mock_resp = MagicMock() + mock_resp.json.return_value = {} + mock_resp.headers = {"retry-after": "100"} + + rate_limit_err = litellm.RateLimitError( + message="Rate limit exceeded", + response=mock_resp, + model="openai/gpt-4o", + llm_provider="openai", + ) + rate_limit_err.status_code = 429 + + slept_delays = [] + + async def mock_acompletion(*args, **kwargs): + raise rate_limit_err + + async def mock_sleep(delay): + slept_delays.append(delay) + + with ( + patch("cecli.llm.litellm.acompletion", side_effect=mock_acompletion), + patch("asyncio.sleep", side_effect=mock_sleep), + ): + _hash, resp = await model.send_completion( + messages=[{"role": "user", "content": "hi"}], functions=None, stream=False + ) + + # Should not sleep and should return model error response immediately + assert len(slept_delays) == 0 + assert "Model API Response Error" in resp.choices[0].message.content + + asyncio.run(run_test()) + + +def test_simple_send_with_retries_retry_after_integration(): + async def run_test(): + model = Model("openai/gpt-4o") + + mock_resp = MagicMock() + mock_resp.json.return_value = {} + mock_resp.headers = {"retry-after": "1.5"} + + rate_limit_err = litellm.RateLimitError( + message="Rate limit exceeded", + response=mock_resp, + model="openai/gpt-4o", + llm_provider="openai", + ) + rate_limit_err.status_code = 429 + + call_count = 0 + slept_delays = [] + + async def mock_send_completion(*args, **kwargs): + nonlocal call_count + call_count += 1 + if call_count == 1: + raise rate_limit_err + return ( + "hash", + MagicMock(choices=[MagicMock(message=MagicMock(content="generated commit"))]), + ) + + def mock_sleep(delay): + slept_delays.append(delay) + + with ( + patch.object(model, "send_completion", side_effect=mock_send_completion), + patch("time.sleep", side_effect=mock_sleep), + ): + result = await model.simple_send_with_retries( + messages=[{"role": "user", "content": "generate commit"}] + ) + + assert len(slept_delays) == 1 + assert slept_delays[0] == 1.5 + assert result == "generated commit" + + asyncio.run(run_test()) + + +def test_simple_send_with_retries_gemini_payload_integration(): + async def run_test(): + model = Model("gemini/gemini-2.5-flash") + + mock_resp = MagicMock() + mock_resp.json.return_value = { + "error": { + "code": 429, + "message": "Resource exhausted", + "details": [{"retryDelay": "2.0s"}], + } + } + + rate_limit_err = litellm.RateLimitError( + message="Resource exhausted", + response=mock_resp, + model="gemini/gemini-2.5-flash", + llm_provider="gemini", + ) + rate_limit_err.status_code = 429 + + call_count = 0 + slept_delays = [] + + async def mock_send_completion(*args, **kwargs): + nonlocal call_count + call_count += 1 + if call_count == 1: + raise rate_limit_err + return ( + "hash", + MagicMock(choices=[MagicMock(message=MagicMock(content="summary output"))]), + ) + + def mock_sleep(delay): + slept_delays.append(delay) + + with ( + patch.object(model, "send_completion", side_effect=mock_send_completion), + patch("time.sleep", side_effect=mock_sleep), + ): + result = await model.simple_send_with_retries( + messages=[{"role": "user", "content": "summarize"}] + ) + + assert len(slept_delays) == 1 + assert slept_delays[0] == 2.0 + assert result == "summary output" + + asyncio.run(run_test()) + + +def test_simple_send_with_retries_exceeds_retry_timeout(): + async def run_test(): + model = Model("openai/gpt-4o") + + mock_resp = MagicMock() + mock_resp.json.return_value = {} + mock_resp.headers = {"retry-after": "100"} + + rate_limit_err = litellm.RateLimitError( + message="Rate limit exceeded", + response=mock_resp, + model="openai/gpt-4o", + llm_provider="openai", + ) + rate_limit_err.status_code = 429 + + slept_delays = [] + + async def mock_send_completion(*args, **kwargs): + raise rate_limit_err + + def mock_sleep(delay): + slept_delays.append(delay) + + with ( + patch.object(model, "send_completion", side_effect=mock_send_completion), + patch("time.sleep", side_effect=mock_sleep), + ): + result = await model.simple_send_with_retries( + messages=[{"role": "user", "content": "test"}] + ) + + assert len(slept_delays) == 0 + assert result is None + + asyncio.run(run_test())