diff --git a/openkb/agent/compiler.py b/openkb/agent/compiler.py index d0c9f878d..2312f6902 100644 --- a/openkb/agent/compiler.py +++ b/openkb/agent/compiler.py @@ -260,6 +260,71 @@ # --------------------------------------------------------------------------- +def _should_retry_exception(exc: Exception) -> bool: + """Determine whether an exception is retryable (transient error). + + Returns True for temporary API/network errors that may succeed on retry: + - Timeout (client-side or server-side) + - APIError 5xx (server errors) + - RateLimitError (429) + - ConnectionError / ServiceUnavailableError + + Returns False for permanent errors that won't be fixed by retry: + - TruncatedResponseError (model hit max_tokens) + - ValueError, TypeError (malformed input/output) + - AuthenticationError (credentials issue) + - BadRequestError (invalid parameters) + - Unknown error types (conservative approach) + """ + exc_type_name = type(exc).__name__ + + # ===== RETRYABLE (transient errors) ===== + + # Timeout (network/gateway timeout) + if "Timeout" in exc_type_name: + return True + + # Generic API errors (5xx range, but not 4xx) + if "APIError" in exc_type_name: + # Don't retry if it's a BadRequest/Invalid error (4xx) + if "Invalid" not in exc_type_name and "BadRequest" not in exc_type_name: + return True + + # Rate limiting (429) + if "RateLimitError" in exc_type_name or "Rate" in exc_type_name: + return True + + # Connection errors + if "ConnectionError" in exc_type_name: + return True + + # Service unavailable + if "ServiceUnavailable" in exc_type_name: + return True + + # ===== NOT RETRYABLE (permanent errors) ===== + + # Model hit max_tokens limit + if isinstance(exc, TruncatedResponseError): + return False + + # Content validation failures + if "ValueError" in exc_type_name or "TypeError" in exc_type_name: + return False + + # Authentication failures + if "Auth" in exc_type_name or "Permission" in exc_type_name: + return False + + # Bad parameters/requests + if "BadRequest" in exc_type_name or "Invalid" in exc_type_name: + return False + + # ===== UNKNOWN: Conservative approach ===== + # Don't retry errors we don't recognize + return False + + def _cached_text(text: str) -> list[dict]: """Wrap a text payload into a content-block list with an Anthropic ephemeral cache_control marker. @@ -406,7 +471,11 @@ def _llm_call( bundle=None, **kwargs, ) -> str: - """Single LLM call with animated progress and debug logging.""" + """Single LLM call with animated progress, debug logging, and retry support. + + Transient errors (Timeout, 5xx, 429) are automatically retried by LiteLLM. + Permanent errors (4xx, truncation, validation) are raised immediately. + """ messages = _prepare_messages(model, messages) extra_headers = bundle.extra_headers if bundle is not None else get_extra_headers() if extra_headers: @@ -417,6 +486,15 @@ def _llm_call( if bundle is not None: kwargs.setdefault("api_key", bundle.api_key) kwargs.setdefault("base_url", bundle.base_url) + + # Retry configuration for transient errors (fixed: 2 retries). Uses + # LiteLLM's recognized ``num_retries`` kwarg — NOT ``retries``, which + # LiteLLM does not treat as an internal control parameter. An + # unrecognized kwarg falls through as a provider request-body field, + # which strict-mode proxies reject with e.g. "retries: Extra inputs + # are not permitted" (#233). + kwargs.setdefault("num_retries", 2) + logger.debug("LLM request [%s]:\n%s", step_name, _fmt_messages(messages)) if kwargs: logger.debug("LLM kwargs [%s]: %s", step_name, kwargs) @@ -425,7 +503,27 @@ def _llm_call( spinner.start() t0 = time.time() - response = litellm.completion(model=model, messages=messages, **kwargs) + try: + response = litellm.completion(model=model, messages=messages, **kwargs) + except Exception as exc: + # NEW: Better error logging with retry context + if _should_retry_exception(exc): + logger.warning( + "LLM [%s] failed with transient error (retries applied by LiteLLM): %s", + step_name, + exc, + exc_info=False, # Don't spam stack traces for known transient errors + ) + else: + logger.warning( + "LLM [%s] failed with permanent error (no retry): %s", + step_name, + exc, + exc_info=True, # Full trace for unexpected errors + ) + spinner.stop("[FAILED]") + raise + content = response.choices[0].message.content or "" truncated = _warn_if_truncated(response, step_name, kwargs.get("max_tokens")) @@ -449,7 +547,11 @@ async def _llm_call_async( bundle=None, **kwargs, ) -> str: - """Async LLM call with timing output and debug logging.""" + """Async LLM call with timing output, debug logging, and retry support. + + Transient errors (Timeout, 5xx, 429) are automatically retried by LiteLLM. + Permanent errors (4xx, truncation, validation) are raised immediately. + """ messages = _prepare_messages(model, messages) extra_headers = bundle.extra_headers if bundle is not None else get_extra_headers() if extra_headers: @@ -460,13 +562,41 @@ async def _llm_call_async( if bundle is not None: kwargs.setdefault("api_key", bundle.api_key) kwargs.setdefault("base_url", bundle.base_url) + + # Retry configuration for transient errors (fixed: 2 retries). Uses + # LiteLLM's recognized ``num_retries`` kwarg — NOT ``retries``, which + # LiteLLM does not treat as an internal control parameter. An + # unrecognized kwarg falls through as a provider request-body field, + # which strict-mode proxies reject with e.g. "retries: Extra inputs + # are not permitted" (#233). + kwargs.setdefault("num_retries", 2) + logger.debug("LLM request [%s]:\n%s", step_name, _fmt_messages(messages)) if kwargs: logger.debug("LLM kwargs [%s]: %s", step_name, kwargs) t0 = time.time() - response = await litellm.acompletion(model=model, messages=messages, **kwargs) + try: + response = await litellm.acompletion(model=model, messages=messages, **kwargs) + except Exception as exc: + # NEW: Better error logging with retry context + if _should_retry_exception(exc): + logger.warning( + "LLM [%s] failed with transient error (retries applied by LiteLLM): %s", + step_name, + exc, + exc_info=False, + ) + else: + logger.warning( + "LLM [%s] failed with permanent error (no retry): %s", + step_name, + exc, + exc_info=True, + ) + raise + content = response.choices[0].message.content or "" truncated = _warn_if_truncated(response, step_name, kwargs.get("max_tokens")) diff --git a/tests/test_compiler_retry.py b/tests/test_compiler_retry.py new file mode 100644 index 000000000..9ab7e3c69 --- /dev/null +++ b/tests/test_compiler_retry.py @@ -0,0 +1,180 @@ +"""Tests for LLM retry logic in compiler.py.""" + +import asyncio +from unittest.mock import AsyncMock, MagicMock, patch + +from openkb.agent.compiler import ( + TruncatedResponseError, + _llm_call, + _llm_call_async, + _should_retry_exception, +) + + +# Custom exception classes for testing (so we can control the type name) +class TimeoutError(Exception): + """Simulates litellm.Timeout.""" + + pass + + +class APIError(Exception): + """Simulates litellm.APIError (5xx).""" + + pass + + +class InvalidAPIError(APIError): + """Simulates InvalidAPIError (4xx).""" + + pass + + +class BadRequestError(Exception): + """Simulates BadRequestError.""" + + pass + + +class RateLimitError(Exception): + """Simulates litellm.RateLimitError.""" + + pass + + +class AuthenticationError(Exception): + """Simulates AuthenticationError.""" + + pass + + +class PermissionError(Exception): + """Simulates PermissionError.""" + + pass + + +class ServiceUnavailableError(Exception): + """Simulates ServiceUnavailableError.""" + + pass + + +class TestShouldRetryException: + """Test the exception filtering logic for retry decisions.""" + + def test_retryable_timeout(self): + """Timeout should be retryable.""" + exc = TimeoutError("Gateway Timeout") + assert _should_retry_exception(exc) is True + + def test_retryable_api_error_5xx(self): + """5xx API errors should be retryable.""" + exc = APIError("503 Service Unavailable") + assert _should_retry_exception(exc) is True + + def test_not_retryable_invalid_api_error(self): + """InvalidAPIError (4xx) should NOT be retryable.""" + exc = InvalidAPIError("400 Bad Request") + assert _should_retry_exception(exc) is False + + def test_retryable_rate_limit(self): + """Rate limit errors should be retryable.""" + exc = RateLimitError("429 Too Many Requests") + assert _should_retry_exception(exc) is True + + def test_retryable_connection_error(self): + """Connection errors should be retryable.""" + exc = ConnectionError("Connection refused") + assert _should_retry_exception(exc) is True + + def test_retryable_service_unavailable(self): + """Service unavailable errors should be retryable.""" + exc = ServiceUnavailableError("Service down") + assert _should_retry_exception(exc) is True + + def test_not_retryable_truncation(self): + """Truncated output should NOT be retryable.""" + exc = TruncatedResponseError("hit length limit") + assert _should_retry_exception(exc) is False + + def test_not_retryable_value_error(self): + """ValueError should NOT be retryable.""" + exc = ValueError("empty content") + assert _should_retry_exception(exc) is False + + def test_not_retryable_type_error(self): + """TypeError should NOT be retryable.""" + exc = TypeError("malformed") + assert _should_retry_exception(exc) is False + + def test_not_retryable_auth_error(self): + """Authentication errors should NOT be retryable.""" + exc = AuthenticationError("invalid API key") + assert _should_retry_exception(exc) is False + + def test_not_retryable_permission_error(self): + """Permission errors should NOT be retryable.""" + exc = PermissionError("forbidden") + assert _should_retry_exception(exc) is False + + def test_not_retryable_bad_request(self): + """BadRequest errors should NOT be retryable.""" + exc = BadRequestError("invalid params") + assert _should_retry_exception(exc) is False + + def test_not_retryable_unknown(self): + """Unknown errors should NOT be retried (conservative).""" + + class WeirdCustomError(Exception): + pass + + exc = WeirdCustomError("something weird") + assert _should_retry_exception(exc) is False + + def test_not_retryable_generic_exception(self): + """Generic Exception without special name should NOT be retried.""" + exc = Exception("generic error") + assert _should_retry_exception(exc) is False + + +def _fake_response(): + choice = MagicMock() + choice.message.content = "ok" + choice.finish_reason = "stop" + resp = MagicMock() + resp.choices = [choice] + return resp + + +class TestRetryKwargForwarding: + """Regression tests for #233: the retry kwarg forwarded to LiteLLM must be + ``num_retries`` (LiteLLM's recognized internal control parameter), not + ``retries``. An unrecognized kwarg falls through as a provider + request-body field, which strict-mode proxies reject. + """ + + def test_llm_call_forwards_num_retries_not_retries(self): + with patch( + "openkb.agent.compiler.litellm.completion", return_value=_fake_response() + ) as completion: + _llm_call("gpt-4o", [{"role": "user", "content": "hi"}], "step") + assert completion.call_args.kwargs["num_retries"] == 2 + assert "retries" not in completion.call_args.kwargs + + def test_llm_call_does_not_override_explicit_num_retries(self): + with patch( + "openkb.agent.compiler.litellm.completion", return_value=_fake_response() + ) as completion: + _llm_call("gpt-4o", [{"role": "user", "content": "hi"}], "step", num_retries=5) + assert completion.call_args.kwargs["num_retries"] == 5 + + def test_llm_call_async_forwards_num_retries_not_retries(self): + with patch( + "openkb.agent.compiler.litellm.acompletion", + new_callable=AsyncMock, + return_value=_fake_response(), + ) as acompletion: + asyncio.run(_llm_call_async("gpt-4o", [{"role": "user", "content": "hi"}], "step")) + assert acompletion.call_args.kwargs["num_retries"] == 2 + assert "retries" not in acompletion.call_args.kwargs