diff --git a/litellm/utils.py b/litellm/utils.py index b7caf0edd7..4367ec789b 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -780,9 +780,9 @@ def function_setup( # noqa: PLR0915 coroutine_checker = get_coroutine_checker_fn() ## DYNAMIC CALLBACKS ## - dynamic_callbacks: Optional[List[Union[str, Callable, "CustomLogger"]]] = ( - kwargs.pop("callbacks", None) - ) + dynamic_callbacks: Optional[ + List[Union[str, Callable, "CustomLogger"]] + ] = kwargs.pop("callbacks", None) all_callbacks = get_dynamic_callbacks(dynamic_callbacks=dynamic_callbacks) if len(all_callbacks) > 0: @@ -1143,6 +1143,14 @@ def function_setup( # noqa: PLR0915 litellm_params: Dict[str, Any] = {"api_base": ""} if "metadata" in kwargs: litellm_params["metadata"] = kwargs["metadata"] + if "litellm_metadata" in kwargs and isinstance(kwargs["litellm_metadata"], dict): + litellm_params["litellm_metadata"] = kwargs["litellm_metadata"].copy() + # For endpoints like /v1/messages that use "litellm_metadata" instead + # of "metadata" (to avoid conflicting with provider API metadata fields), + # populate litellm_params["metadata"] so callbacks (e.g. Langfuse) that + # read API key info from litellm_params["metadata"] see the fields. + if not litellm_params.get("metadata"): + litellm_params["metadata"] = kwargs["litellm_metadata"].copy() logging_obj.update_environment_variables( model=model, @@ -1682,9 +1690,9 @@ def client(original_function): # noqa: PLR0915 exception=e, retry_policy=kwargs.get("retry_policy"), ) - kwargs["retry_policy"] = ( - reset_retry_policy() - ) # prevent infinite loops + kwargs[ + "retry_policy" + ] = reset_retry_policy() # prevent infinite loops litellm.num_retries = ( None # set retries to None to prevent infinite loops ) @@ -1731,9 +1739,9 @@ def client(original_function): # noqa: PLR0915 exception=e, retry_policy=kwargs.get("retry_policy"), ) - kwargs["retry_policy"] = ( - reset_retry_policy() - ) # prevent infinite loops + kwargs[ + "retry_policy" + ] = reset_retry_policy() # prevent infinite loops litellm.num_retries = ( None # set retries to None to prevent infinite loops ) @@ -3686,10 +3694,10 @@ def pre_process_non_default_params( if "response_format" in non_default_params: if provider_config is not None: - non_default_params["response_format"] = ( - provider_config.get_json_schema_from_pydantic_object( - response_format=non_default_params["response_format"] - ) + non_default_params[ + "response_format" + ] = provider_config.get_json_schema_from_pydantic_object( + response_format=non_default_params["response_format"] ) else: non_default_params["response_format"] = type_to_response_format_param( @@ -3818,16 +3826,16 @@ def pre_process_optional_params( True # so that main.py adds the function call to the prompt ) if "tools" in non_default_params: - optional_params["functions_unsupported_model"] = ( - non_default_params.pop("tools") - ) + optional_params[ + "functions_unsupported_model" + ] = non_default_params.pop("tools") non_default_params.pop( "tool_choice", None ) # causes ollama requests to hang elif "functions" in non_default_params: - optional_params["functions_unsupported_model"] = ( - non_default_params.pop("functions") - ) + optional_params[ + "functions_unsupported_model" + ] = non_default_params.pop("functions") elif ( litellm.add_function_to_prompt ): # if user opts to add it to prompt instead @@ -7428,9 +7436,9 @@ class ModelResponseIterator: if convert_to_delta is True: _stream_response = ModelResponseStream() _stream_response.choices[0].delta.content = model_response.choices[0].message.content # type: ignore - self.model_response: Union[ModelResponse, ModelResponseStream] = ( - _stream_response - ) + self.model_response: Union[ + ModelResponse, ModelResponseStream + ] = _stream_response else: self.model_response = model_response self.is_done = False @@ -7901,7 +7909,10 @@ class ProviderConfigManager: # Simple provider mappings (no model parameter needed) LlmProviders.DEEPSEEK: (lambda: litellm.DeepSeekChatConfig(), False), LlmProviders.GROQ: (lambda: litellm.GroqChatConfig(), False), - LlmProviders.BEDROCK_MANTLE: (lambda: litellm.BedrockMantleChatConfig(), False), + LlmProviders.BEDROCK_MANTLE: ( + lambda: litellm.BedrockMantleChatConfig(), + False, + ), LlmProviders.A2A: (lambda: litellm.A2AConfig(), False), LlmProviders.BYTEZ: (lambda: litellm.BytezChatConfig(), False), LlmProviders.DATABRICKS: (lambda: litellm.DatabricksConfig(), False), diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 28624ea8b2..f4aeb27b31 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -1,4 +1,3 @@ -import json import os import sys from unittest.mock import MagicMock, patch @@ -190,8 +189,13 @@ async def test_datadog_logger_not_shadowed_by_llm_obs(monkeypatch): # Regression check: we expect a distinct DataDogLogger, not the LLM Obs logger assert type(datadog_logger) is DataDogLogger - assert any(isinstance(cb, DataDogLLMObsLogger) for cb in logging_module._in_memory_loggers) - assert any(type(cb) is DataDogLogger for cb in logging_module._in_memory_loggers) + assert any( + isinstance(cb, DataDogLLMObsLogger) + for cb in logging_module._in_memory_loggers + ) + assert any( + type(cb) is DataDogLogger for cb in logging_module._in_memory_loggers + ) finally: logging_module._in_memory_loggers.clear() @@ -202,7 +206,9 @@ async def test_logfire_logger_accepts_env_vars_for_base_url(monkeypatch): # Required env vars for Logfire integration monkeypatch.setenv("LOGFIRE_TOKEN", "test-token") - monkeypatch.setenv("LOGFIRE_BASE_URL", "https://logfire-api-custom.pydantic.dev") # no trailing slash on purpose + monkeypatch.setenv( + "LOGFIRE_BASE_URL", "https://logfire-api-custom.pydantic.dev" + ) # no trailing slash on purpose # Import after env vars are set (important if module-level caching exists) from litellm.integrations.opentelemetry import OpenTelemetry # logger class @@ -221,7 +227,9 @@ async def test_logfire_logger_accepts_env_vars_for_base_url(monkeypatch): # Sanity: we got the right logger type and it is cached assert type(logger) is OpenTelemetry - assert any(type(cb) is OpenTelemetry for cb in logging_module._in_memory_loggers) + assert any( + type(cb) is OpenTelemetry for cb in logging_module._in_memory_loggers + ) # Core regression check: base URL env var should influence the exporter endpoint. # @@ -232,7 +240,9 @@ async def test_logfire_logger_accepts_env_vars_for_base_url(monkeypatch): or getattr(logger, "config", None) or getattr(logger, "_otel_config", None) ) - assert cfg is not None, "Expected OpenTelemetry logger to keep an otel config on the instance" + assert ( + cfg is not None + ), "Expected OpenTelemetry logger to keep an otel config on the instance" endpoint = getattr(cfg, "endpoint", None) or getattr(cfg, "otlp_endpoint", None) assert endpoint is not None, "Expected otel config to expose the OTLP endpoint" @@ -297,7 +307,7 @@ async def test_logging_non_streaming_request(): mock_response="Hello, world!", ) await asyncio.sleep(1) - + # Filter calls to only count the one with the expected input message "Hey" # Bridge models may make internal calls that also log, so we filter by the actual input calls_with_expected_input = [] @@ -307,13 +317,13 @@ async def test_logging_non_streaming_request(): first_message_content = messages[0].get("content") if first_message_content == "Hey": calls_with_expected_input.append(call) - + # Assert that we have exactly one call with the expected input assert len(calls_with_expected_input) == 1, ( f"Expected 1 call with input 'Hey', but got {len(calls_with_expected_input)}. " f"Total calls: {mock_async_log_success_event.call_count}" ) - + # Use the filtered call for assertions call_args = calls_with_expected_input[0] standard_logging_object = call_args.kwargs["kwargs"][ @@ -326,14 +336,18 @@ async def test_logging_non_streaming_request(): @pytest.mark.parametrize("async_flag", ["acompletion", "aresponses"]) -def test_success_handler_skips_sync_callbacks_for_async_requests(logging_obj, async_flag): +def test_success_handler_skips_sync_callbacks_for_async_requests( + logging_obj, async_flag +): """Ensure sync success callbacks are skipped when async call type flags are set.""" from litellm.integrations.custom_logger import CustomLogger class DummyLogger(CustomLogger): pass - logging_obj.stream = False # simulate non-streaming request where sync callbacks would normally run + logging_obj.stream = ( + False # simulate non-streaming request where sync callbacks would normally run + ) logging_obj.model_call_details["litellm_params"] = {async_flag: True} logging_obj.litellm_params = logging_obj.model_call_details["litellm_params"] @@ -523,7 +537,7 @@ def test_success_handler_runs_guardrail_logging_hook_when_enabled(logging_obj): assert guardrail_call_kwargs["event_type"] == GuardrailEventHooks.logging_only guardrail.logging_hook.assert_called_once() assert logging_obj.model_call_details.get("guardrail_hook_ran") is True - + def test_get_user_agent_tags(): from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup @@ -676,21 +690,29 @@ def test_get_request_tags_does_not_mutate_original_tags(): ) # Verify the original tags list was NOT mutated - assert original_tags == ["custom-tag-1", "custom-tag-2"], ( - f"Original tags list was mutated: {original_tags}" - ) - assert metadata["tags"] == ["custom-tag-1", "custom-tag-2"], ( - f"metadata['tags'] was mutated: {metadata['tags']}" - ) + assert original_tags == [ + "custom-tag-1", + "custom-tag-2", + ], f"Original tags list was mutated: {original_tags}" + assert metadata["tags"] == [ + "custom-tag-1", + "custom-tag-2", + ], f"metadata['tags'] was mutated: {metadata['tags']}" # Verify each returned list has exactly 2 User-Agent tags (not duplicated) user_agent_count_1 = len([t for t in tags1 if t.startswith("User-Agent:")]) user_agent_count_2 = len([t for t in tags2 if t.startswith("User-Agent:")]) user_agent_count_3 = len([t for t in tags3 if t.startswith("User-Agent:")]) - assert user_agent_count_1 == 2, f"Expected 2 User-Agent tags, got {user_agent_count_1}" - assert user_agent_count_2 == 2, f"Expected 2 User-Agent tags, got {user_agent_count_2}" - assert user_agent_count_3 == 2, f"Expected 2 User-Agent tags, got {user_agent_count_3}" + assert ( + user_agent_count_1 == 2 + ), f"Expected 2 User-Agent tags, got {user_agent_count_1}" + assert ( + user_agent_count_2 == 2 + ), f"Expected 2 User-Agent tags, got {user_agent_count_2}" + assert ( + user_agent_count_3 == 2 + ), f"Expected 2 User-Agent tags, got {user_agent_count_3}" # Verify all returned lists are independent (different objects) assert tags1 is not tags2 @@ -795,7 +817,6 @@ def test_get_extra_header_tags(): def test_response_cost_calculator_with_response_cost_in_hidden_params(logging_obj): from litellm import Router - from litellm.litellm_core_utils.litellm_logging import Logging router = Router( model_list=[ @@ -933,7 +954,6 @@ async def test_e2e_generate_cold_storage_object_key_successful(): with patch("litellm.cold_storage_custom_logger", return_value="s3"), patch( "litellm.integrations.s3.get_s3_object_key" ) as mock_get_s3_key: - # Mock the S3 object key generation to return a predictable result mock_get_s3_key.return_value = ( "2025-01-15/time-10-30-45-123456_chatcmpl-test-12345.json" @@ -981,7 +1001,6 @@ async def test_e2e_generate_cold_storage_object_key_with_custom_logger_s3_path() ) as mock_get_logger, patch( "litellm.integrations.s3.get_s3_object_key" ) as mock_get_s3_key: - # Setup mocks mock_get_logger.return_value = mock_custom_logger mock_get_s3_key.return_value = ( @@ -1033,7 +1052,6 @@ async def test_e2e_generate_cold_storage_object_key_with_logger_no_s3_path(): ) as mock_get_logger, patch( "litellm.integrations.s3.get_s3_object_key" ) as mock_get_s3_key: - # Setup mocks mock_get_logger.return_value = mock_custom_logger mock_get_s3_key.return_value = ( @@ -1279,9 +1297,9 @@ async def test_async_success_handler_sets_standard_logging_object_for_pass_throu "standard_logging_object should be set for pass-through endpoints " "even when complete_streaming_response is None" ) - assert logging_obj.model_call_details["standard_logging_object"] is not None, ( - "standard_logging_object should not be None for pass-through endpoints" - ) + assert ( + logging_obj.model_call_details["standard_logging_object"] is not None + ), "standard_logging_object should not be None for pass-through endpoints" # Verify that async_complete_streaming_response was set to prevent re-processing # This is consistent with the existing code pattern for regular streaming @@ -1289,15 +1307,15 @@ async def test_async_success_handler_sets_standard_logging_object_for_pass_throu "async_complete_streaming_response should be set to prevent re-processing, " "consistent with the existing code pattern" ) - assert logging_obj.model_call_details["async_complete_streaming_response"] is result, ( - "async_complete_streaming_response should be set to the result" - ) + assert ( + logging_obj.model_call_details["async_complete_streaming_response"] is result + ), "async_complete_streaming_response should be set to the result" # Verify that response_cost is set to None (cost calculation not possible for pass-through) # This is consistent with the error handling in the non-pass-through code path - assert "response_cost" in logging_obj.model_call_details, ( - "response_cost should be set for pass-through endpoints" - ) + assert ( + "response_cost" in logging_obj.model_call_details + ), "response_cost should be set for pass-through endpoints" assert logging_obj.model_call_details["response_cost"] is None, ( "response_cost should be None for pass-through endpoints since " "StandardPassThroughResponseObject doesn't have standard usage info" @@ -1356,10 +1374,14 @@ async def test_async_success_handler_prevents_reprocessing_for_pass_through_endp # Verify first call set the values assert "standard_logging_object" in logging_obj.model_call_details assert "async_complete_streaming_response" in logging_obj.model_call_details - first_standard_logging_object = logging_obj.model_call_details["standard_logging_object"] + first_standard_logging_object = logging_obj.model_call_details[ + "standard_logging_object" + ] # Second call - should return early due to async_complete_streaming_response guard - with patch.object(logging_obj, "get_combined_callback_list", return_value=[]) as mock_callbacks: + with patch.object( + logging_obj, "get_combined_callback_list", return_value=[] + ) as mock_callbacks: await logging_obj.async_success_handler( result=result, start_time=start_time, @@ -1370,9 +1392,10 @@ async def test_async_success_handler_prevents_reprocessing_for_pass_through_endp mock_callbacks.assert_not_called() # Verify standard_logging_object wasn't modified by second call - assert logging_obj.model_call_details["standard_logging_object"] is first_standard_logging_object, ( - "standard_logging_object should not be modified on re-processing" - ) + assert ( + logging_obj.model_call_details["standard_logging_object"] + is first_standard_logging_object + ), "standard_logging_object should not be modified on re-processing" @pytest.mark.asyncio @@ -1433,9 +1456,11 @@ async def test_async_success_handler_sets_standard_logging_object_for_streaming_ "standard_logging_object should be set for streaming pass-through endpoints " "even when the response cannot be parsed into a ModelResponse" ) - assert logging_obj.model_call_details["standard_logging_object"] is not None, ( - "standard_logging_object should not be None for streaming pass-through endpoints" - ) + assert ( + logging_obj.model_call_details["standard_logging_object"] is not None + ), "standard_logging_object should not be None for streaming pass-through endpoints" + + def test_get_error_information_error_code_priority(): """ Test get_error_information prioritizes 'code' attribute over 'status_code' attribute @@ -1680,3 +1705,131 @@ async def test_async_success_handler_preserves_response_cost_for_pass_through_en slo = logging_obj.model_call_details.get("standard_logging_object") assert slo is not None assert slo["response_cost"] > 0 + + +def test_function_setup_litellm_metadata_populates_metadata(): + """ + Test that function_setup() properly handles litellm_metadata (used by /v1/messages, + /batches, /responses, /files endpoints) and populates litellm_params["metadata"] + so callbacks like Langfuse can read API key fields. + + This is the root cause of: Claude Code requests missing user_api_key_hash in Langfuse. + """ + import litellm + + test_api_key_hash = "sk-hashed-1234567890abcdef" + test_team_id = "team-test-123" + test_key_alias = "my-test-key" + + # Simulate what happens for /v1/messages: metadata is in "litellm_metadata", not "metadata" + kwargs = { + "model": "claude-3-5-sonnet", + "messages": [{"role": "user", "content": "hello"}], + "litellm_call_id": "test-call-id-123", + "litellm_metadata": { + "user_api_key_hash": test_api_key_hash, + "user_api_key_alias": test_key_alias, + "user_api_key_team_id": test_team_id, + "user_api_key_user_id": "user-123", + "user_api_key": test_api_key_hash, + }, + } + + logging_obj, returned_kwargs = litellm.utils.function_setup( + original_function="anthropic_messages", + rules_obj=litellm.utils.Rules(), + start_time=time.time(), + **kwargs, + ) + + # litellm_params["metadata"] must contain the API key fields + litellm_params = logging_obj.model_call_details.get("litellm_params", {}) + metadata = litellm_params.get("metadata") + assert metadata is not None, "litellm_params['metadata'] should not be None" + assert isinstance(metadata, dict), "litellm_params['metadata'] should be a dict" + assert metadata.get("user_api_key_hash") == test_api_key_hash + assert metadata.get("user_api_key_alias") == test_key_alias + assert metadata.get("user_api_key_team_id") == test_team_id + + # litellm_metadata should also be preserved + litellm_metadata = litellm_params.get("litellm_metadata") + assert litellm_metadata is not None + assert litellm_metadata.get("user_api_key_hash") == test_api_key_hash + + # metadata should be a COPY, not an alias — mutating one must not affect the other + assert ( + metadata is not litellm_metadata + ), "litellm_params['metadata'] should be a copy, not the same object" + + +def test_function_setup_metadata_takes_precedence_over_litellm_metadata(): + """ + Test that when BOTH metadata and litellm_metadata are present (e.g., user sets + Anthropic API metadata AND proxy adds litellm_metadata), metadata is used as + litellm_params["metadata"] and litellm_metadata is stored separately. + """ + import litellm + + kwargs = { + "model": "claude-3-5-sonnet", + "messages": [{"role": "user", "content": "hello"}], + "litellm_call_id": "test-call-id-456", + "metadata": { + "user_id": "anthropic-user-id", + }, + "litellm_metadata": { + "user_api_key_hash": "sk-hashed-xyz", + "user_api_key_team_id": "team-xyz", + }, + } + + logging_obj, _ = litellm.utils.function_setup( + original_function="anthropic_messages", + rules_obj=litellm.utils.Rules(), + start_time=time.time(), + **kwargs, + ) + + litellm_params = logging_obj.model_call_details.get("litellm_params", {}) + + # When both are present, metadata should be the explicit "metadata" dict + metadata = litellm_params.get("metadata") + assert metadata is not None + assert metadata.get("user_id") == "anthropic-user-id" + + # litellm_metadata should be preserved separately for merge_litellm_metadata() + litellm_metadata = litellm_params.get("litellm_metadata") + assert litellm_metadata is not None + assert litellm_metadata.get("user_api_key_hash") == "sk-hashed-xyz" + + +def test_function_setup_empty_metadata_falls_back_to_litellm_metadata(): + """ + Test that when metadata is explicitly set to {} (empty dict), litellm_metadata + is still used to populate litellm_params["metadata"] so API key fields are visible. + """ + import litellm + + kwargs = { + "model": "claude-3-5-sonnet", + "messages": [{"role": "user", "content": "hello"}], + "litellm_call_id": "test-call-id-789", + "metadata": {}, + "litellm_metadata": { + "user_api_key_hash": "sk-hashed-empty-test", + "user_api_key_team_id": "team-empty-test", + }, + } + + logging_obj, _ = litellm.utils.function_setup( + original_function="anthropic_messages", + rules_obj=litellm.utils.Rules(), + start_time=time.time(), + **kwargs, + ) + + litellm_params = logging_obj.model_call_details.get("litellm_params", {}) + metadata = litellm_params.get("metadata") + assert metadata is not None + assert metadata.get("user_api_key_hash") == "sk-hashed-empty-test" + assert metadata.get("user_api_key_team_id") == "team-empty-test"