mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(proxy): record header-derived spend tags on pass-through routes
Co-authored-by: Quinn Xu <satonaoyukii@gmail.com> Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
276fc9c63a
commit
65bb28221c
3 changed files with 102 additions and 2 deletions
|
|
@ -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 = []
|
||||
|
|
|
|||
76
tests/integration/spend/test_passthrough_spend_tags.py
Normal file
76
tests/integration/spend/test_passthrough_spend_tags.py
Normal file
|
|
@ -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]
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue