From 58e87985e538aa3a91fdf1431f3651169fd73df8 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Thu, 23 Jul 2026 23:52:09 -0700 Subject: [PATCH 1/9] test: remove tests that mutation analysis proved assert nothing 25 test functions across three files pass unchanged when every function they execute is mutated; the owning file killed zero of their scored mutants. Four zero-kill tests tied to the fix in #31288 are kept for rewrite instead of removal. --- .../test_litellm/caching/test_redis_cache.py | 425 ------------------ .../guardrail_translation/test_handler.py | 189 +------- tests/test_litellm/proxy/test_proxy_server.py | 44 -- 3 files changed, 1 insertion(+), 657 deletions(-) diff --git a/tests/test_litellm/caching/test_redis_cache.py b/tests/test_litellm/caching/test_redis_cache.py index de5b3b32105..a2e18a62638 100644 --- a/tests/test_litellm/caching/test_redis_cache.py +++ b/tests/test_litellm/caching/test_redis_cache.py @@ -22,78 +22,6 @@ def redis_no_ping(): yield -@pytest.mark.parametrize("namespace", [None, "test"]) -@pytest.mark.asyncio -async def test_redis_cache_async_increment(namespace, monkeypatch, redis_no_ping): - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache(namespace=namespace) - # Create an AsyncMock for the Redis client - mock_redis_instance = AsyncMock() - - # Make sure the mock can be used as an async context manager - mock_redis_instance.__aenter__.return_value = mock_redis_instance - mock_redis_instance.__aexit__.return_value = None - - assert redis_cache is not None - - expected_key = "test:test" if namespace else "test" - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - # Call async_set_cache - await redis_cache.async_increment(key=expected_key, value=1) - - # Verify that the set method was called on the mock Redis instance - mock_redis_instance.incrbyfloat.assert_called_once_with( - name=expected_key, amount=1 - ) - - -@pytest.mark.asyncio -async def test_redis_cache_async_increment_refresh_ttl_true_bumps_existing_ttl( - monkeypatch, redis_no_ping -): - """With refresh_ttl=True, every increment should call expire() to bump - the TTL, even when the key already has a TTL (counter-style use).""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - mock_redis_instance = AsyncMock() - mock_redis_instance.__aenter__.return_value = mock_redis_instance - mock_redis_instance.__aexit__.return_value = None - mock_redis_instance.ttl.return_value = 42 # key already has ~42s left - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - await redis_cache.async_increment( - key="spend:team_member:u:t", value=0.05, refresh_ttl=True - ) - - mock_redis_instance.expire.assert_awaited_once_with("spend:team_member:u:t", 60) - - -@pytest.mark.asyncio -async def test_redis_cache_async_increment_default_does_not_bump_existing_ttl( - monkeypatch, redis_no_ping -): - """Default (refresh_ttl=False) preserves window-style semantics: TTL is - set only on first creation, never refreshed (used by rate-limit windows).""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - mock_redis_instance = AsyncMock() - mock_redis_instance.__aenter__.return_value = mock_redis_instance - mock_redis_instance.__aexit__.return_value = None - mock_redis_instance.ttl.return_value = 42 # key already has ~42s left - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - await redis_cache.async_increment(key="rate_limit:window", value=1) - - mock_redis_instance.expire.assert_not_awaited() - - @pytest.mark.parametrize("namespace", [None, "litellm"]) @pytest.mark.asyncio async def test_async_delete_cache_applies_namespace( @@ -140,42 +68,6 @@ async def test_redis_client_init_with_socket_timeout(monkeypatch, redis_no_ping) assert client.connection_pool.connection_kwargs["socket_timeout"] == 1.0 -@pytest.mark.asyncio -async def test_redis_cache_async_batch_get_cache(monkeypatch, redis_no_ping): - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - - # Create an AsyncMock for the Redis client - mock_redis_instance = AsyncMock() - - # Make sure the mock can be used as an async context manager - mock_redis_instance.__aenter__.return_value = mock_redis_instance - mock_redis_instance.__aexit__.return_value = None - - # Setup the return value for mget - mock_redis_instance.mget.return_value = [ - b'{"key1": "value1"}', - None, - b'{"key3": "value3"}', - ] - - test_keys = ["key1", "key2", "key3"] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - # Call async_batch_get_cache - result = await redis_cache.async_batch_get_cache(key_list=test_keys) - - # Verify mget was called with the correct keys - mock_redis_instance.mget.assert_called_once() - - # Check that results were properly decoded - assert result["key1"] == {"key1": "value1"} - assert result["key2"] is None - assert result["key3"] == {"key3": "value3"} - - @pytest.mark.asyncio async def test_handle_lpop_count_for_older_redis_versions(monkeypatch): """Test the helper method that handles LPOP with count for Redis versions < 7.0""" @@ -202,41 +94,6 @@ async def test_handle_lpop_count_for_older_redis_versions(monkeypatch): assert mock_pipeline.execute.call_count == 2 -@pytest.mark.asyncio -async def test_async_rpush_pipeline_executes_all_operations(monkeypatch, redis_no_ping): - """Verify that multiple rpush ops are batched into a single pipeline execute""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - - mock_redis_instance = AsyncMock() - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - mock_pipeline.rpush = MagicMock() - mock_pipeline.execute = AsyncMock(return_value=[3, 5, 1]) - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - from litellm.types.caching import RedisPipelineRpushOperation - - rpush_list = [ - RedisPipelineRpushOperation(key="key1", values=["a", "b"]), - RedisPipelineRpushOperation(key="key2", values=["c"]), - RedisPipelineRpushOperation(key="key3", values=["d", "e", "f"]), - ] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - result = await redis_cache.async_rpush_pipeline(rpush_list=rpush_list) - - assert result == [3, 5, 1] - assert mock_pipeline.rpush.call_count == 3 - mock_pipeline.rpush.assert_any_call("key1", "a", "b") - mock_pipeline.rpush.assert_any_call("key2", "c") - mock_pipeline.rpush.assert_any_call("key3", "d", "e", "f") - mock_pipeline.execute.assert_called_once() - - @pytest.mark.asyncio async def test_async_rpush_pipeline_empty_list_returns_empty( monkeypatch, redis_no_ping @@ -256,183 +113,6 @@ async def test_async_rpush_pipeline_empty_list_returns_empty( mock_redis_instance.pipeline.assert_not_called() -@pytest.mark.asyncio -async def test_async_rpush_pipeline_raises_on_redis_error(monkeypatch, redis_no_ping): - """Pipeline errors should propagate""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - - mock_redis_instance = AsyncMock() - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - mock_pipeline.rpush = MagicMock() - mock_pipeline.execute = AsyncMock(side_effect=ConnectionError("Redis down")) - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - from litellm.types.caching import RedisPipelineRpushOperation - - rpush_list = [RedisPipelineRpushOperation(key="key1", values=["a"])] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - with pytest.raises(ConnectionError, match="Redis down"): - await redis_cache.async_rpush_pipeline(rpush_list=rpush_list) - - -@pytest.mark.asyncio -async def test_async_lpop_pipeline_single_round_trip(monkeypatch, redis_no_ping): - """Verify that multiple lpop ops are batched into a single pipeline execute""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - redis_cache.redis_version = "7.0.0" - - mock_redis_instance = AsyncMock() - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - mock_pipeline.lpop = MagicMock() - mock_pipeline.execute = AsyncMock( - return_value=[ - [b"val1", b"val2"], # key1 results - None, # key2 empty - [b"val3"], # key3 results - ] - ) - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - from litellm.types.caching import RedisPipelineLpopOperation - - lpop_list = [ - RedisPipelineLpopOperation(key="key1", count=10), - RedisPipelineLpopOperation(key="key2", count=10), - RedisPipelineLpopOperation(key="key3", count=5), - ] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - results = await redis_cache.async_lpop_pipeline(lpop_list=lpop_list) - - assert len(results) == 3 - assert results[0] == ["val1", "val2"] - assert results[1] is None - assert results[2] == ["val3"] - mock_pipeline.execute.assert_called_once() - - -@pytest.mark.asyncio -async def test_async_lpop_pipeline_redis_lt7_regroups_flat_results( - monkeypatch, redis_no_ping -): - """Verify Redis < 7 fallback issues individual LPOPs and regroups correctly""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - redis_cache.redis_version = "6.2.0" - - mock_redis_instance = AsyncMock() - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - mock_pipeline.lpop = MagicMock() - - # With count=3 for key1 and count=2 for key2, we get 5 individual LPOP commands - # Simulate: key1 has 2 values then None, key2 has 1 value then None - mock_pipeline.execute = AsyncMock( - return_value=[ - b"val1", - b"val2", - None, # 3 LPOPs for key1 - b"val3", - None, # 2 LPOPs for key2 - ] - ) - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - from litellm.types.caching import RedisPipelineLpopOperation - - lpop_list = [ - RedisPipelineLpopOperation(key="key1", count=3), - RedisPipelineLpopOperation(key="key2", count=2), - ] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - results = await redis_cache.async_lpop_pipeline(lpop_list=lpop_list) - - assert len(results) == 2 - assert results[0] == ["val1", "val2"] # 2 values, None filtered out - assert results[1] == ["val3"] # 1 value, None filtered out - # All 5 individual LPOPs should be queued, but only 1 execute() call - assert mock_pipeline.lpop.call_count == 5 - mock_pipeline.execute.assert_called_once() - - -@pytest.mark.asyncio -async def test_async_rpush_pipeline_raises_on_per_command_error( - monkeypatch, redis_no_ping -): - """Verify that per-command errors in pipeline results are raised, not silently dropped""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - - mock_redis_instance = AsyncMock() - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - mock_pipeline.rpush = MagicMock() - # Simulate: first RPUSH succeeds, second returns a per-command error - mock_pipeline.execute = AsyncMock(return_value=[3, Exception("WRONGTYPE")]) - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - from litellm.types.caching import RedisPipelineRpushOperation - - rpush_list = [ - RedisPipelineRpushOperation(key="key1", values=["a"]), - RedisPipelineRpushOperation(key="key2", values=["b"]), - ] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - with pytest.raises(Exception, match="WRONGTYPE"): - await redis_cache.async_rpush_pipeline(rpush_list=rpush_list) - - -@pytest.mark.asyncio -async def test_async_lpop_pipeline_raises_on_per_command_error( - monkeypatch, redis_no_ping -): - """Verify that per-command errors in LPOP pipeline results are raised, not silently dropped""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - redis_cache.redis_version = "7.0.0" - - mock_redis_instance = AsyncMock() - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - mock_pipeline.lpop = MagicMock() - # Simulate: first LPOP succeeds, second returns a per-command error - mock_pipeline.execute = AsyncMock(return_value=[[b"val1"], Exception("WRONGTYPE")]) - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - from litellm.types.caching import RedisPipelineLpopOperation - - lpop_list = [ - RedisPipelineLpopOperation(key="key1", count=10), - RedisPipelineLpopOperation(key="key2", count=10), - ] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - with pytest.raises(Exception, match="WRONGTYPE"): - await redis_cache.async_lpop_pipeline(lpop_list=lpop_list) - - @pytest.mark.asyncio async def test_async_lpop_pipeline_empty_list(monkeypatch, redis_no_ping): """Empty lpop_list should return empty list without touching Redis""" @@ -450,111 +130,6 @@ async def test_async_lpop_pipeline_empty_list(monkeypatch, redis_no_ping): mock_redis_instance.pipeline.assert_not_called() -@pytest.mark.asyncio -async def test_async_lpop_pipeline_propagates_redis_exception( - monkeypatch, redis_no_ping -): - """Pipeline errors should propagate""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - redis_cache.redis_version = "7.0.0" - - mock_redis_instance = AsyncMock() - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - mock_pipeline.lpop = MagicMock() - mock_pipeline.execute = AsyncMock(side_effect=ConnectionError("Redis down")) - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - from litellm.types.caching import RedisPipelineLpopOperation - - lpop_list = [RedisPipelineLpopOperation(key="key1", count=10)] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - with pytest.raises(ConnectionError, match="Redis down"): - await redis_cache.async_lpop_pipeline(lpop_list=lpop_list) - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "redis_version", - [ - # Standard cases - "7.0.0", # Standard Redis string version - 7.0, # Valkey/ElastiCache float version (THE BUG this fix addresses) - 7, # Integer version (e.g., from some Redis forks) - # Version < 7 - "6", # String without dots, version < 7 - # Malformed versions (fallback to 7) - "latest", # Non-numeric version - "", # Empty string - -7.0, # Negative float - # Format variations - " 7.0.0 ", # Whitespace (should be stripped) - "7.0.0-rc1", # Version with suffix - "10.0.0", # Double digit major version - ], -) -async def test_async_lpop_with_float_redis_version( - monkeypatch, redis_no_ping, redis_version -): - """ - Test async_lpop with various Redis version formats (especially float). - - This test specifically addresses the issue where AWS ElastiCache Valkey - returns redis_version as a float (e.g., 7.0) instead of a string (e.g., "7.0.0"), - which caused a 'float' object has no attribute 'split' error when trying to - use the Redis transaction buffer feature. - - The fix converts the version to a string and handles edge cases like: - - Floats (7.0) and integers (7) - - Strings with/without dots ("7" vs "7.0.0") - - Malformed versions ("v7.0.0", "latest") - fallback to version 7 - - Whitespace (" 7.0.0 ") - - Negative versions (fallback to version 7) - - Related: Database deadlock issues when use_redis_transaction_buffer is enabled. - """ - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - - # Create RedisCache instance - redis_cache = RedisCache() - redis_cache.redis_version = redis_version # Set the version to test - - # Create an AsyncMock for the Redis client - mock_redis_instance = AsyncMock() - mock_redis_instance.__aenter__.return_value = mock_redis_instance - mock_redis_instance.__aexit__.return_value = None - - # Mock lpop to return a test value (Redis >= 7.0 behavior) - mock_redis_instance.lpop.return_value = [b"value1", b"value2"] - - # Mock pipeline for Redis < 7.0 (used when major_version < 7) - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - # Make pipeline() a regular method (not async) that returns the mock - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - # Mock handle_lpop_count_for_older_redis_versions for Redis < 7 - with patch.object( - redis_cache, - "handle_lpop_count_for_older_redis_versions", - return_value=[b"value1", b"value2"], - ): - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - # Call async_lpop with count - this should not raise AttributeError - result = await redis_cache.async_lpop(key="test_key", count=2) - - # Verify the method completed without error - assert result is not None - - # LIT-3374: the namespace must be applied uniformly across every key-taking # Redis operation, not just get/set/increment. Before the fix these paths wrote # or read raw keys, so with a namespace configured the prefixed keys other diff --git a/tests/test_litellm/llms/pass_through/guardrail_translation/test_handler.py b/tests/test_litellm/llms/pass_through/guardrail_translation/test_handler.py index f8bd83fc7df..1043c26c6ec 100644 --- a/tests/test_litellm/llms/pass_through/guardrail_translation/test_handler.py +++ b/tests/test_litellm/llms/pass_through/guardrail_translation/test_handler.py @@ -1,15 +1,10 @@ """ -Tests for LlmPassthroughRouteHandler and the guardrail_translation_mappings registry. +Tests for the guardrail_translation_mappings registry. Validates: - allm_passthrough_route is registered in the mappings (regression: this was the bug) -- Bedrock provider is dispatched to BedrockPassthroughGuardrailHandler -- Unknown provider skips apply_guardrail """ -import pytest -from unittest.mock import AsyncMock, MagicMock, patch - from litellm.llms.pass_through.guardrail_translation import ( guardrail_translation_mappings, ) @@ -40,185 +35,3 @@ class TestRegistry: is PassThroughEndpointHandler ) - -def _make_guardrail() -> MagicMock: - g = MagicMock() - g.guardrail_name = "test-guard" - g.apply_guardrail = AsyncMock(return_value={"texts": []}) - g.skip_system_message_in_guardrail = False - g.skip_tool_message_in_guardrail = False - return g - - -class TestLlmPassthroughRouteHandlerInput: - @pytest.mark.asyncio - async def test_bedrock_provider_delegates_to_bedrock_handler(self): - handler = LlmPassthroughRouteHandler() - data = { - "custom_llm_provider": "bedrock", - "endpoint": "model/anthropic.claude-3-sonnet/converse", - "data": {"messages": [{"role": "user", "content": [{"text": "hi"}]}]}, - } - guardrail = _make_guardrail() - - await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) - - guardrail.apply_guardrail.assert_called_once() - - @pytest.mark.asyncio - async def test_unknown_provider_skips_apply_guardrail(self): - handler = LlmPassthroughRouteHandler() - data = { - "custom_llm_provider": "some_unknown_provider", - "endpoint": "v1/chat/completions", - "data": {"messages": [{"role": "user", "content": "hi"}]}, - } - guardrail = _make_guardrail() - - result = await handler.process_input_messages( - data=data, guardrail_to_apply=guardrail - ) - - guardrail.apply_guardrail.assert_not_called() - assert result is data - - @pytest.mark.asyncio - async def test_missing_provider_skips(self): - handler = LlmPassthroughRouteHandler() - data = {"endpoint": "foo/bar", "data": {}} - guardrail = _make_guardrail() - - result = await handler.process_input_messages( - data=data, guardrail_to_apply=guardrail - ) - - guardrail.apply_guardrail.assert_not_called() - assert result is data - - -class TestLlmPassthroughRouteHandlerOutput: - @pytest.mark.asyncio - async def test_bedrock_provider_delegates_output_to_bedrock_handler(self): - handler = LlmPassthroughRouteHandler() - response = { - "output": { - "message": { - "role": "assistant", - "content": [{"text": "hello"}], - } - } - } - request_data = { - "custom_llm_provider": "bedrock", - "endpoint": "model/anthropic.claude-3-sonnet/converse", - } - guardrail = _make_guardrail() - - await handler.process_output_response( - response=response, - guardrail_to_apply=guardrail, - request_data=request_data, - ) - - guardrail.apply_guardrail.assert_called_once() - - @pytest.mark.asyncio - async def test_unknown_provider_skips_output(self): - handler = LlmPassthroughRouteHandler() - response = {"some": "response"} - request_data = {"custom_llm_provider": "unknown"} - guardrail = _make_guardrail() - - result = await handler.process_output_response( - response=response, - guardrail_to_apply=guardrail, - request_data=request_data, - ) - - guardrail.apply_guardrail.assert_not_called() - assert result is response - - -class TestDeAnonymizeEventStream: - @pytest.mark.asyncio - async def test_bedrock_provider_dispatches_to_handler(self): - body = b"original-stream-bytes" - expected = b"de-anonymized-bytes" - proxy_logging_obj = MagicMock() - user_api_key_dict = MagicMock() - - with patch( - "litellm.llms.bedrock.passthrough.guardrail_translation.handler." - "BedrockPassthroughGuardrailHandler.de_anonymize_event_stream", - new=AsyncMock(return_value=expected), - ) as mock_handler: - result = await LlmPassthroughRouteHandler.de_anonymize_event_stream( - body_bytes=body, - proxy_logging_obj=proxy_logging_obj, - user_api_key_dict=user_api_key_dict, - data={"custom_llm_provider": "bedrock"}, - ) - - mock_handler.assert_awaited_once() - assert result == expected - - @pytest.mark.asyncio - async def test_unknown_provider_returns_original_bytes(self): - body = b"original-stream-bytes" - - result = await LlmPassthroughRouteHandler.de_anonymize_event_stream( - body_bytes=body, - proxy_logging_obj=MagicMock(), - user_api_key_dict=MagicMock(), - data={"custom_llm_provider": "anthropic"}, - ) - - assert result is body - - @pytest.mark.asyncio - async def test_missing_provider_returns_original_bytes(self): - body = b"original-stream-bytes" - - result = await LlmPassthroughRouteHandler.de_anonymize_event_stream( - body_bytes=body, - proxy_logging_obj=MagicMock(), - user_api_key_dict=MagicMock(), - data={}, - ) - - assert result is body - - -class TestSupportsEventStreamDeAnonymization: - def test_bedrock_converse_stream_is_supported(self): - assert ( - LlmPassthroughRouteHandler.supports_event_stream_de_anonymization( - "bedrock", "model/us.amazon.nova-lite-v1:0/converse-stream" - ) - is True - ) - - def test_bedrock_invoke_stream_is_not_supported(self): - assert ( - LlmPassthroughRouteHandler.supports_event_stream_de_anonymization( - "bedrock", - "model/us.amazon.nova-lite-v1:0/invoke-with-response-stream", - ) - is False - ) - - def test_unknown_provider_is_not_supported(self): - assert ( - LlmPassthroughRouteHandler.supports_event_stream_de_anonymization( - "anthropic", "model/foo/converse-stream" - ) - is False - ) - - def test_missing_provider_is_not_supported(self): - assert ( - LlmPassthroughRouteHandler.supports_event_stream_de_anonymization( - None, "model/foo/converse-stream" - ) - is False - ) diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 54db0c0fd4f..040fba7c53a 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -1206,50 +1206,6 @@ def test_team_info_masking(): assert "public-test-key" not in str(exc_info.value) -def test_embedding_input_array_of_tokens(client_no_auth): - """ - Test to bypass decoding input as array of tokens for selected providers - - Ref: https://github.com/BerriAI/litellm/issues/10113 - """ - from litellm.proxy import proxy_server - - # The client_no_auth fixture should initialize the router - # Assert this to catch any router initialization regressions - assert proxy_server.llm_router is not None, ( - "llm_router is None after client_no_auth fixture initialized. " - "This indicates a router initialization issue that should be investigated." - ) - - try: - with mock.patch.object( - proxy_server.llm_router, - "aembedding", - return_value=example_embedding_result, - ) as mock_aembedding: - test_data = { - "model": "vllm_embed_model", - "input": [[2046, 13269, 158208]], - } - - response = client_no_auth.post("/v1/embeddings", json=test_data) - - # Assert that aembedding was called, and that input was not modified - mock_aembedding.assert_called_once() - call_args, call_kwargs = mock_aembedding.call_args - assert call_kwargs["model"] == "vllm_embed_model" - assert call_kwargs["input"] == [[2046, 13269, 158208]] - - assert response.status_code == 200 - result = response.json() - print(len(result["data"][0]["embedding"])) - assert ( - len(result["data"][0]["embedding"]) > 10 - ) # this usually has len==1536 so - except Exception as e: - pytest.fail(f"LiteLLM Proxy test failed. Exception - {str(e)}") - - @pytest.mark.asyncio async def test_get_all_team_models(): """ From 7d9eec623081f72abeb3770001828ab24b490503 Mon Sep 17 00:00:00 2001 From: shivam Date: Fri, 24 Jul 2026 20:32:09 +0000 Subject: [PATCH 2/9] fix(proxy): return 400 instead of 500 for chat completions without messages Router.acompletion() takes messages positionally, so splatting a body that omits it raised a TypeError that the generic handler mapped to a 500. Validate the required body param at the routing boundary and raise the existing 400 contract instead. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/route_llm_request.py | 28 ++++++- .../proxy/test_route_llm_request.py | 77 ++++++++++++++++--- 2 files changed, 94 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 25fa0819930..1f5aacc2115 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -1,5 +1,5 @@ import asyncio -from typing import TYPE_CHECKING, Any, Literal, Optional +from typing import TYPE_CHECKING, Any, Literal, Mapping, Optional import httpx from fastapi import HTTPException, status @@ -145,6 +145,30 @@ class ProxyModelNotFoundError(HTTPException): super().__init__(status_code=status.HTTP_400_BAD_REQUEST, detail=detail) +REQUIRED_BODY_PARAM_BY_ROUTE: Mapping[str, str] = { + "acompletion": "messages", + "aembedding": "input", +} + + +class ProxyMissingRequiredParamError(HTTPException): + def __init__(self, route: str, param: str): + detail = {"error": f"{route}: Missing required parameter: '{param}'."} + super().__init__(status_code=status.HTTP_400_BAD_REQUEST, detail=detail) + self.type = "invalid_request_error" + self.param = param + + +def raise_if_required_body_param_missing(route_type: str, data: Mapping[str, object]) -> None: + required_param = REQUIRED_BODY_PARAM_BY_ROUTE.get(route_type) + if required_param is None or data.get(required_param) is not None: + return + raise ProxyMissingRequiredParamError( + route=ROUTE_ENDPOINT_MAPPING.get(route_type, route_type), + param=required_param, + ) + + def get_team_id_from_data(data: dict) -> Optional[str]: """ Get the team id from the data's metadata or litellm_metadata params. @@ -353,6 +377,8 @@ async def route_request( """ Common helper to route the request """ + raise_if_required_body_param_missing(route_type=route_type, data=data) + await add_shared_session_to_data(data) # Strip router-internal mock_testing_* flags. Combined with an diff --git a/tests/test_litellm/proxy/test_route_llm_request.py b/tests/test_litellm/proxy/test_route_llm_request.py index f506b9665a6..93b3ef1cce8 100644 --- a/tests/test_litellm/proxy/test_route_llm_request.py +++ b/tests/test_litellm/proxy/test_route_llm_request.py @@ -12,24 +12,25 @@ from litellm.proxy.route_llm_request import ProxyModelNotFoundError, route_reque @pytest.mark.parametrize( - "route_type", + "route_type, required_body_params", [ - "atext_completion", - "acompletion", - "aembedding", - "aimage_generation", - "aspeech", - "atranscription", - "amoderation", - "arerank", + ("atext_completion", {}), + ("acompletion", {"messages": [{"role": "user", "content": "Hello"}]}), + ("aembedding", {"input": "Hello"}), + ("aimage_generation", {}), + ("aspeech", {}), + ("atranscription", {}), + ("amoderation", {}), + ("arerank", {}), ], ) @pytest.mark.asyncio -async def test_route_request_dynamic_credentials(route_type): +async def test_route_request_dynamic_credentials(route_type, required_body_params): data = { "model": "openai/gpt-4o-mini-2024-07-18", "api_key": "my-bad-key", "api_base": "https://api.openai.com/v1 ", + **required_body_params, } llm_router = MagicMock() # Ensure that the dynamic method exists on the llm_router mock. @@ -887,3 +888,59 @@ async def test_route_request_override_enable_tag_filtering_beats_body_value(): call_kwargs = llm_router.acompletion.call_args[1] assert call_kwargs["enable_tag_filtering"] is True + + +@pytest.mark.parametrize( + "route_type, param, route", + [ + ("acompletion", "messages", "/chat/completions"), + ("aembedding", "input", "/embeddings"), + ], +) +@pytest.mark.parametrize("data_extra", [{}, {"messages": None, "input": None}]) +def test_raise_if_required_body_param_missing_rejects_missing_param(route_type, param, route, data_extra): + from litellm.proxy.route_llm_request import ( + ProxyMissingRequiredParamError, + raise_if_required_body_param_missing, + ) + + with pytest.raises(ProxyMissingRequiredParamError) as exc_info: + raise_if_required_body_param_missing(route_type=route_type, data={"model": "gpt-4o", **data_extra}) + + assert exc_info.value.status_code == 400 + assert exc_info.value.param == param + assert exc_info.value.type == "invalid_request_error" + assert exc_info.value.detail == {"error": f"{route}: Missing required parameter: '{param}'."} + + +@pytest.mark.parametrize( + "route_type, data", + [ + ("acompletion", {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}]}), + ("acompletion", {"model": "gpt-4o", "messages": []}), + ("atext_completion", {"model": "gpt-4o"}), + ("aembedding", {"model": "text-embedding-3-small", "input": "hi"}), + ("arerank", {"model": "rerank-model"}), + ("aimage_generation", {"model": "dall-e-3"}), + ], +) +def test_raise_if_required_body_param_missing_allows_valid_requests(route_type, data): + from litellm.proxy.route_llm_request import raise_if_required_body_param_missing + + raise_if_required_body_param_missing(route_type=route_type, data=data) + + +@pytest.mark.asyncio +async def test_route_request_rejects_chat_completion_without_messages(): + """A /chat/completions body without `messages` used to splat into + Router.acompletion() and surface the resulting TypeError as a 500.""" + from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError + + llm_router = MagicMock() + + with pytest.raises(ProxyMissingRequiredParamError) as exc_info: + await route_request({"model": "gpt-4o"}, llm_router, None, "acompletion") + + assert exc_info.value.status_code == 400 + assert exc_info.value.param == "messages" + llm_router.acompletion.assert_not_called() From a376f724002185917340d7d74252b5ef1894c5cb Mon Sep 17 00:00:00 2001 From: shivam Date: Fri, 24 Jul 2026 20:53:05 +0000 Subject: [PATCH 3/9] fix(responses): stop treating stream_options as a Responses API param Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/types/llms/openai.py | 1 - ...erimental_pass_through_messages_handler.py | 64 +++++++++++++++++++ .../test_responses_api_request_body.py | 27 ++++++++ 3 files changed, 91 insertions(+), 1 deletion(-) diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 9f689a2dd31..2263d53182a 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -1171,7 +1171,6 @@ class ResponsesAPIOptionalRequestParams(TypedDict, total=False): max_tool_calls: Optional[int] prompt_cache_key: Optional[str] prompt_cache_retention: Optional[str] - stream_options: Optional[dict] top_logprobs: Optional[int] partial_images: Optional[int] # Number of partial images to generate (1-3) for streaming image generation context_management: Optional[List[ContextManagementEntry]] diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py index 3327fc39f73..8875a75e86f 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py @@ -2,6 +2,7 @@ import json import os import sys +import httpx import pytest from fastapi.testclient import TestClient @@ -9,6 +10,7 @@ sys.path.insert(0, os.path.abspath("../../../../..")) from unittest.mock import AsyncMock, MagicMock, patch +import litellm from litellm.anthropic_interface import messages from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.types.utils import Delta, ModelResponse, StreamingChoices @@ -37,6 +39,68 @@ def test_anthropic_experimental_pass_through_messages_handler(): assert mock_responses.call_args.kwargs["api_key"] == "test-api-key" +@pytest.mark.asyncio +async def test_openai_model_does_not_forward_stream_options_to_responses_api(): + """ + Regression test for LIT-4779. `always_include_stream_usage` injects + stream_options={'include_usage': True} into every streaming request, but OpenAI + models on /v1/messages go to the Responses API, which 400s on that param. + """ + responses_payload = { + "id": "resp_stream_options", + "object": "response", + "created_at": 1734366691, + "status": "completed", + "model": "gpt-5.5", + "output": [ + { + "type": "message", + "id": "msg_1", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "hi", "annotations": []}], + } + ], + "parallel_tool_calls": True, + "usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2}, + "error": None, + "incomplete_details": None, + "instructions": None, + "metadata": None, + "temperature": None, + "tool_choice": "auto", + "tools": [], + "top_p": None, + "max_output_tokens": None, + "previous_response_id": None, + "reasoning": None, + "truncation": None, + "user": None, + } + + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.text = json.dumps(responses_payload) + mock_response.headers = httpx.Headers({}) + mock_response.json.return_value = responses_payload + + with patch.object(AsyncHTTPHandler, "post", new_callable=AsyncMock) as mock_post: + mock_post.return_value = mock_response + + await litellm.anthropic.messages.acreate( + max_tokens=100, + messages=[{"role": "user", "content": "Hello, how are you?"}], + model="openai/gpt-5.5", + api_key="test-api-key", + stream_options={"include_usage": True}, + ) + + mock_post.assert_called_once() + post_kwargs = mock_post.call_args.kwargs + request_body = post_kwargs["json"] if "json" in post_kwargs else json.loads(post_kwargs["data"]) + assert "stream_options" not in request_body + + def test_anthropic_experimental_pass_through_messages_handler_dynamic_api_key_and_api_base_and_custom_values(): """ Test that api key, api base, and extra kwargs are forwarded to litellm.completion for Azure models. diff --git a/tests/test_litellm/responses/test_responses_api_request_body.py b/tests/test_litellm/responses/test_responses_api_request_body.py index 44dfa240d42..2922c9738aa 100644 --- a/tests/test_litellm/responses/test_responses_api_request_body.py +++ b/tests/test_litellm/responses/test_responses_api_request_body.py @@ -198,6 +198,33 @@ async def test_aresponses_azure_shell_tool_400_maps_to_bad_request_error(): assert "not supported" in str(excinfo.value).lower() +@pytest.mark.asyncio +async def test_aresponses_drops_stream_options(): + """ + stream_options is a Chat Completions param; the Responses API rejects it with + "Unknown parameter: 'stream_options.include_usage'". It must never reach the wire. + """ + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = MockResponse( + _minimal_responses_api_payload("resp_stream_options_test", "gpt-5.5"), 200 + ) + + await litellm.aresponses( + model="openai/gpt-5.5", + api_key="fake-api-key", + input="hi", + stream_options={"include_usage": True}, + ) + + mock_post.assert_called_once() + post_kwargs = mock_post.call_args.kwargs + request_body = post_kwargs["json"] if "json" in post_kwargs else json.loads(post_kwargs["data"]) + assert "stream_options" not in request_body + + @pytest.mark.asyncio async def test_aresponses_request_level_drop_params_drops_bedrock_mantle_service_tier( monkeypatch, From 9f9714c209b3a00083d5bb3f93c1894aa93841dc Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Fri, 24 Jul 2026 16:05:18 -0700 Subject: [PATCH 4/9] test: remove five more zero-kill tests from test_http_handler The http_handler pair only received its full mutation verdict after the first removal batch landed; these five tests pass unchanged when every function they execute is mutated and the owning file killed none of their scored mutants. The ssl tests excluded from mutation scoring are untouched. --- .../llms/custom_httpx/test_http_handler.py | 111 ------------------ 1 file changed, 111 deletions(-) diff --git a/tests/test_litellm/llms/custom_httpx/test_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_http_handler.py index 7bd1d7a6031..87d67e0e8b7 100644 --- a/tests/test_litellm/llms/custom_httpx/test_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_http_handler.py @@ -181,28 +181,6 @@ async def test_force_ipv4_transport(): litellm.disable_aiohttp_transport = original_disable -@pytest.mark.asyncio -async def test_ssl_context_transport(): - """Test transport creation with SSL context""" - # Create a test SSL context - ssl_context = ssl.create_default_context() - - transport = AsyncHTTPHandler._create_async_transport(ssl_context=ssl_context) - assert transport is not None - - try: - if isinstance(transport, LiteLLMAiohttpTransport): - # Get the client session and verify SSL context is passed through - client_session = transport._get_valid_client_session() - assert isinstance(client_session, ClientSession) - assert isinstance(client_session.connector, TCPConnector) - # Verify the connector has SSL context set by checking if it's using SSL - assert client_session.connector._ssl is not None - finally: - if isinstance(transport, LiteLLMAiohttpTransport): - await transport.aclose() - - @pytest.mark.asyncio async def test_aiohttp_disabled_transport(): """Test transport creation with aiohttp disabled""" @@ -339,44 +317,6 @@ async def test_ssl_context_with_shared_session(): litellm.disable_aiohttp_transport = original_disable -@pytest.mark.asyncio -async def test_aiohttp_transport_trust_env_setting(monkeypatch): - """Test that trust_env setting is properly configured in aiohttp transport""" - transports = [] - try: - # Test 1: Default trust_env behavior - transport = AsyncHTTPHandler._create_aiohttp_transport() - transports.append(transport) - client_session = transport._get_valid_client_session() - - # Default should be False (litellm.aiohttp_trust_env default) - default_trust_env = getattr(litellm, "aiohttp_trust_env", False) - assert client_session._trust_env == default_trust_env - - # Test 2: Environment variable override - monkeypatch.setenv("AIOHTTP_TRUST_ENV", "True") - transport_with_env = AsyncHTTPHandler._create_aiohttp_transport() - transports.append(transport_with_env) - client_session_with_env = transport_with_env._get_valid_client_session() - - # Should be True when environment variable is set - assert client_session_with_env._trust_env is True - - # Test 3: Verify environment variable with False value - monkeypatch.setenv("AIOHTTP_TRUST_ENV", "False") - transport_with_false_env = AsyncHTTPHandler._create_aiohttp_transport() - transports.append(transport_with_false_env) - client_session_with_false_env = ( - transport_with_false_env._get_valid_client_session() - ) - - # Should respect the litellm.aiohttp_trust_env setting when env var is False - assert client_session_with_false_env._trust_env == default_trust_env - finally: - for t in transports: - await t.aclose() - - def test_get_ssl_configuration(): """Test that get_ssl_configuration() returns a proper SSL context with certifi CA bundle when no environment variables are set.""" @@ -443,36 +383,6 @@ async def test_create_aiohttp_transport_with_shared_session(): assert not callable(transport.client) # Should not be callable -@pytest.mark.asyncio -async def test_create_aiohttp_transport_without_shared_session(): - """Test that _create_aiohttp_transport creates new session when none provided""" - from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler - - # Test without shared session - transport = AsyncHTTPHandler._create_aiohttp_transport(shared_session=None) - - # Verify the transport uses a lambda function (for backward compatibility) - assert callable(transport.client) # Should be a lambda function - - -@pytest.mark.asyncio -async def test_create_aiohttp_transport_with_closed_session(): - """Test that _create_aiohttp_transport creates new session when shared session is closed""" - from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler - - # Create a mock closed session - mock_session = MockClientSession() - mock_session.closed = True - - # Test with closed session - transport = AsyncHTTPHandler._create_aiohttp_transport( - shared_session=mock_session # type: ignore - ) - - # Verify the transport creates a new session (lambda function) - assert callable(transport.client) # Should be a lambda function - - @pytest.mark.asyncio async def test_async_handler_with_shared_session(): """Test AsyncHTTPHandler initialization with shared session""" @@ -622,27 +532,6 @@ async def test_session_reuse_integration(): await client2.close() -@pytest.mark.asyncio -async def test_session_validation(): - """Test that session validation works correctly""" - from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler - - # Test with None session - transport1 = AsyncHTTPHandler._create_aiohttp_transport(shared_session=None) - assert callable(transport1.client) # Should create lambda - - # Test with closed session - mock_closed_session = MockClientSession() - mock_closed_session.closed = True - transport2 = AsyncHTTPHandler._create_aiohttp_transport(shared_session=mock_closed_session) # type: ignore - assert callable(transport2.client) # Should create lambda - - # Test with valid session - mock_valid_session = MockClientSession() - transport3 = AsyncHTTPHandler._create_aiohttp_transport(shared_session=mock_valid_session) # type: ignore - assert transport3.client is mock_valid_session # Should reuse session - - @pytest.mark.parametrize( "env_curve,litellm_curve,expected_curve,should_call", [ From 3476240f11bbe288d3e79f47b4e91973ca57dc6c Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Fri, 24 Jul 2026 16:07:23 -0700 Subject: [PATCH 5/9] fix(ui): keep entity usage tabs aligned with their panels Tremor's TabPanels hands each child an index via React.Children.map, while the selected index comes from HeadlessUI counting only real Tab elements. An empty fragment, false, or null still consumes a panel index but contributes no tab, so the team-only Agent Activity conditional made the two lists drift for every non-team entity type: Key Activity resolved to the empty slot and rendered nothing at all, and Endpoint Activity rendered the key metrics Drive both lists from a single tab array so adding or removing a conditional tab touches one place and the indices cannot diverge --- .../EntityUsage/EntityUsage.test.tsx | 61 +- .../components/EntityUsage/EntityUsage.tsx | 618 +++++++++--------- 2 files changed, 364 insertions(+), 315 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx index cbf3a2cc1f6..89c38c6274f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx @@ -25,8 +25,17 @@ vi.mock("@/components/networking", () => ({ // Mock the child components to simplify testing vi.mock("@/components/activity_metrics", () => ({ - ActivityMetrics: () =>
Activity Metrics
, - processActivityData: () => ({ data: [], metadata: {} }), + ActivityMetrics: ({ modelMetrics }: { modelMetrics?: { __source?: string } }) => ( +
+ Activity Metrics + {`metrics-source:${modelMetrics?.__source ?? "none"}`} +
+ ), + processActivityData: (_data: unknown, key: string) => ({ __source: key }), +})); + +vi.mock("../EndpointUsage/EndpointUsage", () => ({ + default: () =>
Endpoint Usage Panel
, })); vi.mock("@/components/UsagePage/components/EntityUsage/TopKeyView", () => ({ @@ -481,6 +490,54 @@ describe("EntityUsage", () => { expect(screen.getAllByText("Activity Metrics")[1]).toBeInTheDocument(); }); + const selectedPanels = (container: HTMLElement) => + Array.from(container.querySelectorAll("div.tremor-TabPanel-root")).filter( + (panel) => panel.getAttribute("aria-selected") === "true", + ); + + it.each([ + ["Cost", "Tag Spend Overview"], + ["Model Activity", "metrics-source:models"], + ["Key Activity", "metrics-source:api_keys"], + ["Endpoint Activity", "Endpoint Usage Panel"], + ])("shows only the %s panel for a non-team entity type", async (tabLabel, marker) => { + const { container } = render(); + + await waitFor(() => { + expect(mockTagDailyActivityCall).toHaveBeenCalled(); + }); + + act(() => { + fireEvent.click(screen.getByText(tabLabel)); + }); + + const selected = selectedPanels(container); + expect(selected).toHaveLength(1); + expect(selected[0].textContent).toContain(marker); + }); + + it.each([ + ["Cost", "Team Spend Overview"], + ["Model Activity", "metrics-source:models"], + ["Agent Activity", "metrics-source:entities"], + ["Key Activity", "metrics-source:api_keys"], + ["Endpoint Activity", "Endpoint Usage Panel"], + ])("shows only the %s panel for the team entity type", async (tabLabel, marker) => { + const { container } = render(); + + await waitFor(() => { + expect(mockTeamDailyActivityCall).toHaveBeenCalled(); + }); + + act(() => { + fireEvent.click(screen.getByText(tabLabel)); + }); + + const selected = selectedPanels(container); + expect(selected).toHaveLength(1); + expect(selected[0].textContent).toContain(marker); + }); + it("should handle empty data gracefully", async () => { const emptyData = { results: [], diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx index 534e2be7fe8..e330983b6f9 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx @@ -25,7 +25,7 @@ import { } from "@tremor/react"; import { ExportOutlined, LoadingOutlined } from "@ant-design/icons"; import { Alert, Button } from "antd"; -import React, { useMemo, useState } from "react"; +import React, { type ReactNode, useMemo, useState } from "react"; import TeamMultiSelect from "@/components/common_components/team_multi_select"; import { ActivityMetrics, processActivityData } from "@/components/activity_metrics"; import { UsageExportHeader } from "@/components/EntityUsageExport"; @@ -406,6 +406,304 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti const capitalizedEntityLabel = entityType.charAt(0).toUpperCase() + entityType.slice(1); + const costPanel = ( + + {/* Total Spend Card */} + + + {capitalizedEntityLabel} Spend Overview + + + Total Spend + + ${formatNumberWithCommas(spendData.metadata.total_spend, 2)} + + + + Total Requests + {spendData.metadata.total_api_requests.toLocaleString()} + + + Successful Requests + + {spendData.metadata.total_successful_requests.toLocaleString()} + + + + Failed Requests + + {spendData.metadata.total_failed_requests.toLocaleString()} + + + + Total Tokens + {spendData.metadata.total_tokens.toLocaleString()} + + + + + + {/* Daily Spend Chart */} + + + + Daily Spend + + + new Date(a.date).getTime() - new Date(b.date).getTime())} + index="date" + categories={["metrics.spend"]} + colors={["cyan"]} + valueFormatter={valueFormatterSpend} + yAxisWidth={100} + showLegend={false} + customTooltip={({ payload, active }) => { + if (!active || !payload?.[0]) return null; + const data = payload[0].payload; + const entityCount = Object.keys(data.breakdown.entities || {}).length; + return ( +
+

{data.date}

+

Total Spend: ${formatNumberWithCommas(data.metrics.spend, 2)}

+

Total Requests: {data.metrics.api_requests}

+

Successful: {data.metrics.successful_requests}

+

Failed: {data.metrics.failed_requests}

+

Total Tokens: {data.metrics.total_tokens}

+

+ Total {capitalizedEntityLabel}s: {entityCount} +

+
+

Spend by {capitalizedEntityLabel}:

+ {Object.entries(data.breakdown.entities || {}) + .sort(([, a], [, b]) => { + const spendA = (a as EntityMetrics).metrics.spend; + const spendB = (b as EntityMetrics).metrics.spend; + return spendB - spendA; + }) + .slice(0, 5) + .map(([entity, entityData]) => { + const metrics = entityData as EntityMetrics; + return ( +

+ {getEntityLabel(entity, metrics.metadata)}: $ + {formatNumberWithCommas(metrics.metrics.spend, 2)} +

+ ); + })} + {entityCount > 5 &&

...and {entityCount - 5} more

} +
+
+ ); + }} + /> +
+
+ + + {/* Entity Breakdown Section */} + + +
+
+ Spend Per {capitalizedEntityLabel} + Showing Top 5 by Spend +
+ Get Started by Tracking cost per {capitalizedEntityLabel} + + here + +
+
+ + + { + if (!active || !payload?.[0]) return null; + const data = payload[0].payload; + return ( +
+

{data.metadata.alias}

+

Spend: ${formatNumberWithCommas(data.metrics.spend, 4)}

+

Requests: {data.metrics.api_requests.toLocaleString()}

+

+ Successful: {data.metrics.successful_requests.toLocaleString()} +

+

Failed: {data.metrics.failed_requests.toLocaleString()}

+

Tokens: {data.metrics.total_tokens.toLocaleString()}

+
+ ); + }} + /> + + +
+ + + + {capitalizedEntityLabel} + Spend + Successful + Failed + Tokens + + + + {getEntityBreakdown() + .filter((entity) => entity.metrics.spend > 0) + .map((entity) => ( + + {entity.metadata.alias} + + + + + {entity.metrics.successful_requests.toLocaleString()} + + + {entity.metrics.failed_requests.toLocaleString()} + + {entity.metrics.total_tokens.toLocaleString()} + + ))} + +
+
+ +
+
+
+ + + {/* Top API Keys */} + + + Top Virtual Keys + + + + + {/* Top Models */} + + + {entityType === "agent" ? "Top Agents" : "Top Models"} + + + + + {/* Top Agents - only for team entity type */} + {entityType === "team" && ( + + + Top Agents Driving Spend + + + + )} + + {/* Spend by Provider */} + + +
+ Provider Usage + + + `$${formatNumberWithCommas(value, 2)}`} + colors={["cyan", "blue", "indigo", "violet", "purple"]} + showLabel + startAngle={90} + endAngle={-270} + /> + + + + + + Provider + Spend + Successful + Failed + Tokens + + + + {getProviderSpend().map((provider) => ( + + +
+ {provider.provider && } + {provider.provider} +
+
+ + + + + {provider.successful_requests.toLocaleString()} + + {provider.failed_requests.toLocaleString()} + {provider.tokens.toLocaleString()} +
+ ))} +
+
+ +
+
+
+ +
+ ); + + const tabs: readonly { key: string; label: string; content: ReactNode }[] = [ + { key: "cost", label: "Cost", content: costPanel }, + { + key: "models", + label: entityType === "agent" ? "Request / Token Consumption" : "Model Activity", + content: , + }, + ...(entityType === "team" + ? [{ key: "agents", label: "Agent Activity", content: }] + : []), + { + key: "keys", + label: "Key Activity", + content: , + }, + { key: "endpoints", label: "Endpoint Activity", content: }, + ]; + return (
{isFetchingMore && ( @@ -501,320 +799,14 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti /> - Cost - {entityType === "agent" ? "Request / Token Consumption" : "Model Activity"} - {entityType === "team" ? Agent Activity : <>} - Key Activity - Endpoint Activity + {tabs.map(({ key, label }) => ( + {label} + ))} - - - {/* Total Spend Card */} - - - {capitalizedEntityLabel} Spend Overview - - - Total Spend - - ${formatNumberWithCommas(spendData.metadata.total_spend, 2)} - - - - Total Requests - - {spendData.metadata.total_api_requests.toLocaleString()} - - - - Successful Requests - - {spendData.metadata.total_successful_requests.toLocaleString()} - - - - Failed Requests - - {spendData.metadata.total_failed_requests.toLocaleString()} - - - - Total Tokens - - {spendData.metadata.total_tokens.toLocaleString()} - - - - - - - {/* Daily Spend Chart */} - - - - Daily Spend - - - new Date(a.date).getTime() - new Date(b.date).getTime(), - )} - index="date" - categories={["metrics.spend"]} - colors={["cyan"]} - valueFormatter={valueFormatterSpend} - yAxisWidth={100} - showLegend={false} - customTooltip={({ payload, active }) => { - if (!active || !payload?.[0]) return null; - const data = payload[0].payload; - const entityCount = Object.keys(data.breakdown.entities || {}).length; - return ( -
-

{data.date}

-

- Total Spend: ${formatNumberWithCommas(data.metrics.spend, 2)} -

-

Total Requests: {data.metrics.api_requests}

-

Successful: {data.metrics.successful_requests}

-

Failed: {data.metrics.failed_requests}

-

Total Tokens: {data.metrics.total_tokens}

-

- Total {capitalizedEntityLabel}s: {entityCount} -

-
-

Spend by {capitalizedEntityLabel}:

- {Object.entries(data.breakdown.entities || {}) - .sort(([, a], [, b]) => { - const spendA = (a as EntityMetrics).metrics.spend; - const spendB = (b as EntityMetrics).metrics.spend; - return spendB - spendA; - }) - .slice(0, 5) - .map(([entity, entityData]) => { - const metrics = entityData as EntityMetrics; - return ( -

- {getEntityLabel(entity, metrics.metadata)}: $ - {formatNumberWithCommas(metrics.metrics.spend, 2)} -

- ); - })} - {entityCount > 5 && ( -

...and {entityCount - 5} more

- )} -
-
- ); - }} - /> -
-
- - - {/* Entity Breakdown Section */} - - -
-
- Spend Per {capitalizedEntityLabel} - Showing Top 5 by Spend -
- Get Started by Tracking cost per {capitalizedEntityLabel} - - here - -
-
- - - { - if (!active || !payload?.[0]) return null; - const data = payload[0].payload; - return ( -
-

{data.metadata.alias}

-

Spend: ${formatNumberWithCommas(data.metrics.spend, 4)}

-

Requests: {data.metrics.api_requests.toLocaleString()}

-

- Successful: {data.metrics.successful_requests.toLocaleString()} -

-

Failed: {data.metrics.failed_requests.toLocaleString()}

-

Tokens: {data.metrics.total_tokens.toLocaleString()}

-
- ); - }} - /> - - -
- - - - {capitalizedEntityLabel} - Spend - Successful - Failed - Tokens - - - - {getEntityBreakdown() - .filter((entity) => entity.metrics.spend > 0) - .map((entity) => ( - - {entity.metadata.alias} - - - - - {entity.metrics.successful_requests.toLocaleString()} - - - {entity.metrics.failed_requests.toLocaleString()} - - {entity.metrics.total_tokens.toLocaleString()} - - ))} - -
-
- -
-
-
- - - {/* Top API Keys */} - - - Top Virtual Keys - - - - - {/* Top Models */} - - - {entityType === "agent" ? "Top Agents" : "Top Models"} - - - - - {/* Top Agents - only for team entity type */} - {entityType === "team" && ( - - - Top Agents Driving Spend - - - - )} - - {/* Spend by Provider */} - - -
- Provider Usage - - - `$${formatNumberWithCommas(value, 2)}`} - colors={["cyan", "blue", "indigo", "violet", "purple"]} - showLabel - startAngle={90} - endAngle={-270} - /> - - - - - - Provider - Spend - Successful - Failed - Tokens - - - - {getProviderSpend().map((provider) => ( - - -
- {provider.provider && } - {provider.provider} -
-
- - - - - {provider.successful_requests.toLocaleString()} - - - {provider.failed_requests.toLocaleString()} - - {provider.tokens.toLocaleString()} -
- ))} -
-
- -
-
-
- -
-
- - - - {entityType === "team" ? ( - - - - ) : ( - <> - )} - - - - - - + {tabs.map(({ key, content }) => ( + {content} + ))}
From 770f41b5fa0d42c9c6d9f7c822f486e8671f00f5 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Thu, 23 Jul 2026 21:01:37 -0700 Subject: [PATCH 6/9] fix(guardrails): keep guardrail information in spend logs when the caller sends its own metadata The guardrail-information writer picked its metadata bucket with a hand-rolled precedence that preferred a caller-supplied `metadata` field, while every reader resolves the bucket through `get_metadata_variable_name_from_kwargs`, which prefers `litellm_metadata`. The two rules agree only when the caller sends no `metadata` of its own. Routes in `LITELLM_METADATA_ROUTES` seed `litellm_metadata`, so on /v1/messages and /v1/responses a caller that sends `metadata` sent the entry to a dict nothing reads; the spend log then reported `guardrail_status: not_run` with no `guardrail_information` even though the guardrail ran and the `x-litellm-applied-guardrails` header was present. Give the resolver one owner. `get_or_create_metadata_bucket` moves from the proxy layer into core_helpers next to the resolver it calls, so `litellm/integrations` can reach it without a proxy dependency, and the byte-identical duplicate of `get_metadata_variable_name_from_kwargs` in callback_utils is deleted. The writer now shares that owner with `add_guardrail_to_applied_guardrails_header`, so the response header and the spend log can no longer disagree. Two readers had to move with it or the fix would be a no-op on the affected routes. `_sync_guardrail_info_to_logging_obj`, which bridges request_data into the spend-log payload for passthrough routes, picked the first truthy bucket, so a non-empty caller `metadata` short-circuited it. The otel failure-path span reader `_emit_guardrail_spans_from_request_data` read a hard-coded `metadata` key, which also dropped the span whenever the entry lived in `litellm_metadata`. Model Armor already resolved the bucket for its file-scan results but wrote its text-scan and post-call results, and read them back in `_process_response`, through a hard-coded `metadata` key; on a seeded route that split the record so a file scan's evidence never reached the logger. All four Model Armor sites now use the shared resolver. The unified guardrail hook seeds `litellm_metadata` on every route, so the OpenAI moderation entry lands there too; spend-log output is unchanged because `merge_litellm_metadata` reads both buckets. --- litellm/integrations/custom_guardrail.py | 21 ++++----- litellm/integrations/opentelemetry.py | 11 +++-- litellm/litellm_core_utils/core_helpers.py | 19 ++++++++ litellm/proxy/common_utils/callback_utils.py | 46 ++++-------------- .../model_armor/model_armor.py | 19 +++++--- .../pass_through_endpoints.py | 13 ++--- .../integrations/test_custom_guardrail.py | 47 +++++++++++++++++++ .../test_guardrail_logging_sync.py | 27 +++++++++-- .../test_otel_guardrail_violation_spans.py | 42 +++++++++++++++++ .../litellm_core_utils/test_core_helpers.py | 41 ++++++++++++++++ .../openai/test_moderations.py | 11 +++-- .../test_openai_moderation_streaming.py | 11 +++-- .../guardrail_hooks/test_model_armor.py | 38 +++++++++++++++ 13 files changed, 271 insertions(+), 75 deletions(-) diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index f639ad49d5e..fc62a7ad59c 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -17,7 +17,11 @@ from typing import ( ) from litellm._logging import verbose_logger -from litellm.litellm_core_utils.core_helpers import redact_nested_match_and_regex_keys +from litellm.litellm_core_utils.core_helpers import ( + get_metadata_variable_name_from_kwargs, + get_or_create_metadata_bucket, + redact_nested_match_and_regex_keys, +) from litellm.caching import DualCache from litellm.integrations.custom_logger import CustomLogger from litellm.secret_managers.main import str_to_bool @@ -954,17 +958,8 @@ class CustomGuardrail(CustomLogger): # should not happen container[key] = [existing, slg] - if "metadata" in request_data: - if request_data["metadata"] is None: - request_data["metadata"] = {} - _append_guardrail_info(request_data["metadata"]) - elif "litellm_metadata" in request_data: - _append_guardrail_info(request_data["litellm_metadata"]) - else: - # Ensure guardrail info is always logged (e.g. proxy may not have set - # metadata yet). Attach to "metadata" so spend log / standard logging see it. - request_data["metadata"] = {} - _append_guardrail_info(request_data["metadata"]) + _, metadata_bucket = get_or_create_metadata_bucket(request_data) + _append_guardrail_info(metadata_bucket) _guardrail_self_recorded.set(True) @@ -1223,7 +1218,7 @@ def _sync_guardrail_info_to_logging_obj(request_data: dict, logging_obj: object) """ if logging_obj is None: return - meta_src = request_data.get("metadata") or request_data.get("litellm_metadata") or {} + meta_src = request_data.get(get_metadata_variable_name_from_kwargs(request_data)) or {} slg_info = meta_src.get("standard_logging_guardrail_information") if not slg_info: return diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index fea55cd1db4..12465377b51 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -883,8 +883,8 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): request_data: dict, parent_span: Optional[Any], ) -> None: - """Emit ``guardrail`` spans from ``request_data["metadata"] - ["standard_logging_guardrail_information"]``. + """Emit ``guardrail`` spans from the request's proxy-internal metadata bucket + (``standard_logging_guardrail_information``). Routed through ``_create_guardrail_span`` so the dedupe state in ``_otel_internal`` is honoured — if ``_handle_failure`` already @@ -892,7 +892,12 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): """ from opentelemetry import trace as _trace - metadata = (request_data or {}).get("metadata") or {} + from litellm.litellm_core_utils.core_helpers import ( + get_metadata_variable_name_from_kwargs, + ) + + request_data = request_data or {} + metadata = request_data.get(get_metadata_variable_name_from_kwargs(request_data)) or {} guardrail_information = metadata.get("standard_logging_guardrail_information") if not guardrail_information: return diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py index 88dddb59cc7..cecc35ee1c1 100644 --- a/litellm/litellm_core_utils/core_helpers.py +++ b/litellm/litellm_core_utils/core_helpers.py @@ -195,6 +195,25 @@ def get_metadata_variable_name_from_kwargs( return "litellm_metadata" if "litellm_metadata" in kwargs else "metadata" +def get_or_create_metadata_bucket( + request_data: dict, +) -> tuple[Literal["metadata", "litellm_metadata"], dict]: + """ + Return the proxy-internal metadata bucket for this request, creating it if absent. + + Batch/file routes store proxy state in ``litellm_metadata`` so the OpenAI + ``metadata`` field can remain provider-safe (string values only). Every writer and + reader of proxy-internal metadata resolves the bucket through here, so a caller that + supplies its own ``metadata`` field cannot split them across two dicts. + """ + metadata_key = get_metadata_variable_name_from_kwargs(request_data) + metadata_bucket = request_data.get(metadata_key) + if not isinstance(metadata_bucket, dict): + metadata_bucket = {} + request_data[metadata_key] = metadata_bucket + return metadata_key, metadata_bucket + + def get_litellm_metadata_from_kwargs(kwargs: dict): """ Helper to get litellm metadata from all litellm request kwargs diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index a9c2a12aff7..33bca782e0b 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -1,12 +1,16 @@ import copy import os -from typing import TYPE_CHECKING, Any, Callable, Dict, Iterable, List, Literal, Optional +from typing import TYPE_CHECKING, Any, Callable, Dict, Iterable, List, Optional import litellm from litellm import get_secret from litellm._logging import verbose_proxy_logger from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.core_helpers import ( + get_metadata_variable_name_from_kwargs, + get_or_create_metadata_bucket, +) from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker from litellm.proxy._types import CommonProxyErrors, LiteLLMPromptInjectionParams from litellm.proxy.common_utils.encrypt_decrypt_utils import ( @@ -406,23 +410,6 @@ def get_logging_caching_headers(request_data: Dict) -> Optional[Dict]: return headers -def get_metadata_variable_name_from_kwargs( - kwargs: dict, -) -> Literal["metadata", "litellm_metadata"]: - """ - Helper to return what the "metadata" field should be called in the request data - - - New endpoints return `litellm_metadata` - - Old endpoints return `metadata` - - Context: - - LiteLLM used `metadata` as an internal field for storing metadata - - OpenAI then started using this field for their metadata - - LiteLLM is now moving to using `litellm_metadata` for our metadata - """ - return "litellm_metadata" if "litellm_metadata" in kwargs else "metadata" - - LITELLM_PROXY_INTERNAL_METADATA_KEYS = frozenset( { "applied_policies", @@ -450,23 +437,6 @@ LITELLM_PROXY_INTERNAL_METADATA_KEYS = frozenset( ) -def _get_or_create_proxy_metadata_bucket( - request_data: Dict, -) -> tuple[Literal["metadata", "litellm_metadata"], dict]: - """ - Return the proxy-internal metadata bucket for this request. - - Batch/file routes store proxy state in ``litellm_metadata`` so the OpenAI - ``metadata`` field can remain provider-safe (string values only). - """ - metadata_key = get_metadata_variable_name_from_kwargs(request_data) - metadata_bucket = request_data.get(metadata_key) - if not isinstance(metadata_bucket, dict): - metadata_bucket = {} - request_data[metadata_key] = metadata_bucket - return metadata_key, metadata_bucket - - def sanitize_openai_provider_metadata( metadata: Optional[Dict[str, Any]], ) -> Optional[Dict[str, str]]: @@ -496,7 +466,7 @@ def sanitize_openai_provider_metadata( def add_guardrail_to_applied_guardrails_header(request_data: Dict, guardrail_name: Optional[str]): if guardrail_name is None: return - _, _metadata = _get_or_create_proxy_metadata_bucket(request_data) + _, _metadata = get_or_create_metadata_bucket(request_data) if "applied_guardrails" in _metadata: if guardrail_name not in _metadata["applied_guardrails"]: _metadata["applied_guardrails"].append(guardrail_name) @@ -513,7 +483,7 @@ def add_policy_to_applied_policies_header(request_data: Dict, policy_name: Optio """ if policy_name is None: return - _, _metadata = _get_or_create_proxy_metadata_bucket(request_data) + _, _metadata = get_or_create_metadata_bucket(request_data) if "applied_policies" in _metadata: if policy_name not in _metadata["applied_policies"]: _metadata["applied_policies"].append(policy_name) @@ -531,7 +501,7 @@ def add_policy_sources_to_metadata(request_data: Dict, policy_sources: Dict[str, """ if not policy_sources: return - _, _metadata = _get_or_create_proxy_metadata_bucket(request_data) + _, _metadata = get_or_create_metadata_bucket(request_data) existing = _metadata.get("policy_sources", {}) if not isinstance(existing, dict): existing = {} diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py index 28b9dec100f..f2c5a95202b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py @@ -30,6 +30,10 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) +from litellm.litellm_core_utils.core_helpers import ( + get_metadata_variable_name_from_kwargs, + get_or_create_metadata_bucket, +) from litellm.llms.vertex_ai.vertex_llm_base import VertexBase from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.guardrails.guardrail_hooks.model_armor.file_scanning import ( @@ -432,7 +436,11 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): Override to store only the Model Armor API response, not the entire data dict. This prevents circular references in logging. """ - metadata = (request_data.get("metadata") or {}) if isinstance(request_data, dict) else {} + metadata = ( + request_data.get(get_metadata_variable_name_from_kwargs(request_data)) or {} + if isinstance(request_data, dict) + else {} + ) guardrail_response = metadata.get("_model_armor_response", {}) # Determine status – default to "success" but prefer the explicit value if present. @@ -471,7 +479,6 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): blocking, while fail_on_error still governs real Model Armor API errors. """ from litellm.proxy.common_utils.callback_utils import ( - _get_or_create_proxy_metadata_bucket, add_guardrail_to_applied_guardrails_header, ) @@ -491,7 +498,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name) # Use the same metadata bucket the header helper writes to, so the logged Model Armor # payload and status land where _process_response reads them on every route. - _, metadata = _get_or_create_proxy_metadata_bucket(data) + _, metadata = get_or_create_metadata_bucket(data) fail_on_error = bool(self.optional_params.get("fail_on_error", True)) if unscannable_references > 0: @@ -607,7 +614,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): # overwritten by another coroutine. blocked = self._should_block_content(armor_response, allow_sanitization=self.mask_request_content) if isinstance(data, dict): - metadata = data.setdefault("metadata", {}) # ensures metadata exists and is unique per request + _, metadata = get_or_create_metadata_bucket(data) # ensures metadata exists and is unique per request # Accumulate so a prior file scan on the same request is not overwritten by this text scan. metadata["_model_armor_response"] = self._append_armor_response( metadata.get("_model_armor_response"), @@ -702,7 +709,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): blocked = self._should_block_content(armor_response, allow_sanitization=self.mask_request_content) # Store the armor response for logging if isinstance(data, dict): - metadata = data.setdefault("metadata", {}) + _, metadata = get_or_create_metadata_bucket(data) # Accumulate so a prior file scan on the same request is not overwritten by this text scan. metadata["_model_armor_response"] = self._append_armor_response( metadata.get("_model_armor_response"), @@ -868,7 +875,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): # Attach Model Armor response & status to this request's metadata to avoid race conditions if isinstance(request_data, dict): - metadata = request_data.setdefault("metadata", {}) + _, metadata = get_or_create_metadata_bucket(request_data) metadata["_model_armor_response"] = self._build_logging_response(armor_response) metadata["_model_armor_status"] = ( "blocked" if self._should_block_content(armor_response) else "success" diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index acb2e50c79b..9364d7eae3a 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -38,6 +38,10 @@ from litellm._uuid import uuid from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.core_helpers import ( + get_metadata_variable_name_from_kwargs, + get_or_create_metadata_bucket, +) from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.litellm_core_utils.safe_json_dumps import safe_dumps @@ -668,16 +672,13 @@ def _carry_guardrail_logging_info(request_data: dict, guardrail_data: Optional[d """ if guardrail_data is None: return - source_metadata = guardrail_data.get("metadata") - if not isinstance(source_metadata, dict): - return + source_key = get_metadata_variable_name_from_kwargs(guardrail_data) + source_metadata = guardrail_data.get(source_key) or {} entries = source_metadata.get("standard_logging_guardrail_information") if not entries: return - metadata = request_data.get("metadata") - if not isinstance(metadata, dict): - metadata = request_data["metadata"] = {} + _, metadata = get_or_create_metadata_bucket(request_data) metadata.setdefault("standard_logging_guardrail_information", list(entries)) diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py index 4ea79f9e2a4..0ef26c1c5f4 100644 --- a/tests/test_litellm/integrations/test_custom_guardrail.py +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -654,6 +654,53 @@ class TestGuardrailLoggingAggregation: assert len(info) == 2 assert info[1]["guardrail_name"] == "test_guardrail" + def test_caller_metadata_does_not_divert_the_entry_from_the_reader(self): + """A caller-supplied `metadata` field must not send the entry to a bucket the + spend log never reads. Routes in LITELLM_METADATA_ROUTES (/v1/messages, + /v1/responses, batches, files) seed `litellm_metadata`, and Claude Code sends + `metadata.user_id`, so both keys are present on the same request.""" + request_data = { + "metadata": {"user_id": "device-account-session"}, + "litellm_metadata": {"user_api_key_hash": "abc"}, + } + + self._invoke_add_log(request_data) + + assert ( + "standard_logging_guardrail_information" not in request_data["metadata"] + ), "entry landed in the caller's metadata, where the spend log does not read it" + info = request_data["litellm_metadata"][ + "standard_logging_guardrail_information" + ] + assert len(info) == 1 + assert info[0]["guardrail_name"] == "test_guardrail" + + def test_entry_and_applied_guardrails_header_share_one_bucket(self): + """The x-litellm-applied-guardrails writer and the guardrail-info writer must + resolve the same bucket, otherwise the response header and the spend log + disagree about whether the guardrail ran.""" + from litellm.proxy.common_utils.callback_utils import ( + add_guardrail_to_applied_guardrails_header, + ) + + request_data = { + "metadata": {"user_id": "device-account-session"}, + "litellm_metadata": {}, + } + + self._invoke_add_log(request_data) + add_guardrail_to_applied_guardrails_header( + request_data=request_data, guardrail_name="test_guardrail" + ) + + buckets = { + key + for key in ("metadata", "litellm_metadata") + for field in ("standard_logging_guardrail_information", "applied_guardrails") + if field in request_data[key] + } + assert buckets == {"litellm_metadata"} + class TestGuardrailOtelSpanEmission: """Recording a guardrail emits its otel span inline, so every guardrail diff --git a/tests/test_litellm/integrations/test_guardrail_logging_sync.py b/tests/test_litellm/integrations/test_guardrail_logging_sync.py index 5dcd1114b3d..f9e1a3efbd0 100644 --- a/tests/test_litellm/integrations/test_guardrail_logging_sync.py +++ b/tests/test_litellm/integrations/test_guardrail_logging_sync.py @@ -59,8 +59,11 @@ def test_syncs_from_metadata_key(): assert result == [entry] -def test_metadata_wins_over_litellm_metadata(): - """metadata key takes precedence over litellm_metadata when both are present.""" +def test_litellm_metadata_wins_over_caller_metadata(): + """When both keys are present the helper must read the bucket the writer used, + which get_or_create_metadata_bucket resolves to litellm_metadata. Reading the + caller's metadata instead is how a guardrail entry went missing from spend logs + on the routes that seed litellm_metadata.""" entry_meta = _make_slg_entry("from-metadata") entry_lm = _make_slg_entry("from-litellm_metadata") request_data = { @@ -74,7 +77,25 @@ def test_metadata_wins_over_litellm_metadata(): result = logging_obj.litellm_params["metadata"].get( "standard_logging_guardrail_information" ) - assert result == [entry_meta] + assert result == [entry_lm] + + +def test_syncs_when_caller_sends_its_own_metadata(): + """The Claude Code shape: caller metadata present, guardrail entry in the seeded + litellm_metadata bucket. The entry must still reach the spend-log payload.""" + entry = _make_slg_entry() + request_data = { + "metadata": {"user_id": "device-account-session"}, + "litellm_metadata": {"standard_logging_guardrail_information": [entry]}, + } + logging_obj = _FakeLogging() + + _sync_guardrail_info_to_logging_obj(request_data, logging_obj) + + result = logging_obj.litellm_params["metadata"].get( + "standard_logging_guardrail_information" + ) + assert result == [entry] def test_noop_when_no_guardrail_info(): diff --git a/tests/test_litellm/integrations/test_otel_guardrail_violation_spans.py b/tests/test_litellm/integrations/test_otel_guardrail_violation_spans.py index ace9399cf53..c3e9d67ddad 100644 --- a/tests/test_litellm/integrations/test_otel_guardrail_violation_spans.py +++ b/tests/test_litellm/integrations/test_otel_guardrail_violation_spans.py @@ -279,6 +279,48 @@ class TestGuardrailSpanOnViolation(unittest.TestCase): parent_span.context.span_id, ) + def test_post_call_failure_hook_emits_span_when_caller_sends_metadata(self): + """On routes that seed ``litellm_metadata`` the guardrail entry lives there, + not in the caller's own ``metadata`` field. Reading a hard-coded ``metadata`` + key drops the span for exactly the requests that carry both.""" + otel, provider, exporter = _make_otel() + parent_span = provider.get_tracer(__name__).start_span(PROXY_SPAN_NAME) + + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-test", + parent_otel_span=parent_span, + request_route="/v1/messages", + ) + + request_data = { + "model": "claude-haiku", + "messages": [{"role": "user", "content": "Hello"}], + "metadata": {"user_id": "device-account-session"}, + "litellm_metadata": { + "standard_logging_guardrail_information": [ + _slg_entry("guardrail_intervened", _bedrock_block_response()) + ], + }, + } + + _run( + otel.async_post_call_failure_hook( + request_data=request_data, + original_exception=Exception("guardrail blocked"), + user_api_key_dict=user_api_key_dict, + ) + ) + + guardrail_spans = [ + s for s in exporter.get_finished_spans() if s.name == GUARDRAIL_SPAN_NAME + ] + self.assertEqual( + len(guardrail_spans), + 1, + "the guardrail span must be emitted from the resolved metadata bucket, " + "not from a hard-coded 'metadata' key", + ) + def test_handle_failure_and_post_call_failure_hook_dedupe(self): """When _handle_failure and async_post_call_failure_hook BOTH fire for the same request (the production flow on a guardrail block), diff --git a/tests/test_litellm/litellm_core_utils/test_core_helpers.py b/tests/test_litellm/litellm_core_utils/test_core_helpers.py index b67ea91bb0b..b4f539da286 100644 --- a/tests/test_litellm/litellm_core_utils/test_core_helpers.py +++ b/tests/test_litellm/litellm_core_utils/test_core_helpers.py @@ -4,12 +4,53 @@ import pytest from litellm.litellm_core_utils.core_helpers import ( _FINISH_REASON_MAP, + get_or_create_metadata_bucket, map_finish_reason, reconstruct_model_name, redact_nested_match_and_regex_keys, ) +class TestGetOrCreateMetadataBucket: + """The single owner every guardrail writer and reader shares, so the response + header and the spend log can never disagree about which dict a record lives in.""" + + def test_prefers_litellm_metadata_when_both_present(self): + request_data = {"metadata": {"user_id": "caller"}, "litellm_metadata": {}} + + key, bucket = get_or_create_metadata_bucket(request_data) + + assert key == "litellm_metadata" + assert bucket is request_data["litellm_metadata"] + + def test_uses_metadata_when_litellm_metadata_absent(self): + request_data = {"metadata": {"user_id": "caller"}} + + key, bucket = get_or_create_metadata_bucket(request_data) + + assert key == "metadata" + assert bucket is request_data["metadata"] + + def test_creates_the_bucket_in_place_when_missing(self): + request_data: dict = {} + + key, bucket = get_or_create_metadata_bucket(request_data) + + assert key == "metadata" + assert request_data["metadata"] is bucket + bucket["k"] = "v" + assert request_data["metadata"]["k"] == "v" + + def test_replaces_a_non_dict_bucket(self): + request_data = {"litellm_metadata": None} + + key, bucket = get_or_create_metadata_bucket(request_data) + + assert key == "litellm_metadata" + assert isinstance(request_data["litellm_metadata"], dict) + assert bucket is request_data["litellm_metadata"] + + def test_reconstruct_model_name_prefers_deployment_value(): """Ensure deployment metadata wins when reconstructing the model name.""" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py index fe6cb98d1f5..9002d1f81a3 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py @@ -729,10 +729,15 @@ async def test_openai_moderation_post_call_request_data_passthrough(): mock_make_request.assert_called_once() - # Guardrail info in the REAL request_data (not a throwaway) - guardrail_info_list = request_data["metadata"].get( - "standard_logging_guardrail_information" + # Guardrail info in the REAL request_data (not a throwaway). The unified hook + # seeds litellm_metadata, so read the bucket the resolver names rather than + # assuming "metadata"; the spend log reads it the same way. + from litellm.litellm_core_utils.core_helpers import ( + get_metadata_variable_name_from_kwargs, ) + + bucket = request_data[get_metadata_variable_name_from_kwargs(request_data)] + guardrail_info_list = bucket.get("standard_logging_guardrail_information") assert guardrail_info_list is not None assert isinstance(guardrail_info_list[0]["guardrail_response"], dict) assert "results" in guardrail_info_list[0]["guardrail_response"] diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_openai_moderation_streaming.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_openai_moderation_streaming.py index 0358ca998aa..914af0e2368 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_openai_moderation_streaming.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_openai_moderation_streaming.py @@ -259,10 +259,15 @@ async def test_openai_moderation_streaming_end_of_stream_request_data_passthroug ): pass - # Verify guardrail info reached the REAL request_data (not a throwaway) - guardrail_info_list = request_data["metadata"].get( - "standard_logging_guardrail_information" + # Verify guardrail info reached the REAL request_data (not a throwaway). The + # unified hook seeds litellm_metadata, so read the bucket the resolver names + # rather than assuming "metadata"; the spend log reads it the same way. + from litellm.litellm_core_utils.core_helpers import ( + get_metadata_variable_name_from_kwargs, ) + + bucket = request_data[get_metadata_variable_name_from_kwargs(request_data)] + guardrail_info_list = bucket.get("standard_logging_guardrail_information") assert ( guardrail_info_list is not None ), "Guardrail info should be in request_data after streaming" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py index 18b5bd92411..89b6af27719 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py @@ -3502,6 +3502,44 @@ async def test_single_scan_response_stays_a_dict(): assert isinstance(request_data["metadata"]["_model_armor_response"], dict) +@pytest.mark.asyncio +async def test_scan_result_reaches_the_logger_on_a_seeded_route(): + """On routes that seed `litellm_metadata` the scan result must land in that bucket + and be found by `_process_response`. Writing the file-scan result through the shared + resolver while the text-scan writers and the reader used a hard-coded `metadata` key + split the record in two, so the logged guardrail payload came back empty.""" + guardrail = _make_guardrail() + pdf_b64 = base64.b64encode(PDF_BYTES).decode("utf-8") + request_data = { + "model": "claude-haiku", + "messages": [_file_message(pdf_b64)], + "metadata": {"user_id": "device-account-session"}, + "litellm_metadata": {"guardrails": ["model-armor-test"]}, + } + + with patch.object( + guardrail.async_handler, + "post", + AsyncMock(return_value=_armor_response(blocked=False)), + ): + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(), + cache=MagicMock(spec=DualCache), + data=request_data, + call_type="completion", + ) + + assert "_model_armor_response" not in request_data["metadata"] + assert "_model_armor_response" in request_data["litellm_metadata"] + + before = len(request_data["litellm_metadata"].get("standard_logging_guardrail_information", [])) + guardrail._process_response(response=None, request_data=request_data) + + logged = request_data["litellm_metadata"]["standard_logging_guardrail_information"] + assert len(logged) == before + 1 + assert logged[-1]["guardrail_response"], "the logger recorded an empty Model Armor payload" + + @pytest.mark.asyncio async def test_pre_call_blocks_supported_document_with_undecodable_base64(): """A supported document whose inline base64 will not decode cannot be scanned, so it fails closed.""" From 9777e9524a641f56195f191746d53df0764c94aa Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Thu, 23 Jul 2026 11:33:07 -0700 Subject: [PATCH 7/9] fix(guardrails): stop reporting a no-op guardrail as applied on passthrough On passthrough requests the shared guardrail plumbing still dispatches headroom's pre_call apply_guardrail, but the passthrough translation hands it only `texts` and no `structured_messages`, so it early-returns a no-op. The @log_guardrail_information decorator then synthesized an "allow"/"success" StandardLoggingGuardrailInformation entry, and the unified hook added the guardrail to applied_guardrails, so spend logs reported the compression guardrail as succeeded even though nothing ran. Add a records_own_guardrail_information flag for guardrails that log their own execution (headroom). The decorator skips the synthetic success entry for them, and the unified hook lists such a guardrail in applied_guardrails only when it actually recorded a run. A guardrail that owns its logging must record every outcome it runs, so headroom now records a guardrail_failed_to_respond entry on the fail_open path (compression attempted, service unreachable, request forwarded uncompressed) instead of leaving it unlogged; fail_closed is still recorded by the decorator's error path, and a genuine no-op stays not_run. --- litellm/integrations/custom_guardrail.py | 14 ++- .../guardrail_hooks/headroom/headroom.py | 19 ++- .../unified_guardrail/unified_guardrail.py | 12 +- .../integrations/test_custom_guardrail.py | 59 +++++++++- .../guardrail_hooks/test_headroom.py | 76 +++++++++++- .../test_unified_guardrail.py | 111 +++++++++++++++++- 6 files changed, 277 insertions(+), 14 deletions(-) diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index f639ad49d5e..cebf4a311f7 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -107,6 +107,8 @@ class CustomGuardrail(CustomLogger): # If True, during_call runs async_moderation_hook instead of the unified apply_guardrail path. use_native_during_call_hook: ClassVar[bool] = False + records_own_guardrail_information: ClassVar[bool] = False + def __init__( self, guardrail_name: Optional[str] = None, @@ -1256,6 +1258,14 @@ def log_guardrail_information(func): so it stays correct when guardrails run concurrently (asyncio copies the context into each gathered task): counting shared entries would let one guardrail's append hide another guardrail's missing record. + + A guardrail that only records an entry when it actually runs (e.g. + ``HeadroomGuardrail``, which returns the inputs untouched on an endpoint + whose payload it cannot act on) sets ``records_own_guardrail_information = + True`` so the auto-record is skipped even on the return paths where it + recorded nothing; otherwise a no-op early return would be logged as an + "allow"/"success" run even though the guardrail did nothing. The exception + branch below still records so a genuine failure is not lost. """ import functools import inspect @@ -1291,7 +1301,7 @@ def log_guardrail_information(func): self_recorded_token = _guardrail_self_recorded.set(False) try: response = await func(*args, **kwargs) - if _guardrail_self_recorded.get(): + if self.records_own_guardrail_information or _guardrail_self_recorded.get(): return response return self._process_response( response=response, @@ -1333,7 +1343,7 @@ def log_guardrail_information(func): self_recorded_token = _guardrail_self_recorded.set(False) try: response = func(*args, **kwargs) - if _guardrail_self_recorded.get(): + if self.records_own_guardrail_information or _guardrail_self_recorded.get(): return response return self._process_response( response=response, diff --git a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py index 2d67c22f0aa..ca3bb0ee361 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py +++ b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py @@ -4,7 +4,7 @@ import json import re import time import uuid -from typing import TYPE_CHECKING, Any, List, Literal, Optional +from typing import TYPE_CHECKING, Any, ClassVar, List, Literal, Optional import httpx from fastapi import HTTPException @@ -209,6 +209,8 @@ def _build_responses_followup_items( class HeadroomGuardrail(CustomGuardrail): + records_own_guardrail_information: ClassVar[bool] = True + @classmethod def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]: return [ @@ -481,7 +483,21 @@ class HeadroomGuardrail(CustomGuardrail): ) end_time = time.time() + from litellm.proxy.common_utils.callback_utils import ( + add_guardrail_to_applied_guardrails_header, + ) + if not compression_succeeded: + self.add_standard_logging_guardrail_information_to_request_data( + guardrail_json_response={"error": "headroom compression unavailable; request forwarded uncompressed"}, + request_data=request_data, + guardrail_status="guardrail_failed_to_respond", + guardrail_provider=HEADROOM_GUARDRAIL_PROVIDER, + start_time=start_time, + end_time=end_time, + duration=end_time - start_time, + ) + add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name) return {**inputs, "structured_messages": compressed} # pyright: ignore[reportReturnType] self.add_standard_logging_guardrail_information_to_request_data( @@ -493,6 +509,7 @@ class HeadroomGuardrail(CustomGuardrail): end_time=end_time, duration=end_time - start_time, ) + add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name) hashes = extract_hashes_from_messages(compressed) if not hashes: diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 8e3abfbf159..d4d23cd2e37 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -147,8 +147,10 @@ class UnifiedLLMGuardrails(CustomLogger): litellm_logging_obj=data.get("litellm_logging_obj"), ) - # Add guardrail to applied guardrails header - add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=guardrail_to_apply.guardrail_name) + if not guardrail_to_apply.records_own_guardrail_information: + add_guardrail_to_applied_guardrails_header( + request_data=data, guardrail_name=guardrail_to_apply.guardrail_name + ) return data async def async_moderation_hook( @@ -274,8 +276,10 @@ class UnifiedLLMGuardrails(CustomLogger): if e.original_response is None: e.original_response = response raise - # Add guardrail to applied guardrails header - add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=guardrail_to_apply.guardrail_name) + if not guardrail_to_apply.records_own_guardrail_information: + add_guardrail_to_applied_guardrails_header( + request_data=data, guardrail_name=guardrail_to_apply.guardrail_name + ) return response diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py index 4ea79f9e2a4..219908239d4 100644 --- a/tests/test_litellm/integrations/test_custom_guardrail.py +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -3,9 +3,12 @@ from unittest.mock import AsyncMock import pytest -from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + log_guardrail_information, +) from litellm.proxy._types import CallTypes, UserAPIKeyAuth -from litellm.types.utils import GuardrailTracingDetail +from litellm.types.utils import GenericGuardrailAPIInputs, GuardrailTracingDetail class TestCustomGuardrailDeploymentHook: @@ -1947,3 +1950,55 @@ class TestOnlyScanNewMessages: cache.async_set_cache = AsyncMock(side_effect=RuntimeError("redis down")) await guardrail.mark_texts_scanned(texts=["a"], request_data={"litellm_session_id": "s1"}, cache=cache) + + +def _guardrail_entries(request_data: dict) -> list: + container = request_data.get("metadata") or request_data.get("litellm_metadata") or {} + entries = container.get("standard_logging_guardrail_information") + return entries if isinstance(entries, list) else [] + + +class _NoopGuardrail(CustomGuardrail): + """apply_guardrail that returns the inputs untouched and records nothing.""" + + @log_guardrail_information + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): + return inputs + + +class _NoopSelfLoggingGuardrail(_NoopGuardrail): + records_own_guardrail_information = True + + +class TestRecordsOwnGuardrailInformation: + """The @log_guardrail_information decorator must not synthesize an "allow"/"success" + entry for a no-op apply_guardrail when the guardrail sets + records_own_guardrail_information (LIT-4650).""" + + @pytest.mark.asyncio + async def test_default_noop_apply_guardrail_is_auto_logged(self): + guardrail = _NoopGuardrail(guardrail_name="g1") + request_data: dict = {"model": "gpt-4o"} + + await guardrail.apply_guardrail( + inputs=GenericGuardrailAPIInputs(texts=["x"]), + request_data=request_data, + input_type="request", + ) + + entries = _guardrail_entries(request_data) + assert len(entries) == 1 + assert entries[0]["guardrail_status"] == "success" + + @pytest.mark.asyncio + async def test_self_logging_noop_apply_guardrail_is_not_logged(self): + guardrail = _NoopSelfLoggingGuardrail(guardrail_name="g2") + request_data: dict = {"model": "gpt-4o"} + + await guardrail.apply_guardrail( + inputs=GenericGuardrailAPIInputs(texts=["x"]), + request_data=request_data, + input_type="request", + ) + + assert _guardrail_entries(request_data) == [] diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py index 7f412c008ca..776df985d46 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py @@ -114,6 +114,24 @@ def guardrail() -> HeadroomGuardrail: return _make_guardrail() +def _recorded_guardrail_entries(request_data: dict) -> list: + for container_key in ("metadata", "litellm_metadata"): + container = request_data.get(container_key) + if isinstance(container, dict): + entries = container.get("standard_logging_guardrail_information") + if isinstance(entries, list): + return entries + return [] + + +def _applied_guardrails(request_data: dict) -> list: + for container_key in ("metadata", "litellm_metadata"): + container = request_data.get(container_key) + if isinstance(container, dict) and isinstance(container.get("applied_guardrails"), list): + return container["applied_guardrails"] + return [] + + @pytest.mark.asyncio async def test_apply_guardrail_compresses_and_returns_structured_messages( guardrail: HeadroomGuardrail, @@ -123,6 +141,7 @@ async def test_apply_guardrail_compresses_and_returns_structured_messages( structured_messages=ORIGINAL_MESSAGES, ) mock_response = _make_compress_response(COMPRESSED_MESSAGES) + request_data = {"model": "gpt-4o"} with patch.object( guardrail.async_handler, @@ -132,12 +151,19 @@ async def test_apply_guardrail_compresses_and_returns_structured_messages( ): result = await guardrail.apply_guardrail( inputs=inputs, - request_data={"model": "gpt-4o"}, + request_data=request_data, input_type="request", ) assert result.get("structured_messages") == COMPRESSED_MESSAGES + entries = _recorded_guardrail_entries(request_data) + assert len(entries) == 1 + assert entries[0]["guardrail_name"] == "headroom" + assert entries[0]["guardrail_status"] == "success" + assert entries[0]["guardrail_provider"] == "headroom" + assert "headroom" in _applied_guardrails(request_data) + @pytest.mark.asyncio async def test_apply_guardrail_injects_retrieve_tool_when_hashes_present( @@ -719,6 +745,7 @@ async def test_apply_guardrail_bypass_header_skips_compression( mock_post.assert_not_called() assert result.get("structured_messages") == ORIGINAL_MESSAGES + assert _recorded_guardrail_entries(request_data) == [] @pytest.mark.asyncio @@ -729,16 +756,18 @@ async def test_apply_guardrail_response_type_passthrough( texts=["some response text"], structured_messages=ORIGINAL_MESSAGES, ) + request_data: dict = {"model": "gpt-4o"} with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post: result = await guardrail.apply_guardrail( inputs=inputs, - request_data={}, + request_data=request_data, input_type="response", ) mock_post.assert_not_called() assert result is inputs + assert _recorded_guardrail_entries(request_data) == [] @pytest.mark.asyncio @@ -746,16 +775,48 @@ async def test_apply_guardrail_empty_structured_messages_passthrough( guardrail: HeadroomGuardrail, ): inputs = GenericGuardrailAPIInputs(texts=["hello"]) + request_data: dict = {"model": "gpt-4o"} with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post: result = await guardrail.apply_guardrail( inputs=inputs, - request_data={}, + request_data=request_data, input_type="request", ) mock_post.assert_not_called() assert result is inputs + assert _recorded_guardrail_entries(request_data) == [] + assert "headroom" not in _applied_guardrails(request_data) + + +@pytest.mark.asyncio +async def test_passthrough_handler_does_not_log_headroom_as_run( + guardrail: HeadroomGuardrail, +): + """Regression for LIT-4650. + + A passthrough request drives headroom through PassThroughEndpointHandler, which + only supplies `texts` (no `structured_messages`). Headroom cannot compress that + shape and no-ops, so it must not appear in the spend log's + standard_logging_guardrail_information as a successful run. + """ + from litellm.llms.pass_through.guardrail_translation.handler import ( + PassThroughEndpointHandler, + ) + + data = {"model": "gpt-4o", "messages": [{"role": "user", "content": "hello"}]} + + with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post: + await PassThroughEndpointHandler().process_input_messages( + data=data, + guardrail_to_apply=guardrail, + litellm_logging_obj=None, + ) + mock_post.assert_not_called() + + assert _recorded_guardrail_entries(data) == [] + assert "headroom" not in _applied_guardrails(data) @pytest.mark.asyncio @@ -884,6 +945,7 @@ async def test_apply_guardrail_transport_error_fail_open_forwards_uncompressed() texts=["hello"], structured_messages=ORIGINAL_MESSAGES, ) + request_data = {"model": "gpt-4o"} with patch.object( guardrail.async_handler, @@ -893,12 +955,18 @@ async def test_apply_guardrail_transport_error_fail_open_forwards_uncompressed() ): result = await guardrail.apply_guardrail( inputs=inputs, - request_data={}, + request_data=request_data, input_type="request", ) assert result["structured_messages"] == ORIGINAL_MESSAGES + entries = _recorded_guardrail_entries(request_data) + assert len(entries) == 1 + assert entries[0]["guardrail_name"] == "headroom" + assert entries[0]["guardrail_status"] == "guardrail_failed_to_respond" + assert "headroom" in _applied_guardrails(request_data) + @pytest.mark.asyncio async def test_apply_guardrail_http_error_fail_open_forwards_uncompressed(): diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py index e84e9b74201..bf904dbe394 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py @@ -4,7 +4,10 @@ import pytest import litellm from litellm.caching import DualCache -from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + log_guardrail_information, +) from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation from litellm.llms.base_llm.guardrail_translation.utils import ( effective_skip_system_message_for_guardrail, @@ -1490,3 +1493,109 @@ class TestStreamingTransform: # None holdback treated as 0: full text emitted, no crash. assert "".join(_delta_text(i) for i in out) == "ABCDEF" + + +def _applied_guardrails(data: dict) -> list: + for key in ("metadata", "litellm_metadata"): + meta = data.get(key) + if isinstance(meta, dict) and isinstance(meta.get("applied_guardrails"), list): + return meta["applied_guardrails"] + return [] + + +class _TextsOnlyTranslation(BaseTranslation): + """Mimics a passthrough handler: hands the guardrail only `texts`, never + structured_messages, so a structured_messages-based guardrail no-ops.""" + + async def process_input_messages(self, data, guardrail_to_apply, litellm_logging_obj=None): # type: ignore[override] + await guardrail_to_apply.apply_guardrail( + inputs={"texts": ["payload"]}, + request_data=data, + input_type="request", + logging_obj=litellm_logging_obj, + ) + return data + + async def process_output_response( # type: ignore[override] + self, + response, + guardrail_to_apply, + litellm_logging_obj=None, + user_api_key_dict=None, + request_data=None, + ): + return response + + +class _SelfLoggingGuardrail(CustomGuardrail): + records_own_guardrail_information = True + + def __init__(self, *, self_add: bool): + super().__init__(guardrail_name="self-logging") + self._self_add = self_add + + def should_run_guardrail(self, data, event_type): # type: ignore[override] + return True + + @log_guardrail_information + async def apply_guardrail(self, inputs, request_data, input_type, **kwargs): + if self._self_add: + from litellm.proxy.common_utils.callback_utils import ( + add_guardrail_to_applied_guardrails_header, + ) + + add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name) + return inputs + + +class _AutoLoggingGuardrail(CustomGuardrail): + def __init__(self): + super().__init__(guardrail_name="auto-logging") + + def should_run_guardrail(self, data, event_type): # type: ignore[override] + return True + + @log_guardrail_information + async def apply_guardrail(self, inputs, request_data, input_type, **kwargs): + return inputs + + +class TestAppliedGuardrailsReflectsExecution: + """The unified hook must not auto-mark a self-logging guardrail + (records_own_guardrail_information) as applied; such a guardrail owns that + decision and marks itself only when it actually ran (LIT-4650). Ordinary + guardrails are still auto-marked by the hook after dispatch.""" + + @staticmethod + def _data(guardrail): + return { + "guardrail_to_apply": guardrail, + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hello world"}], + } + + async def _run(self, guardrail): + unified_module.endpoint_guardrail_translation_mappings = {CallTypes.pass_through: _TextsOnlyTranslation} + data = self._data(guardrail) + await UnifiedLLMGuardrails().async_pre_call_hook( + user_api_key_dict=None, + cache=DualCache(), + data=data, + call_type=CallTypes.pass_through.value, + ) + return data + + @pytest.mark.asyncio + async def test_self_logging_guardrail_is_not_auto_marked_applied(self): + data = await self._run(_SelfLoggingGuardrail(self_add=False)) + assert "self-logging" not in _applied_guardrails(data) + + @pytest.mark.asyncio + async def test_self_logging_guardrail_that_self_marks_is_applied(self): + data = await self._run(_SelfLoggingGuardrail(self_add=True)) + assert "self-logging" in _applied_guardrails(data) + + @pytest.mark.asyncio + async def test_ordinary_guardrail_is_auto_marked_applied(self): + data = await self._run(_AutoLoggingGuardrail()) + assert "auto-logging" in _applied_guardrails(data) From 6ff88ba5e9e3778e493a3a4dc47fb1b88e1a1c7c Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 24 Jul 2026 16:46:49 -0700 Subject: [PATCH 8/9] fix(responses): strip include_usage from stream_options instead of dropping the param --- .../transformation.py | 7 +- litellm/responses/utils.py | 25 ++++++- litellm/types/llms/openai.py | 5 ++ ...responses_transformation_transformation.py | 74 +++++++++++++++++++ .../test_responses_api_request_body.py | 29 +++++++- 5 files changed, 130 insertions(+), 10 deletions(-) diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index aecb2552b53..3e50fe66039 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -35,6 +35,7 @@ from litellm.responses.sse_output_recovery import ( record_output_item_chunk, record_output_text_chunk, ) +from litellm.responses.utils import normalize_responses_api_stream_options from litellm.types.llms.openai import ( ChatCompletionAnnotation, ChatCompletionReasoningItem, @@ -320,6 +321,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): responses_api_request["tool_choice"] = ( # type: ignore[assignment] self._normalize_tool_choice_for_responses_api(value) ) + elif key == "stream_options": + stream_options = normalize_responses_api_stream_options(value) + if stream_options is not None: + responses_api_request["stream_options"] = stream_options elif key in ResponsesAPIOptionalRequestParams.__annotations__.keys(): responses_api_request[key] = value # type: ignore elif key == "previous_response_id": @@ -360,8 +365,6 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): continue if key == "instructions" and instructions: request_data["instructions"] = instructions - elif key == "stream_options" and isinstance(value, dict): - request_data["stream_options"] = value.get("include_obfuscation") elif key == "user" and isinstance(value, str): # OpenAI API requires user param to be max 64 chars - truncate if longer if len(value) <= 64: diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py index 12c890ec91d..429ddeef36a 100644 --- a/litellm/responses/utils.py +++ b/litellm/responses/utils.py @@ -5,6 +5,7 @@ from typing import ( Dict, Iterable, List, + Mapping, Optional, Type, Union, @@ -24,6 +25,7 @@ from litellm.types.llms.openai import ( ResponseInputParam, ResponsesAPIOptionalRequestParams, ResponsesAPIResponse, + ResponsesAPIStreamOptions, ResponseText, ) from litellm.types.responses.main import DecodedResponseId @@ -35,6 +37,17 @@ from litellm.types.utils import ( ) +def normalize_responses_api_stream_options( + stream_options: object, +) -> ResponsesAPIStreamOptions | None: + if not isinstance(stream_options, Mapping): + return None + include_obfuscation = stream_options.get("include_obfuscation") + if not isinstance(include_obfuscation, bool): + return None + return ResponsesAPIStreamOptions(include_obfuscation=include_obfuscation) + + class ResponsesAPIRequestUtils: """Helper utils for constructing ResponseAPI requests""" @@ -156,15 +169,19 @@ class ResponsesAPIRequestUtils: drop_params=should_drop_params, ) + stream_options = normalize_responses_api_stream_options(mapped_params.get("stream_options")) + params_with_normalized_stream_options = { + **{key: value for key, value in mapped_params.items() if key != "stream_options"}, + **({} if stream_options is None else {"stream_options": stream_options}), + } + # add any allowed_openai_params to the mapped_params - mapped_params = _apply_openai_param_overrides( - optional_params=mapped_params, + return _apply_openai_param_overrides( + optional_params=params_with_normalized_stream_options, non_default_params=non_default_params, allowed_openai_params=allowed_openai_params or [], ) - return mapped_params - @staticmethod def get_requested_response_api_optional_param( params: Dict[str, Any], diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 2263d53182a..314bb653196 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -1145,6 +1145,10 @@ class ContextManagementEntry(TypedDict, total=False): """Token threshold at which compaction is triggered for this entry. Minimum 1000.""" +class ResponsesAPIStreamOptions(TypedDict, total=False): + include_obfuscation: bool + + class ResponsesAPIOptionalRequestParams(TypedDict, total=False): """TypedDict for Optional parameters supported by the responses API.""" @@ -1171,6 +1175,7 @@ class ResponsesAPIOptionalRequestParams(TypedDict, total=False): max_tool_calls: Optional[int] prompt_cache_key: Optional[str] prompt_cache_retention: Optional[str] + stream_options: Optional[ResponsesAPIStreamOptions] top_logprobs: Optional[int] partial_images: Optional[int] # Number of partial images to generate (1-3) for streaming image generation context_management: Optional[List[ContextManagementEntry]] diff --git a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py index 6a1de0586dd..6907e4d0d02 100644 --- a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py +++ b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py @@ -2853,3 +2853,77 @@ def test_streaming_function_call_tool_id_for_degenerate_call_id(): assert stream_tool_id("fc_unique_abc123", "call_0") == "fc_unique_abc123" assert stream_tool_id("fc_2", "call_tokyo") == "call_tokyo" + + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "stream_options,expected_wire_stream_options", + [ + ({"include_usage": True, "include_obfuscation": False}, {"include_obfuscation": False}), + ({"include_usage": True}, None), + ], +) +async def test_acompletion_bridge_normalizes_stream_options_on_the_wire( + stream_options, expected_wire_stream_options +): + """include_usage must be stripped from the /v1/responses body; include_obfuscation must survive as a dict.""" + from unittest.mock import AsyncMock + + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + + responses_payload = { + "id": "resp_bridge_stream_options", + "object": "response", + "created_at": 1734366691, + "status": "completed", + "model": "gpt-5.5", + "output": [ + { + "type": "message", + "id": "msg_1", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "hi", "annotations": []}], + } + ], + "parallel_tool_calls": True, + "usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2}, + "error": None, + "incomplete_details": None, + "instructions": None, + "metadata": None, + "temperature": None, + "tool_choice": "auto", + "tools": [], + "top_p": None, + "max_output_tokens": None, + "previous_response_id": None, + "reasoning": None, + "truncation": None, + "user": None, + } + + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.text = json.dumps(responses_payload) + mock_response.headers = httpx.Headers({}) + mock_response.json.return_value = responses_payload + + with patch.object(AsyncHTTPHandler, "post", new_callable=AsyncMock) as mock_post: + mock_post.return_value = mock_response + + await litellm.acompletion( + model="openai/responses/gpt-5.5", + messages=[{"role": "user", "content": "hi"}], + api_key="fake-api-key", + stream_options=stream_options, + ) + + mock_post.assert_called_once() + post_kwargs = mock_post.call_args.kwargs + request_body = post_kwargs["json"] if "json" in post_kwargs else json.loads(post_kwargs["data"]) + if expected_wire_stream_options is None: + assert "stream_options" not in request_body + else: + assert request_body["stream_options"] == expected_wire_stream_options diff --git a/tests/test_litellm/responses/test_responses_api_request_body.py b/tests/test_litellm/responses/test_responses_api_request_body.py index 2922c9738aa..83b9c34636e 100644 --- a/tests/test_litellm/responses/test_responses_api_request_body.py +++ b/tests/test_litellm/responses/test_responses_api_request_body.py @@ -200,10 +200,7 @@ async def test_aresponses_azure_shell_tool_400_maps_to_bad_request_error(): @pytest.mark.asyncio async def test_aresponses_drops_stream_options(): - """ - stream_options is a Chat Completions param; the Responses API rejects it with - "Unknown parameter: 'stream_options.include_usage'". It must never reach the wire. - """ + """The Responses API rejects include_usage, so include_usage-only stream_options must never reach the wire.""" with patch( "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", new_callable=AsyncMock, @@ -225,6 +222,30 @@ async def test_aresponses_drops_stream_options(): assert "stream_options" not in request_body +@pytest.mark.asyncio +async def test_aresponses_keeps_include_obfuscation_in_stream_options(): + """include_obfuscation is a valid Responses API stream option and must survive the include_usage strip.""" + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = MockResponse( + _minimal_responses_api_payload("resp_stream_options_obfuscation", "gpt-5.5"), 200 + ) + + await litellm.aresponses( + model="openai/gpt-5.5", + api_key="fake-api-key", + input="hi", + stream_options={"include_usage": True, "include_obfuscation": False}, + ) + + mock_post.assert_called_once() + post_kwargs = mock_post.call_args.kwargs + request_body = post_kwargs["json"] if "json" in post_kwargs else json.loads(post_kwargs["data"]) + assert request_body["stream_options"] == {"include_obfuscation": False} + + @pytest.mark.asyncio async def test_aresponses_request_level_drop_params_drops_bedrock_mantle_service_tier( monkeypatch, From 76b0b10908d2aaa779e2452e13968390b5e806f3 Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Fri, 24 Jul 2026 17:13:11 -0700 Subject: [PATCH 9/9] fix(guardrails): add /v1/messages support for Straiker plugin (#34548) * fix(guardrails): add /v1/messages support for Straiker plugin - Pass prepared response data to Anthropic Messages streaming post-call hooks (litellm/llms/anthropic/chat/guardrail_translation/handler.py) - Normalize Straiker request, tool, finish-reason, and mode fields across Chat Completions, Messages, and Responses APIs * fix(guardrails): gate cross-surface message resolution and cover streaming request data Resolve request messages only for surfaces that have a mapped translation handler. The unguarded fallback tried every registered handler in turn, which raised AttributeError out of the guardrail's error handling on list-shaped `input` bodies, and synthesized a chat message that was never sent for bodies it happened to parse. Prepare request data on the mid-stream Anthropic branch as well, matching the terminal branch and the OpenAI handler, so guardrails that scan before end-of-stream still receive identity metadata. Read usage from Anthropic dict responses so non-streaming /v1/messages reports token counts instead of null. Add regression coverage for the streaming request data on both the terminal and mid-stream branches; reverting either now fails. --------- Co-authored-by: cs-mehta --- .../chat/guardrail_translation/handler.py | 16 +- .../guardrail_hooks/straiker/straiker.py | 123 ++++++- .../guardrails/guardrail_hooks/straiker.py | 8 +- .../test_anthropic_guardrail_handler.py | 92 +++++ .../guardrail_hooks/test_straiker.py | 340 +++++++++++++++++- 5 files changed, 556 insertions(+), 23 deletions(-) diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index 7000c20d9c4..90f735707bf 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -600,9 +600,15 @@ class AnthropicMessagesHandler(BaseTranslation): guardrail_inputs["tool_calls"] = tool_calls_list try: + prepared_request_data = self._prepare_request_data( + request_data, + model_response, + user_api_key_dict, + key="response", + ) _guardrailed_inputs = await guardrail_to_apply.apply_guardrail( inputs=guardrail_inputs, - request_data=request_data if request_data is not None else {}, + request_data=prepared_request_data, input_type="response", logging_obj=litellm_logging_obj, ) @@ -618,9 +624,15 @@ class AnthropicMessagesHandler(BaseTranslation): string_so_far = self.get_streaming_string_so_far(responses_so_far) try: + prepared_request_data = self._prepare_request_data( + request_data, + responses_so_far, + user_api_key_dict, + key="responses", + ) _guardrailed_inputs = await guardrail_to_apply.apply_guardrail( inputs={"texts": [string_so_far]}, - request_data=request_data if request_data is not None else {}, + request_data=prepared_request_data, input_type="response", logging_obj=litellm_logging_obj, ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py b/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py index 5c9f93fc2cd..717f5b6c5fe 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py +++ b/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py @@ -8,7 +8,7 @@ from typing import TYPE_CHECKING, Any, Literal, NoReturn from urllib.parse import urlsplit import httpx -from pydantic import ValidationError +from pydantic import BaseModel, TypeAdapter, ValidationError from litellm._logging import verbose_proxy_logger from litellm._version import version as litellm_version @@ -24,11 +24,12 @@ from litellm.integrations.custom_guardrail import ( log_guardrail_information, ) from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider +from litellm.litellm_core_utils.prompt_templates.factory import resolve_structured_messages from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) -from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.guardrails import GuardrailEventHooks, Mode from litellm.types.proxy.guardrails.guardrail_hooks.straiker import ( STRAIKER_WEBHOOK_SCHEMA_VERSION, StraikerGuardrailConfigModel, @@ -42,7 +43,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.straiker import ( StraikerWebhookStream, StraikerWebhookUsage, ) -from litellm.types.utils import GenericGuardrailAPIInputs, Usage +from litellm.types.utils import GenericGuardrailAPIInputs if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -57,6 +58,7 @@ RETRY_STATUS = frozenset({408, 429, 500, 502, 503, 504}) UNREACHABLE_STATUS = frozenset({502, 503, 504}) _APPLICATION_METADATA_KEYS = frozenset({"agent_id", "app_name"}) _OPAQUE_METADATA_SCALAR_TYPES = (str, int, float, bool) +_JSON_DICT_ADAPTER = TypeAdapter(dict[str, object]) @dataclass(frozen=True, slots=True) @@ -137,6 +139,44 @@ def _resolve_destination(request_data: dict) -> str | None: return None +def _route_has_translation(request_data: dict) -> bool: + from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route + from litellm.llms import load_guardrail_translation_mappings + + route = _as_dict(request_data.get("litellm_metadata")).get("user_api_key_request_route") + if not isinstance(route, str) or not route: + return False + mappings = load_guardrail_translation_mappings() + return any(call_type in mappings for call_type in get_call_types_for_route(route) or ()) + + +def _request_structured_messages(request_data: dict) -> list[dict[str, Any]] | None: + messages = request_data.get("messages") + if messages: + return messages if isinstance(messages, list) else None + if not _route_has_translation(request_data): + return None + return resolve_structured_messages(messages=None, request_kwargs=request_data) + + +def _hook_name(value: object) -> str: + return value.value if isinstance(value, GuardrailEventHooks) else str(value) + + +def _configured_modes(event_hook: object) -> list[str] | None: + if isinstance(event_hook, list): + names = [_hook_name(v) for v in event_hook] + elif isinstance(event_hook, (str, GuardrailEventHooks)): + names = [_hook_name(event_hook)] + elif isinstance(event_hook, Mode): + default = event_hook.default if isinstance(event_hook.default, list) else [event_hook.default] + tags = [v for value in event_hook.tags.values() for v in (value if isinstance(value, list) else [value])] + names = [_hook_name(v) for v in (*default, *tags) if v is not None] + else: + return None + return list(dict.fromkeys(names)) or None + + def _resolve_call_surface(logging_obj: LiteLLMLoggingObj | None, request_data: dict) -> str: call_type = ( (getattr(logging_obj, "call_type", None) if logging_obj is not None else None) @@ -146,23 +186,76 @@ def _resolve_call_surface(logging_obj: LiteLLMLoggingObj | None, request_data: d return call_type if isinstance(call_type, str) and call_type else "unknown" +def _jsonable_dict(value: object) -> dict[str, object] | None: + if isinstance(value, BaseModel): + return _JSON_DICT_ADAPTER.validate_python(value.model_dump(mode="json", exclude_none=True)) + if isinstance(value, dict): + return _JSON_DICT_ADAPTER.validate_python(value) + return None + + +def _opaque_dict_list(value: object) -> list[dict[str, object]] | None: + if not isinstance(value, list): + return None + items = tuple(plain for item in value if (plain := _jsonable_dict(item)) is not None) + return list(items) if items else None + + +def _choice_terminal_reason(choice: object) -> str | None: + if isinstance(choice, dict): + return _as_optional_str(choice.get("finish_reason")) or _as_optional_str(choice.get("stop_reason")) + return _as_optional_str(getattr(choice, "finish_reason", None)) or _as_optional_str( + getattr(choice, "stop_reason", None) + ) + + def _response_finish_reason(response: Any) -> str | None: + if response is None: + return None + if isinstance(response, dict): + top = _as_optional_str(response.get("finish_reason")) or _as_optional_str(response.get("stop_reason")) + if top: + return top + choices = response.get("choices") + if not isinstance(choices, list): + return None + for choice in choices: + reason = _choice_terminal_reason(choice) + if reason: + return reason + return None + + top = _as_optional_str(getattr(response, "finish_reason", None)) or _as_optional_str( + getattr(response, "stop_reason", None) + ) + if top: + return top choices = getattr(response, "choices", None) if not isinstance(choices, list): return None for choice in choices: - reason = getattr(choice, "finish_reason", None) - if isinstance(reason, str) and reason: + reason = _choice_terminal_reason(choice) + if reason: return reason return None +def _as_optional_int(value: object) -> int | None: + return value if isinstance(value, int) and not isinstance(value, bool) else None + + +def _usage_token_count(usage: object, openai_key: str, anthropic_key: str) -> int | None: + get = usage.get if isinstance(usage, dict) else lambda key: getattr(usage, key, None) + openai_count = _as_optional_int(get(openai_key)) + return openai_count if openai_count is not None else _as_optional_int(get(anthropic_key)) + + def _build_usage(response: object) -> StraikerWebhookUsage | None: - usage = getattr(response, "usage", None) - if not isinstance(usage, Usage): + usage = response.get("usage") if isinstance(response, dict) else getattr(response, "usage", None) + if usage is None: return None - input_tokens = usage.prompt_tokens - output_tokens = usage.completion_tokens + input_tokens = _usage_token_count(usage, "prompt_tokens", "input_tokens") + output_tokens = _usage_token_count(usage, "completion_tokens", "output_tokens") if input_tokens is None and output_tokens is None: return None return StraikerWebhookUsage(input_tokens=input_tokens, output_tokens=output_tokens) @@ -234,6 +327,8 @@ class StraikerGuardrail(CustomGuardrail): kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) super().__init__(**kwargs) + self.configured_modes = _configured_modes(self.event_hook) + def _webhook_url(self) -> str: return f"{self.api_base}{WEBHOOK_PATH}" @@ -263,6 +358,7 @@ class StraikerGuardrail(CustomGuardrail): ) -> StraikerWebhookContext: return StraikerWebhookContext( call_surface=_resolve_call_surface(logging_obj, request_data), + mode=self.configured_modes, model=model, model_provider=_resolve_provider(request_data, model), destination=_resolve_destination(request_data), @@ -287,9 +383,9 @@ class StraikerGuardrail(CustomGuardrail): content = StraikerWebhookContent( texts=list(inputs.get("texts") or []), images=list(inputs.get("images") or []), - structured_messages=inputs.get("structured_messages"), - tools=inputs.get("tools"), - tool_calls=inputs.get("tool_calls"), + structured_messages=_opaque_dict_list(inputs.get("structured_messages")), + tools=_opaque_dict_list(inputs.get("tools")), + tool_calls=_opaque_dict_list(inputs.get("tool_calls")), ) if input_type == "request": @@ -305,9 +401,8 @@ class StraikerGuardrail(CustomGuardrail): response_obj = request_data.get("response") content.finish_reason = _response_finish_reason(response_obj) - original_messages = request_data.get("messages") request_content = StraikerWebhookContent( - structured_messages=original_messages if isinstance(original_messages, list) else None, + structured_messages=_opaque_dict_list(_request_structured_messages(request_data)), ) phase: Literal["none", "assembled"] = "assembled" if _is_streamed_request(request_data) else "none" event = StraikerWebhookEvent(type="post_call", id=event_id, stream=StraikerWebhookStream(phase=phase)) diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/straiker.py b/litellm/types/proxy/guardrails/guardrail_hooks/straiker.py index b4375237917..0e816985cb0 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/straiker.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/straiker.py @@ -4,9 +4,6 @@ from typing import Literal from pydantic import BaseModel, ConfigDict, Field -from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolCallChunk -from litellm.types.utils import ChatCompletionMessageToolCall - from .base import GuardrailConfigModel StraikerWebhookEventType = Literal["pre_call", "post_call"] @@ -32,9 +29,9 @@ class StraikerWebhookContent(BaseModel): texts: list[str] = Field(default_factory=list) images: list[str] = Field(default_factory=list) - structured_messages: list[AllMessageValues] | None = None + structured_messages: list[dict[str, object]] | None = None tools: list[dict[str, object]] | None = None - tool_calls: list[ChatCompletionToolCallChunk] | list[ChatCompletionMessageToolCall] | None = None + tool_calls: list[dict[str, object]] | None = None finish_reason: str | None = None @@ -45,6 +42,7 @@ class StraikerWebhookUsage(BaseModel): class StraikerWebhookContext(BaseModel): call_surface: str + mode: list[str] | None = None model: str | None = None model_provider: str | None = None destination: str | None = None diff --git a/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py b/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py index 9cd1fbb59a6..48acdd348e9 100644 --- a/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py +++ b/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py @@ -57,6 +57,98 @@ class MockDynamicGuardrail(CustomGuardrail): return inputs +class MockRecordingGuardrail(CustomGuardrail): + """Mock guardrail that records the request_data it was handed.""" + + def __init__(self, guardrail_name: str): + super().__init__(guardrail_name=guardrail_name) + self.request_data: Optional[dict] = None + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional[Any] = None, + ) -> GenericGuardrailAPIInputs: + self.request_data = request_data + return inputs + + +class TestAnthropicMessagesHandlerStreamingRequestData: + """Post-call guardrails on streaming /v1/messages receive the response and identity metadata""" + + @pytest.mark.asyncio + async def test_terminal_chunk_passes_assembled_response_and_metadata(self): + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.utils import Choices, Message, ModelResponse + + handler = AnthropicMessagesHandler() + guardrail = MockRecordingGuardrail(guardrail_name="test") + mock_response = ModelResponse( + id="msg_123", + created=1234567890, + model="claude-sonnet-4-5", + object="chat.completion", + choices=[ + Choices( + finish_reason="stop", + index=0, + message=Message(content="Hello world", role="assistant"), + ) + ], + ) + + with ( + patch.object(handler, "_check_streaming_has_ended", return_value=True), + patch( + "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler._build_complete_streaming_response", + return_value=mock_response, + ), + ): + await handler.process_output_streaming_response( + responses_so_far=[b"data: some chunk"], + guardrail_to_apply=guardrail, + litellm_logging_obj=MagicMock(), + user_api_key_dict=UserAPIKeyAuth(user_id="u-1", team_id="t-1"), + request_data={"model": "claude-sonnet-4-5"}, + ) + + assert guardrail.request_data is not None + assert guardrail.request_data["response"] is mock_response + assert ( + guardrail.request_data["litellm_metadata"]["user_api_key_user_id"] == "u-1" + ) + + @pytest.mark.asyncio + async def test_mid_stream_chunk_passes_responses_so_far_and_metadata(self): + from litellm.proxy._types import UserAPIKeyAuth + + handler = AnthropicMessagesHandler() + guardrail = MockRecordingGuardrail(guardrail_name="test") + responses_so_far = [b"data: some chunk"] + + with ( + patch.object(handler, "_check_streaming_has_ended", return_value=False), + patch.object( + handler, "get_streaming_string_so_far", return_value="partial text" + ), + ): + await handler.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail, + litellm_logging_obj=MagicMock(), + user_api_key_dict=UserAPIKeyAuth(user_id="u-1", team_id="t-1"), + request_data={"model": "claude-sonnet-4-5"}, + ) + + assert guardrail.request_data is not None + assert guardrail.request_data["responses"] is responses_so_far + assert ( + guardrail.request_data["litellm_metadata"]["user_api_key_user_id"] == "u-1" + ) + + class TestAnthropicMessagesHandlerStreamingOutputProcessing: """Test streaming output processing functionality""" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_straiker.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_straiker.py index ca57118ee9d..36a2e205ea7 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_straiker.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_straiker.py @@ -1,4 +1,5 @@ import json +from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock import httpx @@ -8,6 +9,9 @@ from litellm.exceptions import GuardrailRaisedException, ModifyResponseException from litellm.proxy.guardrails.guardrail_hooks.straiker import initialize_guardrail from litellm.proxy.guardrails.guardrail_hooks.straiker.straiker import ( StraikerGuardrail, + _build_usage, + _request_structured_messages, + _response_finish_reason, ) from litellm.proxy.guardrails.guardrail_registry import ( guardrail_class_registry, @@ -17,7 +21,14 @@ from litellm.types.proxy.guardrails.guardrail_hooks.straiker import ( StraikerGuardrailConfigModel, StraikerGuardrailConfigModelOptionalParams, ) -from litellm.types.utils import Choices, Message, ModelResponse, Usage +from litellm.types.utils import ( + ChatCompletionMessageToolCall, + Choices, + Function, + Message, + ModelResponse, + Usage, +) def _mock_response(action: str, turn_id: str = "turn-1", schema_version: str = "1", **extra) -> MagicMock: @@ -208,7 +219,9 @@ async def test_request_envelope_transport_and_shape(): "metadata": {"user_api_key_alias": "team-key", "agent_id": "chatbot-app", "app_name": "Chatbot"}, } - out = await g.apply_guardrail(inputs=inputs, request_data=request_data, input_type="request", logging_obj=_logging_obj()) + out = await g.apply_guardrail( + inputs=inputs, request_data=request_data, input_type="request", logging_obj=_logging_obj() + ) assert out is inputs url = g.async_handler.post.call_args.args[0] @@ -232,6 +245,35 @@ async def test_request_envelope_transport_and_shape(): assert "metadata" not in payload +@pytest.mark.asyncio +async def test_request_envelope_ignores_unsupported_opaque_items(): + g = _make_guardrail() + g.async_handler.post.return_value = _mock_response("NONE") + + await g.apply_guardrail( + inputs={ + "texts": ["hello"], + "tools": [ + object(), + { + "type": "function", + "function": {"name": "get_weather", "parameters": {"type": "object"}}, + }, + ], + }, + request_data={"model": "m", "messages": [{"role": "user", "content": "hello"}]}, + input_type="request", + logging_obj=_logging_obj(), + ) + + assert _posted_payload(g)["request"]["tools"] == [ + { + "type": "function", + "function": {"name": "get_weather", "parameters": {"type": "object"}}, + } + ] + + @pytest.mark.asyncio async def test_webhook_metadata_session_id_and_opaque_passthrough(): g = _make_guardrail() @@ -331,6 +373,64 @@ async def test_context_session_id_from_request_metadata(): assert "metadata" not in payload +@pytest.mark.asyncio +async def test_context_mode_from_string_event_hook(): + g = _make_guardrail(event_hook="pre_call") + g.async_handler.post.return_value = _mock_response("NONE") + await g.apply_guardrail( + inputs={"texts": ["x"]}, + request_data={"model": "m"}, + input_type="request", + logging_obj=_logging_obj(), + ) + assert _posted_payload(g)["context"]["mode"] == ["pre_call"] + + +@pytest.mark.asyncio +async def test_context_mode_from_list_event_hook(): + from litellm.types.guardrails import GuardrailEventHooks + + g = _make_guardrail(event_hook=[GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call]) + g.async_handler.post.return_value = _mock_response("NONE") + await g.apply_guardrail( + inputs={"texts": ["x"]}, + request_data={"model": "m"}, + input_type="request", + logging_obj=_logging_obj(), + ) + assert _posted_payload(g)["context"]["mode"] == ["pre_call", "post_call"] + + +@pytest.mark.asyncio +async def test_context_mode_from_tagged_mode_is_flattened_and_deduped(): + from litellm.types.guardrails import Mode + + g = _make_guardrail( + event_hook=Mode(tags={"team-a": "pre_call", "team-b": ["post_call", "pre_call"]}, default="post_call") + ) + g.async_handler.post.return_value = _mock_response("NONE") + await g.apply_guardrail( + inputs={"texts": ["x"]}, + request_data={"model": "m"}, + input_type="request", + logging_obj=_logging_obj(), + ) + assert _posted_payload(g)["context"]["mode"] == ["post_call", "pre_call"] + + +@pytest.mark.asyncio +async def test_context_mode_omitted_when_event_hook_absent(): + g = _make_guardrail(event_hook=None) + g.async_handler.post.return_value = _mock_response("NONE") + await g.apply_guardrail( + inputs={"texts": ["x"]}, + request_data={"model": "m"}, + input_type="request", + logging_obj=_logging_obj(), + ) + assert "mode" not in _posted_payload(g)["context"] + + @pytest.mark.asyncio async def test_identity_key_and_team_coalesce_alias_over_id(): g = _make_guardrail() @@ -426,6 +526,7 @@ async def test_application_source_from_agent_id(): ) assert _posted_payload(g)["application"] == {"source": "analytics-app", "name": "Analytics"} + @pytest.mark.asyncio async def test_request_block_raises_guardrail_exception_with_reason(): g = _make_guardrail() @@ -560,6 +661,34 @@ async def test_response_envelope_and_block_replaces_response(): assert payload["request"]["structured_messages"] == [{"role": "user", "content": "original prompt"}] +@pytest.mark.asyncio +async def test_post_call_resolves_request_from_responses_input_when_messages_absent(): + g = _make_guardrail() + g.async_handler.post.return_value = _mock_response("NONE") + response = ModelResponse( + choices=[Choices(finish_reason="stop", index=0, message=Message(content="answer", role="assistant"))], + model="gpt-4o-mini", + ) + request_data = { + "model": "gpt-4o-mini", + "input": "responses-surface prompt", + "response": response, + "litellm_metadata": {"user_api_key_request_route": "/v1/responses"}, + } + + await g.apply_guardrail( + inputs={"texts": ["answer"], "model": "gpt-4o-mini"}, + request_data=request_data, + input_type="response", + logging_obj=_logging_obj(), + ) + + payload = _posted_payload(g) + assert payload["event"]["type"] == "post_call" + messages = payload["request"]["structured_messages"] + assert any(m.get("content") == "responses-surface prompt" for m in messages) + + @pytest.mark.asyncio async def test_post_call_fail_closed_raises_modify_response_exception(): g = _make_guardrail(unreachable_fallback="fail_closed") @@ -731,3 +860,210 @@ async def test_unreachable_http_status_fail_closed_blocks(): await g.apply_guardrail( inputs={"texts": ["x"]}, request_data={"model": "m"}, input_type="request", logging_obj=_logging_obj() ) + + +@pytest.mark.asyncio +async def test_post_call_preserves_anthropic_tool_blocks_in_request_messages(): + g = _make_guardrail() + g.async_handler.post.return_value = _mock_response("NONE") + anthropic_messages = [ + {"role": "user", "content": "What's the weather in Paris?"}, + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "toolu_1", + "name": "get_weather", + "input": {"city": "Paris"}, + } + ], + }, + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "toolu_1", + "content": "18C, cloudy", + } + ], + }, + ] + response = { + "id": "msg_1", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "Mild and cloudy."}], + "stop_reason": "end_turn", + "model": "claude-sonnet-5", + } + await g.apply_guardrail( + inputs={"texts": ["Mild and cloudy."], "model": "claude-sonnet-5"}, + request_data={ + "model": "claude-sonnet-5", + "messages": anthropic_messages, + "response": response, + }, + input_type="response", + logging_obj=_logging_obj(), + ) + payload = _posted_payload(g) + assert payload["request"]["structured_messages"] == anthropic_messages + assert payload["response"]["finish_reason"] == "end_turn" + + +@pytest.mark.asyncio +async def test_pre_call_preserves_anthropic_tool_blocks_in_structured_messages(): + g = _make_guardrail() + g.async_handler.post.return_value = _mock_response("NONE") + anthropic_messages = [ + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "toolu_1", + "name": "get_weather", + "input": {"city": "Paris"}, + } + ], + }, + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "toolu_1", + "content": "18C", + } + ], + }, + ] + await g.apply_guardrail( + inputs={"structured_messages": anthropic_messages, "model": "claude-sonnet-5"}, + request_data={"model": "claude-sonnet-5", "messages": anthropic_messages}, + input_type="request", + logging_obj=_logging_obj(), + ) + assert _posted_payload(g)["request"]["structured_messages"] == anthropic_messages + + +@pytest.mark.asyncio +async def test_response_finish_reason_from_openai_choices_still_works(): + g = _make_guardrail() + g.async_handler.post.return_value = _mock_response("NONE") + response = ModelResponse( + choices=[Choices(finish_reason="tool_calls", index=0, message=Message(content=None, role="assistant"))], + model="gpt-4o-mini", + ) + await g.apply_guardrail( + inputs={ + "texts": [], + "tool_calls": [ + ChatCompletionMessageToolCall( + id="c1", + type="function", + function=Function(name="f", arguments="{}"), + ) + ], + }, + request_data={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}], "response": response}, + input_type="response", + logging_obj=_logging_obj(), + ) + payload = _posted_payload(g) + assert payload["response"]["finish_reason"] == "tool_calls" + assert payload["response"]["tool_calls"] == [ + {"id": "c1", "type": "function", "function": {"name": "f", "arguments": "{}"}} + ] + + +@pytest.mark.parametrize( + ("response", "expected"), + [ + (None, None), + ({"choices": "invalid"}, None), + ({"choices": [{"finish_reason": "length"}]}, "length"), + ({"choices": [{"stop_reason": "end_turn"}]}, "end_turn"), + ({"choices": [{}]}, None), + (SimpleNamespace(stop_reason="end_turn"), "end_turn"), + ], +) +def test_response_finish_reason_handles_supported_shapes(response, expected): + assert _response_finish_reason(response) == expected + + +@pytest.mark.parametrize( + "request_data", + [ + {"input": ["ssn 123-45-6789"], "litellm_metadata": {"user_api_key_request_route": "/vllm/v1/embeddings"}}, + {"input": [[1, 2, 3]], "litellm_metadata": {}}, + {"input": "confidential memo", "litellm_metadata": {}}, + {"input": "confidential memo"}, + ], +) +def test_request_messages_not_resolved_for_unmapped_surfaces(request_data): + """Bodies from surfaces without a translation handler yield no messages, and never raise.""" + assert _request_structured_messages(request_data) is None + + +@pytest.mark.parametrize( + ("request_data", "expected"), + [ + ( + {"messages": [{"role": "user", "content": "hi"}], "litellm_metadata": {}}, + [{"role": "user", "content": "hi"}], + ), + ( + { + "input": [{"role": "user", "content": "weather in Paris?"}], + "litellm_metadata": {"user_api_key_request_route": "/v1/responses"}, + }, + [{"role": "user", "content": "weather in Paris?"}], + ), + ], +) +def test_request_messages_resolved_for_mapped_surfaces(request_data, expected): + assert _request_structured_messages(request_data) == expected + + +@pytest.mark.parametrize( + ("response", "expected"), + [ + ({"usage": {"input_tokens": 10, "output_tokens": 5}}, (10, 5)), + ({"usage": {"prompt_tokens": 7, "completion_tokens": 3}}, (7, 3)), + (SimpleNamespace(usage=Usage(prompt_tokens=7, completion_tokens=3)), (7, 3)), + ({"usage": {"prompt_tokens": 0, "input_tokens": 99}}, (0, None)), + ({"usage": {}}, None), + ({}, None), + ], +) +def test_build_usage_handles_openai_and_anthropic_shapes(response, expected): + usage = _build_usage(response) + if expected is None: + assert usage is None + else: + assert (usage.input_tokens, usage.output_tokens) == expected + + +@pytest.mark.asyncio +async def test_anthropic_non_streaming_response_reports_usage(): + g = _make_guardrail() + g.async_handler.post.return_value = _mock_response("NONE") + await g.apply_guardrail( + inputs={"texts": ["hello"]}, + request_data={ + "model": "claude-sonnet-4-5", + "messages": [{"role": "user", "content": "hi"}], + "response": { + "stop_reason": "end_turn", + "usage": {"input_tokens": 10, "output_tokens": 5}, + }, + }, + input_type="response", + logging_obj=_logging_obj(), + ) + payload = _posted_payload(g) + assert payload["usage"] == {"input_tokens": 10, "output_tokens": 5} + assert payload["response"]["finish_reason"] == "end_turn"