diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index a43574b1a04..ad7bdcb91ed 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -6274,7 +6274,7 @@ class StandardLoggingPayloadSetup: return None user_agent_tags: list[str] | None = None headers: Final = proxy_server_request.get("headers", {}) - if headers is not None and isinstance(headers, dict): + if headers is not None and isinstance(headers, Mapping): if "user-agent" in headers: user_agent: Final = headers["user-agent"] if user_agent is not None: @@ -6299,7 +6299,7 @@ class StandardLoggingPayloadSetup: return None headers: Final = proxy_server_request.get("headers", {}) - if not isinstance(headers, dict): + if not isinstance(headers, Mapping): return None header_tags: Final = [] diff --git a/tests/integration/spend/test_passthrough_spend_tags.py b/tests/integration/spend/test_passthrough_spend_tags.py new file mode 100644 index 00000000000..9d89a866cc8 --- /dev/null +++ b/tests/integration/spend/test_passthrough_spend_tags.py @@ -0,0 +1,76 @@ +import json +import uuid +from hashlib import sha256 +from pathlib import Path +from typing import Final + +import pytest +import yaml +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, wire_server + +MODEL: Final = "claude-sonnet-4-5-20250929" +SENT_HEADERS: Final = {"user-agent": "claude-cli/2.0.0", "x-tenant-id": "tenant-a"} +EXPECTED_TAGS: Final = ["User-Agent: claude-cli", "User-Agent: claude-cli/2.0.0", "x-tenant-id: tenant-a"] + + +def _respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/v1/messages", request.target + return Reply( + body=json.dumps( + { + "id": f"msg_{uuid.uuid4().hex}", + "type": "message", + "role": "assistant", + "model": MODEL, + "content": [{"type": "text", "text": "tagged"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 10, "output_tokens": 2}, + } + ).encode() + ) + + +def _request_tags(key: str) -> list[list[str]]: + rows: Final = read_rows( + 'SELECT request_tags FROM "LiteLLM_SpendLogs" WHERE api_key=%s', (sha256(key.encode()).hexdigest(),) + ) + return [ + json.loads(row["request_tags"]) if isinstance(row["request_tags"], str) else row["request_tags"] for row in rows + ] + + +@pytest.mark.parametrize( + "route", [pytest.param("/anthropic/v1/messages", id="passthrough"), pytest.param("/v1/messages", id="unified")] +) +def test_header_derived_spend_tags_are_recorded_on_anthropic_messages_routes( + gateway: Gateway, tmp_path: Path, route: str +) -> None: + with wire_server(_respond) as wire: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"]["extra_spend_tag_headers"] = ["x-tenant-id"] + path: Final = tmp_path / "spend-tag-headers.yaml" + path.write_text(yaml.safe_dump(config)) + environment: Final = {"ANTHROPIC_API_BASE": wire.url, "ANTHROPIC_API_KEY": "synthetic-anthropic-key"} + with owned_proxy(gateway, tmp_path, environment, config=path) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model( + model=f"anthropic/{MODEL}", api_base=wire.url, api_key="synthetic-anthropic-key" + ) + key: Final = scenario.key() + response: Final = candidate.request( + "POST", + route, + { + "model": MODEL if route == "/anthropic/v1/messages" else model, + "max_tokens": 16, + "messages": [{"role": "user", "content": "tag me"}], + }, + key=key, + headers={**SENT_HEADERS, "anthropic-version": "2023-06-01"}, + ) + assert response.status_code == 200, response.text + assert len(wire.drain()) == 1 + assert eventually(lambda: _request_tags(key), lambda tags: len(tags) == 1, seconds=70) == [EXPECTED_TAGS] diff --git a/tests/unit/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py index e07ffe00d4c..494e82d22e7 100644 --- a/tests/unit/litellm_core_utils/test_litellm_logging.py +++ b/tests/unit/litellm_core_utils/test_litellm_logging.py @@ -2959,6 +2959,30 @@ def test_get_extra_header_tags(): delattr(litellm, "extra_spend_tag_headers") +def test_get_request_tags_reads_header_tags_from_starlette_headers(): + from starlette.datastructures import Headers + + import litellm + from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup + + original_extra_headers = getattr(litellm, "extra_spend_tag_headers", None) + original_disable_user_agent = litellm.disable_add_user_agent_to_request_tags + try: + litellm.extra_spend_tag_headers = ["x-tenant-id"] + litellm.disable_add_user_agent_to_request_tags = False + proxy_server_request = {"headers": Headers({"user-agent": "claude-cli/2.0.0", "x-tenant-id": "tenant-a"})} + + assert StandardLoggingPayloadSetup._get_request_tags( + litellm_params={}, proxy_server_request=proxy_server_request + ) == ["User-Agent: claude-cli", "User-Agent: claude-cli/2.0.0", "x-tenant-id: tenant-a"] + finally: + if original_extra_headers is not None: + litellm.extra_spend_tag_headers = original_extra_headers + elif hasattr(litellm, "extra_spend_tag_headers"): + delattr(litellm, "extra_spend_tag_headers") + litellm.disable_add_user_agent_to_request_tags = original_disable_user_agent + + def test_response_cost_calculator_with_response_cost_in_hidden_params(logging_obj): from litellm import Router