From 12691dcce35f4896e5fa8d44e9e534cddfb094a6 Mon Sep 17 00:00:00 2001 From: giulio-leone Date: Wed, 4 Mar 2026 06:24:41 +0100 Subject: [PATCH] fix: WebSearch interception fails with thinking enabled + SpendLimit constraint --- .../websearch_interception/handler.py | 79 +++- litellm/llms/custom_httpx/llm_http_handler.py | 5 +- .../test_websearch_thinking_constraint.py | 439 ++++++++++++++++++ 3 files changed, 509 insertions(+), 14 deletions(-) create mode 100644 tests/test_litellm/integrations/websearch_interception/test_websearch_thinking_constraint.py diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index bef8925e8e..35275b574d 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -7,6 +7,7 @@ server-side using litellm router's search tools. """ import asyncio +import math from typing import Any, Dict, List, Optional, Tuple, Union, cast import litellm @@ -481,6 +482,56 @@ class WebSearchInterceptionLogger(CustomLogger): response_format=response_format, ) + @staticmethod + def _resolve_max_tokens( + optional_params: Dict, + kwargs: Dict, + ) -> int: + """Extract max_tokens and validate against thinking.budget_tokens. + + Anthropic API requires ``max_tokens > thinking.budget_tokens``. + If the constraint is violated, auto-adjust to ``budget_tokens + 1024``. + """ + max_tokens: int = optional_params.get( + "max_tokens", + kwargs.get("max_tokens", 1024), + ) + thinking_param = optional_params.get("thinking") + if thinking_param and isinstance(thinking_param, dict): + budget_tokens = thinking_param.get("budget_tokens") + if ( + budget_tokens is not None + and isinstance(budget_tokens, (int, float)) + and math.isfinite(budget_tokens) + and budget_tokens > 0 + ): + if max_tokens <= budget_tokens: + adjusted = math.ceil(budget_tokens) + 1024 + verbose_logger.warning( + "WebSearchInterception: max_tokens=%s <= thinking.budget_tokens=%s, " + "adjusting to %s to satisfy Anthropic API constraint", + max_tokens, budget_tokens, adjusted, + ) + max_tokens = adjusted + return max_tokens + + @staticmethod + def _prepare_followup_kwargs(kwargs: Dict) -> Dict: + """Build kwargs for the follow-up call, excluding internal keys. + + ``litellm_logging_obj`` MUST be excluded so the follow-up call creates + its own ``Logging`` instance via ``function_setup``. Reusing the + initial call's logging object triggers the dedup flag + (``has_logged_async_success``) which silently prevents the initial + call's spend from being recorded — the root cause of the + SpendLog / AWS billing mismatch. + """ + _internal_keys = {'litellm_logging_obj'} + return { + k: v for k, v in kwargs.items() + if not k.startswith('_websearch_interception') and k not in _internal_keys + } + async def _execute_agentic_loop( self, model: str, @@ -557,13 +608,18 @@ class WebSearchInterceptionLogger(CustomLogger): f"WebSearchInterception: Last message (tool_result): {user_message}" ) + # Correlation context for structured logging + _call_id = ( + getattr(logging_obj, "litellm_call_id", None) + or kwargs.get("litellm_call_id", "unknown") + ) + + full_model_name = model # safe default before try block + # Use anthropic_messages.acreate for follow-up request try: - # Extract max_tokens from optional params or kwargs - # max_tokens is a required parameter for anthropic_messages.acreate() - max_tokens = anthropic_messages_optional_request_params.get( - "max_tokens", - kwargs.get("max_tokens", 1024) # Default to 1024 if not found + max_tokens = self._resolve_max_tokens( + anthropic_messages_optional_request_params, kwargs ) verbose_logger.debug( @@ -576,16 +632,10 @@ class WebSearchInterceptionLogger(CustomLogger): if k != 'max_tokens' } - # Remove internal websearch interception flags from kwargs before follow-up request - # These flags are used internally and should not be passed to the LLM provider - kwargs_for_followup = { - k: v for k, v in kwargs.items() - if not k.startswith('_websearch_interception') - } + kwargs_for_followup = self._prepare_followup_kwargs(kwargs) # Get model from logging_obj.model_call_details["agentic_loop_params"] # This preserves the full model name with provider prefix (e.g., "bedrock/invoke/...") - full_model_name = model if logging_obj is not None: agentic_params = logging_obj.model_call_details.get("agentic_loop_params", {}) full_model_name = agentic_params.get("model", model) @@ -609,7 +659,10 @@ class WebSearchInterceptionLogger(CustomLogger): return final_response except Exception as e: verbose_logger.exception( - f"WebSearchInterception: Follow-up request failed: {str(e)}" + "WebSearchInterception: Follow-up request failed " + "[call_id=%s model=%s messages=%d searches=%d]: %s", + _call_id, full_model_name, len(follow_up_messages), + len(final_search_results), str(e), ) raise diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index b6fcf853ab..1cef3e9ce1 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -4454,8 +4454,11 @@ class BaseLLMHTTPHandler: return agentic_response except Exception as e: + _call_id = getattr(logging_obj, "litellm_call_id", "unknown") verbose_logger.exception( - f"LiteLLM.AgenticHookError: Exception in agentic completion hooks: {str(e)}" + "LiteLLM.AgenticHookError: Exception in agentic completion hooks " + "[call_id=%s model=%s]: %s", + _call_id, model, str(e), ) # Check if we need to convert response to fake stream diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_thinking_constraint.py b/tests/test_litellm/integrations/websearch_interception/test_websearch_thinking_constraint.py new file mode 100644 index 0000000000..476f38f5a2 --- /dev/null +++ b/tests/test_litellm/integrations/websearch_interception/test_websearch_thinking_constraint.py @@ -0,0 +1,439 @@ +""" +Tests for max_tokens vs thinking.budget_tokens constraint validation +in the websearch interception agentic loop. + +Covers: + - M1-I1: max_tokens auto-adjustment when <= thinking.budget_tokens + - M1-I3: Unit tests for thinking parameter validation + - M2-I5/I8: litellm_logging_obj excluded from follow-up kwargs to prevent SpendLog dedup + - M3-I12: Regression tests for error scenarios +""" + +from typing import Any, Dict, List +from unittest.mock import MagicMock, patch + +import pytest + +from litellm.integrations.websearch_interception.handler import ( + WebSearchInterceptionLogger, +) + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + +def _make_tool_calls() -> List[Dict]: + return [ + { + "id": "toolu_01", + "type": "tool_use", + "name": "web_search", + "input": {"query": "litellm spend tracking"}, + } + ] + + +def _make_logging_obj(model: str = "bedrock/us.anthropic.claude-opus-4-6-v1") -> MagicMock: + obj = MagicMock() + obj.model_call_details = { + "agentic_loop_params": {"model": model, "custom_llm_provider": "bedrock"}, + } + return obj + + +# --------------------------------------------------------------------------- +# M1-I1 / M1-I3: max_tokens validation against thinking.budget_tokens +# --------------------------------------------------------------------------- + +class TestThinkingBudgetTokensConstraint: + """Validate that _execute_agentic_loop adjusts max_tokens when <= thinking.budget_tokens.""" + + @pytest.mark.asyncio + async def test_max_tokens_adjusted_when_less_than_budget(self): + """max_tokens < thinking.budget_tokens → auto-adjusted to budget_tokens + 1024.""" + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"]) + captured_kwargs: Dict[str, Any] = {} + + async def _fake_acreate(**kw): + captured_kwargs.update(kw) + return MagicMock() # dummy response + + with patch( + "litellm.integrations.websearch_interception.handler.anthropic_messages.acreate", + side_effect=_fake_acreate, + ), patch.object(logger, "_execute_search", return_value="search result"): + + await logger._execute_agentic_loop( + model="us.anthropic.claude-opus-4-6-v1", + messages=[{"role": "user", "content": "hi"}], + tool_calls=_make_tool_calls(), + thinking_blocks=[], + anthropic_messages_optional_request_params={ + "max_tokens": 1024, + "thinking": {"type": "enabled", "budget_tokens": 5000}, + }, + logging_obj=_make_logging_obj(), + stream=False, + kwargs={}, + ) + + assert captured_kwargs["max_tokens"] == 5000 + 1024 + + @pytest.mark.asyncio + async def test_max_tokens_adjusted_when_equal_to_budget(self): + """max_tokens == thinking.budget_tokens → still adjusted (must be strictly greater).""" + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"]) + captured_kwargs: Dict[str, Any] = {} + + async def _fake_acreate(**kw): + captured_kwargs.update(kw) + return MagicMock() + + with patch( + "litellm.integrations.websearch_interception.handler.anthropic_messages.acreate", + side_effect=_fake_acreate, + ), patch.object(logger, "_execute_search", return_value="search result"): + + await logger._execute_agentic_loop( + model="us.anthropic.claude-opus-4-6-v1", + messages=[{"role": "user", "content": "hi"}], + tool_calls=_make_tool_calls(), + thinking_blocks=[], + anthropic_messages_optional_request_params={ + "max_tokens": 5000, + "thinking": {"type": "enabled", "budget_tokens": 5000}, + }, + logging_obj=_make_logging_obj(), + stream=False, + kwargs={}, + ) + + assert captured_kwargs["max_tokens"] == 5000 + 1024 + + @pytest.mark.asyncio + async def test_max_tokens_unchanged_when_greater_than_budget(self): + """max_tokens > thinking.budget_tokens → no adjustment needed.""" + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"]) + captured_kwargs: Dict[str, Any] = {} + + async def _fake_acreate(**kw): + captured_kwargs.update(kw) + return MagicMock() + + with patch( + "litellm.integrations.websearch_interception.handler.anthropic_messages.acreate", + side_effect=_fake_acreate, + ), patch.object(logger, "_execute_search", return_value="search result"): + + await logger._execute_agentic_loop( + model="us.anthropic.claude-opus-4-6-v1", + messages=[{"role": "user", "content": "hi"}], + tool_calls=_make_tool_calls(), + thinking_blocks=[], + anthropic_messages_optional_request_params={ + "max_tokens": 10000, + "thinking": {"type": "enabled", "budget_tokens": 5000}, + }, + logging_obj=_make_logging_obj(), + stream=False, + kwargs={}, + ) + + assert captured_kwargs["max_tokens"] == 10000 + + @pytest.mark.asyncio + async def test_no_thinking_param_no_adjustment(self): + """No thinking parameter → max_tokens used as-is (default 1024).""" + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"]) + captured_kwargs: Dict[str, Any] = {} + + async def _fake_acreate(**kw): + captured_kwargs.update(kw) + return MagicMock() + + with patch( + "litellm.integrations.websearch_interception.handler.anthropic_messages.acreate", + side_effect=_fake_acreate, + ), patch.object(logger, "_execute_search", return_value="search result"): + + await logger._execute_agentic_loop( + model="us.anthropic.claude-opus-4-6-v1", + messages=[{"role": "user", "content": "hi"}], + tool_calls=_make_tool_calls(), + thinking_blocks=[], + anthropic_messages_optional_request_params={}, + logging_obj=_make_logging_obj(), + stream=False, + kwargs={}, + ) + + assert captured_kwargs["max_tokens"] == 1024 + + @pytest.mark.asyncio + async def test_thinking_without_budget_tokens_no_adjustment(self): + """thinking param exists but has no budget_tokens → max_tokens used as-is.""" + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"]) + captured_kwargs: Dict[str, Any] = {} + + async def _fake_acreate(**kw): + captured_kwargs.update(kw) + return MagicMock() + + with patch( + "litellm.integrations.websearch_interception.handler.anthropic_messages.acreate", + side_effect=_fake_acreate, + ), patch.object(logger, "_execute_search", return_value="search result"): + + await logger._execute_agentic_loop( + model="us.anthropic.claude-opus-4-6-v1", + messages=[{"role": "user", "content": "hi"}], + tool_calls=_make_tool_calls(), + thinking_blocks=[], + anthropic_messages_optional_request_params={ + "max_tokens": 2048, + "thinking": {"type": "enabled"}, + }, + logging_obj=_make_logging_obj(), + stream=False, + kwargs={}, + ) + + assert captured_kwargs["max_tokens"] == 2048 + + +class TestResolveMaxTokensEdgeCases: + """Edge cases for _resolve_max_tokens: infinity, negative, extreme values.""" + + def test_infinity_budget_tokens_no_crash(self): + """float('inf') budget_tokens must not crash with OverflowError.""" + result = WebSearchInterceptionLogger._resolve_max_tokens( + {"max_tokens": 1024, "thinking": {"budget_tokens": float("inf")}}, {} + ) + assert result == 1024 # no adjustment for non-finite values + + def test_negative_infinity_no_crash(self): + result = WebSearchInterceptionLogger._resolve_max_tokens( + {"max_tokens": 1024, "thinking": {"budget_tokens": float("-inf")}}, {} + ) + assert result == 1024 + + def test_nan_budget_tokens_no_crash(self): + result = WebSearchInterceptionLogger._resolve_max_tokens( + {"max_tokens": 1024, "thinking": {"budget_tokens": float("nan")}}, {} + ) + assert result == 1024 + + def test_negative_budget_tokens_no_adjustment(self): + result = WebSearchInterceptionLogger._resolve_max_tokens( + {"max_tokens": 1024, "thinking": {"budget_tokens": -100}}, {} + ) + assert result == 1024 + + def test_zero_budget_tokens_no_adjustment(self): + result = WebSearchInterceptionLogger._resolve_max_tokens( + {"max_tokens": 1024, "thinking": {"budget_tokens": 0}}, {} + ) + assert result == 1024 + + +# --------------------------------------------------------------------------- +# M2-I5 / M2-I8: litellm_logging_obj excluded from follow-up kwargs +# --------------------------------------------------------------------------- + +class TestLoggingObjExcludedFromFollowUp: + """Verify litellm_logging_obj is NOT forwarded to the follow-up acreate() call. + + Passing the same logging object to both initial and follow-up calls causes + the has_logged_async_success dedup flag to fire, silently preventing the + initial call's spend from being recorded in SpendLogs. + """ + + @pytest.mark.asyncio + async def test_litellm_logging_obj_excluded_from_anthropic_followup(self): + """The Anthropic messages follow-up must NOT receive litellm_logging_obj.""" + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"]) + captured_kwargs: Dict[str, Any] = {} + + async def _fake_acreate(**kw): + captured_kwargs.update(kw) + return MagicMock() + + fake_logging_obj = _make_logging_obj() + + with patch( + "litellm.integrations.websearch_interception.handler.anthropic_messages.acreate", + side_effect=_fake_acreate, + ), patch.object(logger, "_execute_search", return_value="search result"): + + await logger._execute_agentic_loop( + model="us.anthropic.claude-opus-4-6-v1", + messages=[{"role": "user", "content": "hi"}], + tool_calls=_make_tool_calls(), + thinking_blocks=[], + anthropic_messages_optional_request_params={"max_tokens": 4096}, + logging_obj=fake_logging_obj, + stream=False, + kwargs={ + "litellm_logging_obj": fake_logging_obj, + "metadata": {"user_api_key": "test-key-hash"}, + "temperature": 0.5, + }, + ) + + # litellm_logging_obj must be absent from the follow-up call + assert "litellm_logging_obj" not in captured_kwargs + # But other kwargs (metadata, temperature) must be preserved + assert captured_kwargs.get("metadata") == {"user_api_key": "test-key-hash"} + assert captured_kwargs.get("temperature") == 0.5 + + @pytest.mark.asyncio + async def test_websearch_flags_also_excluded(self): + """Both _websearch_interception flags and litellm_logging_obj must be excluded.""" + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"]) + captured_kwargs: Dict[str, Any] = {} + + async def _fake_acreate(**kw): + captured_kwargs.update(kw) + return MagicMock() + + with patch( + "litellm.integrations.websearch_interception.handler.anthropic_messages.acreate", + side_effect=_fake_acreate, + ), patch.object(logger, "_execute_search", return_value="search result"): + + await logger._execute_agentic_loop( + model="us.anthropic.claude-opus-4-6-v1", + messages=[{"role": "user", "content": "hi"}], + tool_calls=_make_tool_calls(), + thinking_blocks=[], + anthropic_messages_optional_request_params={"max_tokens": 4096}, + logging_obj=_make_logging_obj(), + stream=False, + kwargs={ + "litellm_logging_obj": MagicMock(), + "_websearch_interception_converted_stream": True, + "_websearch_interception_other": "x", + "api_key": "fake", + }, + ) + + assert "litellm_logging_obj" not in captured_kwargs + assert "_websearch_interception_converted_stream" not in captured_kwargs + assert "_websearch_interception_other" not in captured_kwargs + assert captured_kwargs.get("api_key") == "fake" + + +# --------------------------------------------------------------------------- +# M3-I12: Regression tests for error scenarios +# --------------------------------------------------------------------------- + +class TestFollowUpErrorScenarios: + """Regression tests: the agentic loop must surface errors properly and + not silently swallow them (except at the _call_agentic_completion_hooks + level which intentionally catches to return the initial response).""" + + @pytest.mark.asyncio + async def test_followup_400_raises(self): + """A 400 error from the follow-up call must propagate out of _execute_agentic_loop.""" + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"]) + + async def _fail_acreate(**kw): + raise Exception("max_tokens must be greater than thinking.budget_tokens") + + with patch( + "litellm.integrations.websearch_interception.handler.anthropic_messages.acreate", + side_effect=_fail_acreate, + ), patch.object(logger, "_execute_search", return_value="search result"): + + with pytest.raises(Exception, match="max_tokens must be greater"): + await logger._execute_agentic_loop( + model="us.anthropic.claude-opus-4-6-v1", + messages=[{"role": "user", "content": "hi"}], + tool_calls=_make_tool_calls(), + thinking_blocks=[], + anthropic_messages_optional_request_params={"max_tokens": 4096}, + logging_obj=_make_logging_obj(), + stream=False, + kwargs={}, + ) + + @pytest.mark.asyncio + async def test_search_failure_does_not_crash_loop(self): + """If a search fails, the loop should still attempt the follow-up with error text.""" + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"]) + captured_kwargs: Dict[str, Any] = {} + + async def _fake_acreate(**kw): + captured_kwargs.update(kw) + return MagicMock() + + with patch( + "litellm.integrations.websearch_interception.handler.anthropic_messages.acreate", + side_effect=_fake_acreate, + ), patch.object( + logger, "_execute_search", side_effect=Exception("search API down") + ): + + result = await logger._execute_agentic_loop( + model="us.anthropic.claude-opus-4-6-v1", + messages=[{"role": "user", "content": "hi"}], + tool_calls=_make_tool_calls(), + thinking_blocks=[], + anthropic_messages_optional_request_params={"max_tokens": 4096}, + logging_obj=_make_logging_obj(), + stream=False, + kwargs={}, + ) + + # The follow-up call should have been made (with error text in search results) + assert result is not None + # Messages should contain the error text + follow_up_messages = captured_kwargs.get("messages", []) + assert len(follow_up_messages) > 1 # original + assistant + tool_result + + @pytest.mark.asyncio + async def test_metadata_preserved_after_logging_obj_exclusion(self): + """Proxy metadata (user_api_key, team_id, etc.) must survive in follow-up kwargs + even after litellm_logging_obj is excluded — so the new logging_obj from + function_setup has access to proxy tracking metadata.""" + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"]) + captured_kwargs: Dict[str, Any] = {} + + async def _fake_acreate(**kw): + captured_kwargs.update(kw) + return MagicMock() + + proxy_metadata = { + "user_api_key": "test-proxy-key-hash", + "user_api_key_user_id": "user-123", + "user_api_key_team_id": "team-456", + "user_api_key_org_id": "org-789", + "user_api_key_end_user_id": "end-user-001", + } + + with patch( + "litellm.integrations.websearch_interception.handler.anthropic_messages.acreate", + side_effect=_fake_acreate, + ), patch.object(logger, "_execute_search", return_value="search result"): + + await logger._execute_agentic_loop( + model="us.anthropic.claude-opus-4-6-v1", + messages=[{"role": "user", "content": "hi"}], + tool_calls=_make_tool_calls(), + thinking_blocks=[], + anthropic_messages_optional_request_params={"max_tokens": 4096}, + logging_obj=_make_logging_obj(), + stream=False, + kwargs={ + "litellm_logging_obj": MagicMock(), + "metadata": proxy_metadata, + "litellm_call_id": "call-abc-123", + }, + ) + + # litellm_logging_obj excluded + assert "litellm_logging_obj" not in captured_kwargs + # But ALL proxy metadata must be preserved + assert captured_kwargs.get("metadata") == proxy_metadata + assert captured_kwargs.get("litellm_call_id") == "call-abc-123"