diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 15d1876e5a2..3bd92d21766 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -1048,16 +1048,10 @@ class LiteLLMProxyRequestSetup: if "disable_global_guardrails" in key_metadata and isinstance(key_metadata["disable_global_guardrails"], bool): data[_metadata_variable_name]["disable_global_guardrails"] = key_metadata["disable_global_guardrails"] if "spend_logs_metadata" in key_metadata and isinstance(key_metadata["spend_logs_metadata"], dict): - if "spend_logs_metadata" in data[_metadata_variable_name] and isinstance( - data[_metadata_variable_name]["spend_logs_metadata"], dict - ): - for key, value in key_metadata["spend_logs_metadata"].items(): - if ( - key not in data[_metadata_variable_name]["spend_logs_metadata"] - ): # don't override k-v pair sent by request (user request) - data[_metadata_variable_name]["spend_logs_metadata"][key] = value - else: - data[_metadata_variable_name]["spend_logs_metadata"] = key_metadata["spend_logs_metadata"] + data[_metadata_variable_name]["spend_logs_metadata"] = LiteLLMProxyRequestSetup._merge_spend_logs_metadata( + existing=data[_metadata_variable_name].get("spend_logs_metadata"), + to_add=key_metadata["spend_logs_metadata"], + ) ## KEY-LEVEL DISABLE FALLBACKS if "disable_fallbacks" in key_metadata and isinstance(key_metadata["disable_fallbacks"], bool): @@ -1095,6 +1089,14 @@ class LiteLLMProxyRequestSetup: return final_tags + @staticmethod + def _merge_spend_logs_metadata(existing: dict | None, to_add: dict | None) -> dict | None: + if not isinstance(to_add, dict): + return existing + if not isinstance(existing, dict): + return dict(to_add) + return {**to_add, **existing} + @staticmethod def add_team_based_callbacks_from_config( team_id: str, @@ -1555,16 +1557,19 @@ async def add_litellm_data_to_request( ): data[_metadata_variable_name]["opted_out_global_guardrails"] = team_metadata["opted_out_global_guardrails"] if "spend_logs_metadata" in team_metadata and isinstance(team_metadata["spend_logs_metadata"], dict): - if "spend_logs_metadata" in data[_metadata_variable_name] and isinstance( - data[_metadata_variable_name]["spend_logs_metadata"], dict - ): - for key, value in team_metadata["spend_logs_metadata"].items(): - if ( - key not in data[_metadata_variable_name]["spend_logs_metadata"] - ): # don't override k-v pair sent by request (user request) - data[_metadata_variable_name]["spend_logs_metadata"][key] = value - else: - data[_metadata_variable_name]["spend_logs_metadata"] = team_metadata["spend_logs_metadata"] + data[_metadata_variable_name]["spend_logs_metadata"] = LiteLLMProxyRequestSetup._merge_spend_logs_metadata( + existing=data[_metadata_variable_name].get("spend_logs_metadata"), + to_add=team_metadata["spend_logs_metadata"], + ) + + organization_metadata = user_api_key_dict.organization_metadata or {} + if "spend_logs_metadata" in organization_metadata and isinstance( + organization_metadata["spend_logs_metadata"], dict + ): + data[_metadata_variable_name]["spend_logs_metadata"] = LiteLLMProxyRequestSetup._merge_spend_logs_metadata( + existing=data[_metadata_variable_name].get("spend_logs_metadata"), + to_add=organization_metadata["spend_logs_metadata"], + ) ## PROJECT-LEVEL TAGS project_metadata = user_api_key_dict.project_metadata or {} diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index d2b8b7ec23d..27bfb620dec 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -5118,3 +5118,98 @@ async def test_add_litellm_data_to_request_unions_metadata_tags_with_header_tags tags = updated["litellm_metadata"]["tags"] assert "header-tag" in tags assert "body-tag" in tags + + +def _spend_logs_request_mock(): + request_mock = MagicMock(spec=Request) + request_mock.url = MagicMock() + request_mock.url.path = "/v1/chat/completions" + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + return request_mock + + +@pytest.mark.asyncio +async def test_organization_spend_logs_metadata_is_merged(): + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + metadata={}, + team_metadata={}, + organization_metadata={"spend_logs_metadata": {"cost_center": "org-123"}}, + ) + + updated = await add_litellm_data_to_request( + data={"model": "gpt-3.5-turbo"}, + request=_spend_logs_request_mock(), + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated["metadata"]["spend_logs_metadata"] == {"cost_center": "org-123"} + + +@pytest.mark.asyncio +async def test_spend_logs_metadata_precedence_request_key_team_org(): + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + metadata={"spend_logs_metadata": {"level": "key", "from_key": "k"}}, + team_metadata={"spend_logs_metadata": {"level": "team", "from_team": "t"}}, + organization_metadata={"spend_logs_metadata": {"level": "org", "from_org": "o"}}, + ) + + updated = await add_litellm_data_to_request( + data={ + "model": "gpt-3.5-turbo", + "metadata": {"spend_logs_metadata": {"level": "request", "from_request": "r"}}, + }, + request=_spend_logs_request_mock(), + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + merged = updated["metadata"]["spend_logs_metadata"] + assert merged["level"] == "request" + assert merged["from_request"] == "r" + assert merged["from_key"] == "k" + assert merged["from_team"] == "t" + assert merged["from_org"] == "o" + + +@pytest.mark.asyncio +async def test_key_spend_logs_metadata_not_mutated_by_team_and_org(): + key_metadata = {"spend_logs_metadata": {"from_key": "k"}} + team_metadata = {"spend_logs_metadata": {"from_team": "t"}} + organization_metadata = {"spend_logs_metadata": {"from_org": "o"}} + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + metadata=key_metadata, + team_metadata=team_metadata, + organization_metadata=organization_metadata, + ) + + updated = await add_litellm_data_to_request( + data={"model": "gpt-3.5-turbo"}, + request=_spend_logs_request_mock(), + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated["metadata"]["spend_logs_metadata"] == { + "from_key": "k", + "from_team": "t", + "from_org": "o", + } + assert key_metadata["spend_logs_metadata"] == {"from_key": "k"} + assert team_metadata["spend_logs_metadata"] == {"from_team": "t"} + assert organization_metadata["spend_logs_metadata"] == {"from_org": "o"}