From f123e558cac221a5bbb13280cc4eaa83924abe1e Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 18 Apr 2026 13:15:55 -0700 Subject: [PATCH] fix: address greptile review comments on PR #25729 MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Skip ``kwargs["tools"] = []`` injection when compression is a no-op — Anthropic Messages rejects empty tool arrays on requests that did not originally declare tools. - Move agentic-loop safety guards (fingerprint cycle / max depth) out of the per-callback try/except so they propagate instead of being swallowed by the generic exception handler. Extracted _check_agentic_loop_safety. - Gate generic ``x--session-id`` capture behind the LITELLM_CAPTURE_VENDOR_SESSION_HEADERS env var (off by default) to preserve backwards compatibility; explicit x-litellm-* headers are unaffected. - Fix monkeypatch target in pre-call-hook test to patch the actual module-level binding (litellm.integrations.compression_interception.handler.compress). - Add regression tests for empty-tools skip and opt-in session capture. Co-Authored-By: Claude Opus 4.6 --- .../compression_interception/handler.py | 21 +- litellm/llms/custom_httpx/llm_http_handler.py | 372 ++++++++++-------- litellm/proxy/litellm_pre_call_utils.py | 34 +- .../test_compression_interception_handler.py | 58 ++- .../proxy/test_litellm_pre_call_utils.py | 225 ++++++++--- 5 files changed, 468 insertions(+), 242 deletions(-) diff --git a/litellm/integrations/compression_interception/handler.py b/litellm/integrations/compression_interception/handler.py index fa58e9c7595..925ae1e1042 100644 --- a/litellm/integrations/compression_interception/handler.py +++ b/litellm/integrations/compression_interception/handler.py @@ -106,15 +106,24 @@ class CompressionInterceptionLogger(CustomLogger): embedding_model_params=self.embedding_model_params, ) - kwargs["messages"] = compressed["messages"] - kwargs["tools"] = self._merge_tools( - existing_tools=cast(Optional[List[Dict[str, Any]]], kwargs.get("tools")), - compressed_tools=cast(List[Dict[str, Any]], compressed.get("tools", [])), - ) - cache = cast(Dict[str, str], compressed.get("cache", {})) skip_reason = cast(Optional[str], compressed.get("compression_skipped_reason")) + compressed_tools = cast(List[Dict[str, Any]], compressed.get("tools", [])) + + # Only mutate kwargs when compression actually produced a result. + # If compression was a no-op (below trigger, invalid tool sequence, etc.), + # leave ``messages`` and ``tools`` untouched — injecting an empty + # ``tools: []`` onto a request that originally had no tools breaks + # Anthropic Messages requests. if cache: + kwargs["messages"] = compressed["messages"] + if compressed_tools: + kwargs["tools"] = self._merge_tools( + existing_tools=cast( + Optional[List[Dict[str, Any]]], kwargs.get("tools") + ), + compressed_tools=compressed_tools, + ) call_id = cast(Optional[str], kwargs.get("litellm_call_id")) if not call_id: call_id = str(uuid.uuid4()) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 259fe087885..03b121bd3fa 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -4476,6 +4476,34 @@ class BaseLLMHTTPHandler: fingerprints = list(kwargs.get("_agentic_loop_fingerprints", []) or []) return depth, max(max_loops, 1), fingerprints + @staticmethod + def _check_agentic_loop_safety( + tool_calls: Any, + fingerprints: List[str], + depth: int, + max_loops: int, + model: str, + ) -> str: + """ + Evaluate agentic-loop safety guards (fingerprint cycle / max depth). + + Raises ValueError on abort. Returns the current fingerprint on success. + + These checks must not be swallowed by the per-callback ``except Exception`` + block that wraps callback dispatch — they are bounded-loop / cycle-break + safety rails and must abort the agentic dispatch when they trip. + """ + fingerprint = BaseLLMHTTPHandler._fingerprint_agentic_tools(tool_calls) + if fingerprint in fingerprints: + raise ValueError( + "Agentic loop detected repeated tool-call fingerprint; aborting rerun" + ) + if depth >= max_loops: + raise ValueError( + f"Exceeded max_agentic_loops={max_loops} for model={model}" + ) + return fingerprint + @staticmethod def _fingerprint_agentic_tools(tools: Dict) -> str: try: @@ -4629,95 +4657,109 @@ class BaseLLMHTTPHandler: tools = anthropic_messages_optional_request_params.get("tools", []) depth, max_loops, fingerprints = self._get_agentic_loop_settings(kwargs=kwargs) - for callback in callbacks: + if not isinstance(callback, CustomLogger): + continue + + should_run: bool = False + tool_calls: Any = None try: - if isinstance(callback, CustomLogger): - # First: Check if agentic loop should run - ( - should_run, - tool_calls, - ) = await callback.async_should_run_agentic_loop( - response=response, + # First: Check if agentic loop should run. Wrap in try/except + # to shield from buggy user callbacks — a callback crash should + # not abort the whole request. + ( + should_run, + tool_calls, + ) = await callback.async_should_run_agentic_loop( + response=response, + model=model, + messages=messages, + tools=tools, + stream=stream, + custom_llm_provider=custom_llm_provider, + kwargs=kwargs, + ) + except Exception as e: + _call_id = getattr(logging_obj, "litellm_call_id", "unknown") + verbose_logger.exception( + "LiteLLM.AgenticHookError: Exception in " + "async_should_run_agentic_loop [call_id=%s model=%s]: %s", + _call_id, + model, + str(e), + ) + continue + + if not should_run: + continue + + # Safety guards must run OUTSIDE the callback try/except — they are + # bounded-loop / cycle-break rails that must propagate to the caller. + fingerprint = self._check_agentic_loop_safety( + tool_calls=tool_calls, + fingerprints=fingerprints, + depth=depth, + max_loops=max_loops, + model=model, + ) + + try: + kwargs_with_provider = kwargs.copy() if kwargs else {} + kwargs_with_provider["custom_llm_provider"] = custom_llm_provider + build_plan_overridden = ( + callback.__class__.async_build_agentic_loop_plan + is not CustomLogger.async_build_agentic_loop_plan + ) + if not build_plan_overridden: + return await callback.async_run_agentic_loop( + tools=tool_calls, model=model, messages=messages, - tools=tools, + response=response, + anthropic_messages_provider_config=anthropic_messages_provider_config, + anthropic_messages_optional_request_params=anthropic_messages_optional_request_params, + logging_obj=logging_obj, stream=stream, - custom_llm_provider=custom_llm_provider, - kwargs=kwargs, + kwargs=kwargs_with_provider, ) + plan = await callback.async_build_agentic_loop_plan( + tools=tool_calls, + model=model, + messages=messages, + response=response, + anthropic_messages_provider_config=anthropic_messages_provider_config, + anthropic_messages_optional_request_params=anthropic_messages_optional_request_params, + logging_obj=logging_obj, + stream=stream, + kwargs=kwargs_with_provider, + ) - if should_run: - fingerprint = self._fingerprint_agentic_tools(tool_calls) - if fingerprint in fingerprints: - raise ValueError( - "Agentic loop detected repeated tool-call fingerprint; aborting rerun" - ) - if depth >= max_loops: - raise ValueError( - f"Exceeded max_agentic_loops={max_loops} for model={model}" - ) - - kwargs_with_provider = kwargs.copy() if kwargs else {} - kwargs_with_provider["custom_llm_provider"] = ( - custom_llm_provider - ) - build_plan_overridden = ( - callback.__class__.async_build_agentic_loop_plan - is not CustomLogger.async_build_agentic_loop_plan - ) - if not build_plan_overridden: - return await callback.async_run_agentic_loop( - tools=tool_calls, - model=model, - messages=messages, - response=response, - anthropic_messages_provider_config=anthropic_messages_provider_config, - anthropic_messages_optional_request_params=anthropic_messages_optional_request_params, - logging_obj=logging_obj, - stream=stream, - kwargs=kwargs_with_provider, - ) - - plan = await callback.async_build_agentic_loop_plan( - tools=tool_calls, - model=model, - messages=messages, - response=response, - anthropic_messages_provider_config=anthropic_messages_provider_config, - anthropic_messages_optional_request_params=anthropic_messages_optional_request_params, - logging_obj=logging_obj, - stream=stream, - kwargs=kwargs_with_provider, - ) - - if plan.response_override is not None: - return plan.response_override - if plan.terminate: - verbose_logger.debug( - "Agentic loop terminated by callback=%s reason=%s", - callback.__class__.__name__, - plan.stop_reason, - ) - return response - if not plan.run_agentic_loop: - continue - - return await self._execute_anthropic_agentic_plan( - plan=plan, - model=model, - messages=messages, - anthropic_messages_optional_request_params=anthropic_messages_optional_request_params, - logging_obj=logging_obj, - kwargs=kwargs_with_provider, - depth=depth, - max_loops=max_loops, - fingerprints=fingerprints, - fingerprint=fingerprint, - stream=stream, - ) + if plan.response_override is not None: + return plan.response_override + if plan.terminate: + verbose_logger.debug( + "Agentic loop terminated by callback=%s reason=%s", + callback.__class__.__name__, + plan.stop_reason, + ) + return response + if not plan.run_agentic_loop: + continue + return await self._execute_anthropic_agentic_plan( + plan=plan, + model=model, + messages=messages, + anthropic_messages_optional_request_params=anthropic_messages_optional_request_params, + logging_obj=logging_obj, + kwargs=kwargs_with_provider, + depth=depth, + max_loops=max_loops, + fingerprints=fingerprints, + fingerprint=fingerprint, + stream=stream, + ) except Exception as e: _call_id = getattr(logging_obj, "litellm_call_id", "unknown") verbose_logger.exception( @@ -4794,97 +4836,101 @@ class BaseLLMHTTPHandler: depth, max_loops, fingerprints = self._get_agentic_loop_settings(kwargs=kwargs) for callback in callbacks: - try: - if isinstance(callback, CustomLogger): - # Check if callback has the chat completion agentic loop method - if not hasattr( - callback, "async_should_run_chat_completion_agentic_loop" - ): - continue + if not isinstance(callback, CustomLogger): + continue + if not hasattr(callback, "async_should_run_chat_completion_agentic_loop"): + continue - # First: Check if agentic loop should run - ( - should_run, - tool_calls, - ) = await callback.async_should_run_chat_completion_agentic_loop( - response=response, + should_run: bool = False + tool_calls: Any = None + try: + ( + should_run, + tool_calls, + ) = await callback.async_should_run_chat_completion_agentic_loop( + response=response, + model=model, + messages=messages, + tools=tools, + stream=stream, + custom_llm_provider=custom_llm_provider, + kwargs=kwargs, + ) + except Exception as e: + verbose_logger.exception( + "LiteLLM.AgenticHookError: Exception in " + "async_should_run_chat_completion_agentic_loop: %s", + str(e), + ) + continue + + if not should_run: + continue + + # Safety guards must run OUTSIDE the callback try/except — they are + # bounded-loop / cycle-break rails that must propagate to the caller. + fingerprint = self._check_agentic_loop_safety( + tool_calls=tool_calls, + fingerprints=fingerprints, + depth=depth, + max_loops=max_loops, + model=model, + ) + + try: + kwargs_with_provider = kwargs.copy() if kwargs else {} + kwargs_with_provider["custom_llm_provider"] = custom_llm_provider + build_plan_overridden = ( + callback.__class__.async_build_chat_completion_agentic_loop_plan + is not CustomLogger.async_build_chat_completion_agentic_loop_plan + ) + if not build_plan_overridden: + return await callback.async_run_chat_completion_agentic_loop( + tools=tool_calls, model=model, messages=messages, - tools=tools, + response=response, + optional_params=optional_params, + logging_obj=logging_obj, stream=stream, - custom_llm_provider=custom_llm_provider, - kwargs=kwargs, + kwargs=kwargs_with_provider, ) - if should_run: - fingerprint = self._fingerprint_agentic_tools(tool_calls) - if fingerprint in fingerprints: - raise ValueError( - "Agentic loop detected repeated tool-call fingerprint; aborting rerun" - ) - if depth >= max_loops: - raise ValueError( - f"Exceeded max_agentic_loops={max_loops} for model={model}" - ) + plan = await callback.async_build_chat_completion_agentic_loop_plan( + tools=tool_calls, + model=model, + messages=messages, + response=response, + optional_params=optional_params, + logging_obj=logging_obj, + stream=stream, + kwargs=kwargs_with_provider, + ) - kwargs_with_provider = kwargs.copy() if kwargs else {} - kwargs_with_provider["custom_llm_provider"] = ( - custom_llm_provider - ) - build_plan_overridden = ( - callback.__class__.async_build_chat_completion_agentic_loop_plan - is not CustomLogger.async_build_chat_completion_agentic_loop_plan - ) - if not build_plan_overridden: - return ( - await callback.async_run_chat_completion_agentic_loop( - tools=tool_calls, - model=model, - messages=messages, - response=response, - optional_params=optional_params, - logging_obj=logging_obj, - stream=stream, - kwargs=kwargs_with_provider, - ) - ) - - plan = await callback.async_build_chat_completion_agentic_loop_plan( - tools=tool_calls, - model=model, - messages=messages, - response=response, - optional_params=optional_params, - logging_obj=logging_obj, - stream=stream, - kwargs=kwargs_with_provider, - ) - - if plan.response_override is not None: - return plan.response_override - if plan.terminate: - verbose_logger.debug( - "Agentic chat loop terminated by callback=%s reason=%s", - callback.__class__.__name__, - plan.stop_reason, - ) - return response - if not plan.run_agentic_loop: - continue - - return await self._execute_chat_completion_agentic_plan( - plan=plan, - model=model, - messages=messages, - optional_params=optional_params, - kwargs=kwargs_with_provider, - custom_llm_provider=custom_llm_provider, - depth=depth, - max_loops=max_loops, - fingerprints=fingerprints, - fingerprint=fingerprint, - ) + if plan.response_override is not None: + return plan.response_override + if plan.terminate: + verbose_logger.debug( + "Agentic chat loop terminated by callback=%s reason=%s", + callback.__class__.__name__, + plan.stop_reason, + ) + return response + if not plan.run_agentic_loop: + continue + return await self._execute_chat_completion_agentic_plan( + plan=plan, + model=model, + messages=messages, + optional_params=optional_params, + kwargs=kwargs_with_provider, + custom_llm_provider=custom_llm_provider, + depth=depth, + max_loops=max_loops, + fingerprints=fingerprints, + fingerprint=fingerprint, + ) except Exception as e: verbose_logger.exception( f"LiteLLM.AgenticHookError: Exception in chat completion agentic hooks: {str(e)}" diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index c4f635b69b5..b739ee46af9 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -149,6 +149,23 @@ def _extract_generic_session_id_from_headers( return None +def _is_generic_session_header_capture_enabled() -> bool: + """ + Check whether capturing generic ``x--session-id`` headers as the + LiteLLM trace/session id is opted-in via env var. + + Defaults to False to preserve backwards compatibility — existing deployments + that send vendor session-id headers for non-LiteLLM purposes must not have + those values silently re-used as the call's ``litellm_trace_id`` / + ``litellm_session_id`` (which would regroup their spend logs and traces). + + Set ``LITELLM_CAPTURE_VENDOR_SESSION_HEADERS=true`` to enable. + """ + from litellm.secret_managers.main import get_secret_bool + + return bool(get_secret_bool("LITELLM_CAPTURE_VENDOR_SESSION_HEADERS", False)) + + def get_chain_id_from_headers(headers: Optional[Dict[str, str]]) -> Optional[str]: """ Extract chain id for call chaining from request headers. @@ -156,8 +173,10 @@ def get_chain_id_from_headers(headers: Optional[Dict[str, str]]) -> Optional[str Priority order: 1. ``x-litellm-trace-id`` (explicit, highest priority) 2. ``x-litellm-session-id`` (explicit) - 3. Any ``x--session-id`` header whose value looks like a session id - (alphanumeric / UUID, at least 8 chars). E.g. ``x-claude-code-session-id``. + 3. (OPT-IN) Any ``x--session-id`` header whose value looks like a + session id. E.g. ``x-claude-code-session-id``. Only consulted when the + ``LITELLM_CAPTURE_VENDOR_SESSION_HEADERS`` env var is truthy — keeping + the default behavior backwards-compatible. Header keys are matched case-insensitively so this works with raw header dicts from any transport. @@ -168,11 +187,14 @@ def get_chain_id_from_headers(headers: Optional[Dict[str, str]]) -> Optional[str if not headers: return None normalized = {k.lower(): v for k, v in headers.items() if isinstance(k, str)} - return ( - normalized.get("x-litellm-trace-id") - or normalized.get("x-litellm-session-id") - or _extract_generic_session_id_from_headers(normalized) + explicit = normalized.get("x-litellm-trace-id") or normalized.get( + "x-litellm-session-id" ) + if explicit: + return explicit + if _is_generic_session_header_capture_enabled(): + return _extract_generic_session_id_from_headers(normalized) + return None def safe_add_api_version_from_query_params(data: dict, request: Request): diff --git a/tests/test_litellm/integrations/compression_interception/test_compression_interception_handler.py b/tests/test_litellm/integrations/compression_interception/test_compression_interception_handler.py index d7ce9ea2ab8..56e5a94cd49 100644 --- a/tests/test_litellm/integrations/compression_interception/test_compression_interception_handler.py +++ b/tests/test_litellm/integrations/compression_interception/test_compression_interception_handler.py @@ -59,7 +59,13 @@ async def test_pre_call_hook_compresses_messages_and_injects_tool(monkeypatch): def _fake_compress(**kwargs): return compressed_result - monkeypatch.setattr("litellm.compress", _fake_compress) + # The handler does ``from litellm.compression import compress`` at module + # scope, so we must patch the binding on the handler module — patching + # ``litellm.compress`` has no effect on the already-bound reference. + monkeypatch.setattr( + "litellm.integrations.compression_interception.handler.compress", + _fake_compress, + ) kwargs = { "model": "bedrock/us.anthropic.claude-sonnet-4-5", @@ -84,6 +90,49 @@ async def test_pre_call_hook_compresses_messages_and_injects_tool(monkeypatch): assert result["litellm_call_id"] in logger._compression_cache_by_call_id +@pytest.mark.asyncio +async def test_pre_call_hook_below_trigger_does_not_inject_empty_tools(monkeypatch): + """ + When compression is a no-op (below trigger / invalid tool sequence), the + hook must NOT replace ``messages`` or inject an empty ``tools: []`` onto + a request that originally had no tools — Anthropic Messages rejects + ``tools: []``. + """ + logger = CompressionInterceptionLogger() + original_messages = [{"role": "user", "content": "short prompt"}] + + def _fake_compress_noop(**kwargs): + return { + "messages": original_messages, + "original_tokens": 42, + "compressed_tokens": 42, + "compression_ratio": 0.0, + "cache": {}, + "tools": [], + "compression_skipped_reason": "below_trigger", + } + + monkeypatch.setattr( + "litellm.integrations.compression_interception.handler.compress", + _fake_compress_noop, + ) + + kwargs = { + "model": "bedrock/us.anthropic.claude-sonnet-4-5", + "messages": original_messages, + } + + result = await logger.async_pre_call_deployment_hook( + kwargs=kwargs, call_type=CallTypes.anthropic_messages + ) + + assert result is not None + # Original request had no ``tools`` — skipped compression must leave it that way. + assert "tools" not in result + # Cache must not be populated for a no-op. + assert result.get("litellm_call_id") not in logger._compression_cache_by_call_id + + @pytest.mark.asyncio async def test_should_run_agentic_loop_detects_retrieval_tool_use(): """Test should-run hook returns tool calls for retrieval tool_use blocks.""" @@ -237,7 +286,12 @@ async def test_should_run_agentic_loop_with_custom_type_tools(): "key": { "type": "string", "description": "The identifier of the content to retrieve", - "enum": ["message_0", "HA_UPTIME_ROUTER_SPEC.md", "message_159", "message_160"], + "enum": [ + "message_0", + "HA_UPTIME_ROUTER_SPEC.md", + "message_159", + "message_160", + ], } }, "required": ["key"], diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 8763b12f37b..29490af12c9 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -221,7 +221,10 @@ async def test_add_litellm_data_to_request_user_spend_and_budget(): request_mock.client = MagicMock() request_mock.client.host = "127.0.0.1" - data = {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hello"}]} + data = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "Hello"}], + } user_api_key_dict = UserAPIKeyAuth( api_key="hashed-key", @@ -1023,6 +1026,7 @@ def test_add_headers_to_llm_call_by_model_group_existing_headers_in_data(): # Restore original model_group_settings litellm.model_group_settings = original_model_group_settings + import json import time from typing import Optional @@ -1040,15 +1044,16 @@ class TestCustomLogger(CustomLogger): def __init__(self): self.standard_logging_object: Optional[StandardLoggingPayload] = None super().__init__() - + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): print(f"SUCCESS CALLBACK CALLED! kwargs keys: {list(kwargs.keys())}") self.standard_logging_object = kwargs.get("standard_logging_object") print(f"Captured standard_logging_object: {self.standard_logging_object}") - + async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): print(f"FAILURE CALLBACK CALLED! kwargs keys: {list(kwargs.keys())}") + @pytest.mark.asyncio async def test_add_litellm_metadata_from_request_headers(): """ @@ -1065,8 +1070,16 @@ async def test_add_litellm_metadata_from_request_headers(): try: # Prepare test data (ensure no streaming, add mock_response and api_key to route to litellm.acompletion) - headers = {"x-litellm-spend-logs-metadata": '{"user_id": "12345", "project_id": "proj_abc", "request_type": "chat_completion", "timestamp": "2025-09-02T10:30:00Z"}'} - data = {"model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}], "stream": False, "mock_response": "Hi", "api_key": "fake-key"} + headers = { + "x-litellm-spend-logs-metadata": '{"user_id": "12345", "project_id": "proj_abc", "request_type": "chat_completion", "timestamp": "2025-09-02T10:30:00Z"}' + } + data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}], + "stream": False, + "mock_response": "Hi", + "api_key": "fake-key", + } # Create mock request with headers mock_request = MagicMock(spec=Request) @@ -1078,9 +1091,7 @@ async def test_add_litellm_metadata_from_request_headers(): # Create mock user API key dict mock_user_api_key_dict = UserAPIKeyAuth( - api_key="test-key", - user_id="test-user", - org_id="test-org" + api_key="test-key", user_id="test-user", org_id="test-org" ) # Create mock proxy logging object @@ -1095,7 +1106,7 @@ async def test_add_litellm_metadata_from_request_headers(): async def mock_post_call_success_hook(*args, **kwargs): # Return the response unchanged - return kwargs.get('response', args[2] if len(args) > 2 else None) + return kwargs.get("response", args[2] if len(args) > 2 else None) mock_proxy_logging_obj.during_call_hook = mock_during_call_hook mock_proxy_logging_obj.pre_call_hook = mock_pre_call_hook @@ -1108,10 +1119,15 @@ async def test_add_litellm_metadata_from_request_headers(): general_settings = {} # Create mock select_data_generator with correct signature - def mock_select_data_generator(response=None, user_api_key_dict=None, request_data=None): + def mock_select_data_generator( + response=None, user_api_key_dict=None, request_data=None + ): async def mock_generator(): - yield "data: " + json.dumps({"choices": [{"delta": {"content": "Hello"}}]}) + "\n\n" + yield "data: " + json.dumps( + {"choices": [{"delta": {"content": "Hello"}}]} + ) + "\n\n" yield "data: [DONE]\n\n" + return mock_generator() # Create the processor @@ -1129,22 +1145,28 @@ async def test_add_litellm_metadata_from_request_headers(): select_data_generator=mock_select_data_generator, llm_router=None, model="gpt-4", - is_streaming_request=False + is_streaming_request=False, ) # Sleep for 3 seconds to allow logging to complete await asyncio.sleep(3) # Check if standard_logging_object was set - assert test_logger.standard_logging_object is not None, "standard_logging_object should be populated after LLM request" + assert ( + test_logger.standard_logging_object is not None + ), "standard_logging_object should be populated after LLM request" # Verify the logging object contains expected metadata standard_logging_obj = test_logger.standard_logging_object - print(f"Standard logging object captured: {json.dumps(standard_logging_obj, indent=4, default=str)}") + print( + f"Standard logging object captured: {json.dumps(standard_logging_obj, indent=4, default=str)}" + ) SPEND_LOGS_METADATA = standard_logging_obj["metadata"]["spend_logs_metadata"] - assert SPEND_LOGS_METADATA == dict(json.loads(headers["x-litellm-spend-logs-metadata"])), "spend_logs_metadata should be the same as the headers" + assert SPEND_LOGS_METADATA == dict( + json.loads(headers["x-litellm-spend-logs-metadata"]) + ), "spend_logs_metadata should be the same as the headers" finally: litellm.callbacks = original_callbacks @@ -1191,8 +1213,15 @@ def test_add_litellm_metadata_from_request_headers_both_headers_trace_id_precede assert data["litellm_trace_id"] == "trace-value" -def test_add_litellm_metadata_from_request_headers_generic_session_id_header(): - """A generic x--session-id header is used when no explicit litellm header is set.""" +def test_add_litellm_metadata_from_request_headers_generic_session_id_header( + monkeypatch, +): + """ + A generic x--session-id header is used when no explicit litellm + header is set — only when the opt-in env var is enabled. + """ + monkeypatch.setenv("LITELLM_CAPTURE_VENDOR_SESSION_HEADERS", "true") + headers = {"x-claude-code-session-id": "e96634a3-fa28-4083-b354-55542e2dca01"} data = {"metadata": {}} LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( @@ -1203,8 +1232,31 @@ def test_add_litellm_metadata_from_request_headers_generic_session_id_header(): assert data["litellm_trace_id"] == "e96634a3-fa28-4083-b354-55542e2dca01" -def test_add_litellm_metadata_from_request_headers_explicit_header_beats_generic(): +def test_add_litellm_metadata_from_request_headers_generic_session_id_header_ignored_by_default( + monkeypatch, +): + """ + Default (flag off): a generic x--session-id header must NOT be + treated as a litellm chain id — prior users who sent such headers for + non-LiteLLM purposes continue to work unchanged. + """ + monkeypatch.delenv("LITELLM_CAPTURE_VENDOR_SESSION_HEADERS", raising=False) + + headers = {"x-claude-code-session-id": "e96634a3-fa28-4083-b354-55542e2dca01"} + data = {"metadata": {}} + LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( + headers=headers, data=data, _metadata_variable_name="metadata" + ) + assert "litellm_session_id" not in data + assert "litellm_trace_id" not in data + + +def test_add_litellm_metadata_from_request_headers_explicit_header_beats_generic( + monkeypatch, +): """Explicit x-litellm-trace-id wins over a generic x-*-session-id header.""" + monkeypatch.setenv("LITELLM_CAPTURE_VENDOR_SESSION_HEADERS", "true") + headers = { "x-litellm-trace-id": "explicit-trace-id-value", "x-claude-code-session-id": "e96634a3-fa28-4083-b354-55542e2dca01", @@ -1217,10 +1269,16 @@ def test_add_litellm_metadata_from_request_headers_explicit_header_beats_generic assert data["litellm_trace_id"] == "explicit-trace-id-value" -def test_get_chain_id_from_headers_generic_vendor_session_id(): - """get_chain_id_from_headers picks up any x--session-id with a valid value.""" +def test_get_chain_id_from_headers_generic_vendor_session_id(monkeypatch): + """ + Generic ``x--session-id`` capture is opt-in via + ``LITELLM_CAPTURE_VENDOR_SESSION_HEADERS``; when enabled, valid values are + picked up and explicit headers still take precedence. + """ from litellm.proxy.litellm_pre_call_utils import get_chain_id_from_headers + monkeypatch.setenv("LITELLM_CAPTURE_VENDOR_SESSION_HEADERS", "true") + assert ( get_chain_id_from_headers( {"x-claude-code-session-id": "e96634a3-fa28-4083-b354-55542e2dca01"} @@ -1242,13 +1300,42 @@ def test_get_chain_id_from_headers_generic_vendor_session_id(): ) +def test_get_chain_id_from_headers_generic_vendor_session_id_disabled_by_default( + monkeypatch, +): + """ + Generic vendor session-id capture must stay OFF by default — otherwise + existing deployments that send such headers for non-LiteLLM purposes would + have their spend logs / traces silently regrouped under those IDs. + """ + from litellm.proxy.litellm_pre_call_utils import get_chain_id_from_headers + + monkeypatch.delenv("LITELLM_CAPTURE_VENDOR_SESSION_HEADERS", raising=False) + + # Without the opt-in env var, generic vendor session headers are ignored. + assert ( + get_chain_id_from_headers( + {"x-claude-code-session-id": "e96634a3-fa28-4083-b354-55542e2dca01"} + ) + is None + ) + + # Explicit litellm headers still work (unchanged behavior). + assert ( + get_chain_id_from_headers({"x-litellm-trace-id": "explicit-id-value"}) + == "explicit-id-value" + ) + + def test_get_internal_user_header_from_mapping_returns_expected_header(): mappings = [ {"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"}, {"header_name": "X-OpenWebUI-User-Email", "litellm_user_role": "customer"}, ] - header_name = LiteLLMProxyRequestSetup.get_internal_user_header_from_mapping(mappings) + header_name = LiteLLMProxyRequestSetup.get_internal_user_header_from_mapping( + mappings + ) assert header_name == "X-OpenWebUI-User-Id" @@ -1256,7 +1343,9 @@ def test_get_internal_user_header_from_mapping_none_when_absent(): mappings = [ {"header_name": "X-OpenWebUI-User-Email", "litellm_user_role": "customer"} ] - header_name = LiteLLMProxyRequestSetup.get_internal_user_header_from_mapping(mappings) + header_name = LiteLLMProxyRequestSetup.get_internal_user_header_from_mapping( + mappings + ) assert header_name is None single = {"header_name": "X-Only-Customer", "litellm_user_role": "customer"} @@ -1269,7 +1358,10 @@ def test_add_internal_user_from_user_mapping_sets_user_id_when_header_present(): headers = {"X-OpenWebUI-User-Id": "internal-user-123"} general_settings = { "user_header_mappings": [ - {"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"}, + { + "header_name": "X-OpenWebUI-User-Id", + "litellm_user_role": "internal_user", + }, {"header_name": "X-OpenWebUI-User-Email", "litellm_user_role": "customer"}, ] } @@ -1363,7 +1455,7 @@ async def test_team_guardrails_append_to_key_guardrails(): metadata = updated_data.get("metadata", {}) guardrails = metadata.get("guardrails", []) - + assert "key-guardrail-1" in guardrails assert "key-guardrail-2" in guardrails assert "team-guardrail-1" in guardrails @@ -1392,7 +1484,7 @@ async def test_request_guardrails_do_not_override_key_guardrails(): metadata={"guardrails": ["key-guardrail-1"]}, team_metadata={}, ) - + # Test case: Request with empty guardrails should not result in empty guardrails data_with_empty = { "model": "gpt-3.5-turbo", @@ -1412,7 +1504,7 @@ async def test_request_guardrails_do_not_override_key_guardrails(): _metadata = updated_data_empty.get("metadata", {}) requested_guardrails = _metadata.get("guardrails", []) - + assert "guardrails" not in updated_data_empty assert "key-guardrail-1" in requested_guardrails assert len(requested_guardrails) == 1 @@ -1527,7 +1619,10 @@ def test_update_model_if_key_alias_exists(): assert data["model"] == "xai/grok-4-fast-non-reasoning" # Test case 2: Key alias doesn't exist - data = {"model": "unknown-model", "messages": [{"role": "user", "content": "Hello"}]} + data = { + "model": "unknown-model", + "messages": [{"role": "user", "content": "Hello"}], + } user_api_key_dict = UserAPIKeyAuth( api_key="test-key", aliases={"modelAlias": "xai/grok-4-fast-non-reasoning"}, @@ -1645,16 +1740,22 @@ async def test_embedding_header_forwarding_with_model_group(): # Verify that only x- prefixed headers (except x-stainless) were forwarded forwarded_headers = updated_data["headers"] - assert "X-Custom-Header" in forwarded_headers, "X-Custom-Header should be forwarded" + assert ( + "X-Custom-Header" in forwarded_headers + ), "X-Custom-Header should be forwarded" assert forwarded_headers["X-Custom-Header"] == "custom-value" assert "X-Request-ID" in forwarded_headers, "X-Request-ID should be forwarded" assert forwarded_headers["X-Request-ID"] == "test-request-123" # Verify that authorization header was NOT forwarded (sensitive header) - assert "Authorization" not in forwarded_headers, "Authorization header should not be forwarded" + assert ( + "Authorization" not in forwarded_headers + ), "Authorization header should not be forwarded" # Verify that Content-Type was NOT forwarded (doesn't start with x-) - assert "Content-Type" not in forwarded_headers, "Content-Type should not be forwarded" + assert ( + "Content-Type" not in forwarded_headers + ), "Content-Type should not be forwarded" # Verify original data fields are preserved assert updated_data["model"] == "local-openai/text-embedding-3-small" @@ -1710,8 +1811,9 @@ async def test_embedding_header_forwarding_without_model_group_config(): ) # Verify that headers were NOT added since model is not in forward list - assert "headers" not in updated_data or updated_data.get("headers") is None, \ - "Headers should not be forwarded for models not in forward_client_headers_to_llm_api list" + assert ( + "headers" not in updated_data or updated_data.get("headers") is None + ), "Headers should not be forwarded for models not in forward_client_headers_to_llm_api list" # Verify original data fields are preserved assert updated_data["model"] == "text-embedding-ada-002" @@ -1765,7 +1867,9 @@ async def test_add_guardrails_from_policy_engine(): attachment_registry = get_attachment_registry() attachment_registry._attachments = [ PolicyAttachment(policy="global-baseline", scope="*"), # applies to all - PolicyAttachment(policy="healthcare", teams=["healthcare-team"]), # applies to healthcare team + PolicyAttachment( + policy="healthcare", teams=["healthcare-team"] + ), # applies to healthcare team ] attachment_registry._initialized = True @@ -1808,7 +1912,10 @@ async def test_add_guardrails_from_policy_engine_accepts_dynamic_policies_and_po data = { "model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}], - "policies": ["PII-POLICY-GLOBAL", "HIPAA-POLICY"], # Dynamic policies - should be accepted and removed + "policies": [ + "PII-POLICY-GLOBAL", + "HIPAA-POLICY", + ], # Dynamic policies - should be accepted and removed "metadata": {}, } @@ -1831,7 +1938,9 @@ async def test_add_guardrails_from_policy_engine_accepts_dynamic_policies_and_po ) # Verify that 'policies' was removed from the request body - assert "policies" not in data, "'policies' should be removed from request body to prevent forwarding to LLM provider" + assert ( + "policies" not in data + ), "'policies' should be removed from request body to prevent forwarding to LLM provider" # Verify that other fields are preserved assert "model" in data @@ -1920,7 +2029,9 @@ async def test_bearer_token_not_in_debug_logs(): from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request from litellm.proxy.proxy_server import ProxyConfig - secret_token = "eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxMjM0NTY3ODkwIn0.fakesignature" + secret_token = ( + "eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxMjM0NTY3ODkwIn0.fakesignature" + ) mock_request = MagicMock(spec=Request) mock_request.headers = { @@ -1949,8 +2060,10 @@ async def test_bearer_token_not_in_debug_logs(): logger.setLevel(logging.DEBUG) try: - with patch("litellm.proxy.proxy_server.llm_router", None), \ - patch("litellm.proxy.proxy_server.premium_user", True): + with ( + patch("litellm.proxy.proxy_server.llm_router", None), + patch("litellm.proxy.proxy_server.premium_user", True), + ): await add_litellm_data_to_request( data=data, request=mock_request, @@ -2071,9 +2184,7 @@ def test_resolve_project_model_specific_wins(): "gpt-4": {"azure": {"litellm_credentials": "team-gpt4"}}, "defaultconfig": {"azure": {"litellm_credentials": "team-default"}}, } - result = _resolve_credential_from_model_config( - "gpt-4", project_config, team_config - ) + result = _resolve_credential_from_model_config("gpt-4", project_config, team_config) assert result == "proj-gpt4" @@ -2085,9 +2196,7 @@ def test_resolve_project_default_wins_over_team(): "gpt-4": {"azure": {"litellm_credentials": "team-gpt4"}}, "defaultconfig": {"azure": {"litellm_credentials": "team-default"}}, } - result = _resolve_credential_from_model_config( - "gpt-4", project_config, team_config - ) + result = _resolve_credential_from_model_config("gpt-4", project_config, team_config) assert result == "proj-default" @@ -2142,12 +2251,8 @@ def test_apply_overrides_project_model_specific(setup_test_credentials): }, project_metadata={ "model_config": { - "defaultconfig": { - "azure": {"litellm_credentials": "hotel-rec-azure"} - }, - "gpt-4-vision": { - "azure": {"litellm_credentials": "hotel-rec-vision"} - }, + "defaultconfig": {"azure": {"litellm_credentials": "hotel-rec-azure"}}, + "gpt-4-vision": {"azure": {"litellm_credentials": "hotel-rec-vision"}}, } }, ) @@ -2174,12 +2279,8 @@ def test_apply_overrides_project_default(setup_test_credentials): }, project_metadata={ "model_config": { - "defaultconfig": { - "azure": {"litellm_credentials": "hotel-rec-azure"} - }, - "gpt-4-vision": { - "azure": {"litellm_credentials": "hotel-rec-vision"} - }, + "defaultconfig": {"azure": {"litellm_credentials": "hotel-rec-azure"}}, + "gpt-4-vision": {"azure": {"litellm_credentials": "hotel-rec-vision"}}, } }, ) @@ -2282,9 +2383,7 @@ def test_apply_overrides_missing_credential_name(setup_test_credentials): api_key="test-key", team_metadata={ "model_config": { - "gpt-4": { - "azure": {"litellm_credentials": "nonexistent-credential"} - } + "gpt-4": {"azure": {"litellm_credentials": "nonexistent-credential"}} } }, ) @@ -2323,9 +2422,7 @@ def test_apply_overrides_no_model_in_data(setup_test_credentials): api_key="test-key", team_metadata={ "model_config": { - "defaultconfig": { - "azure": {"litellm_credentials": "some-cred"} - } + "defaultconfig": {"azure": {"litellm_credentials": "some-cred"}} } }, ) @@ -2356,9 +2453,7 @@ def test_apply_overrides_clientside_api_version_preserved(setup_test_credentials api_key="test-key", team_metadata={ "model_config": { - "gpt-4-vision": { - "azure": {"litellm_credentials": "hotel-rec-vision"} - } + "gpt-4-vision": {"azure": {"litellm_credentials": "hotel-rec-vision"}} } }, )