fix(spend_tracking): keep inferred provider out of model reconstruction and OAuth provider lookup

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-09-16 23:36:54 +00:00
parent acd4f0eb04
commit 9af014d75a
2 changed files with 63 additions and 3 deletions

View file

@ -32,6 +32,7 @@ from litellm.litellm_core_utils.core_helpers import (
get_litellm_metadata_from_kwargs,
reconstruct_model_name,
)
from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider
from litellm.litellm_core_utils.internal_call_metadata import is_unbilled_non_inference_call
from litellm.litellm_core_utils.litellm_logging import (
coerce_model_access_groups,
@ -345,6 +346,9 @@ def _sl_attribution_fallback(
def _deployment_provider(deployment: DeploymentTypedDict) -> str | None:
litellm_params: Final = LiteLLM_Params.model_validate(deployment["litellm_params"])
declared: Final = declared_authenticating_provider(litellm_params.model, litellm_params.custom_llm_provider)
if declared is not None:
return declared
try:
_, provider, _, _ = litellm.get_llm_provider(
model=litellm_params.model, custom_llm_provider=litellm_params.custom_llm_provider
@ -468,15 +472,16 @@ def get_logging_payload(
hidden_params: Final = standard_logging_payload.get("hidden_params", {})
litellm_overhead_time_ms = hidden_params.get("litellm_overhead_time_ms")
custom_llm_provider: Final = (
logged_provider: Final = (
kwargs.get("custom_llm_provider")
or _sl_attribution_fallback(standard_logging_payload, "custom_llm_provider")
or _model_group_provider(_model_group, llm_router)
or None
)
custom_llm_provider: Final = logged_provider or _model_group_provider(_model_group, llm_router)
raw_model: Final = cast(str, kwargs.get("model") or "")
resolved_model: Final = (
standard_logging_payload.get("model") if standard_logging_payload is not None else None
) or reconstruct_model_name(raw_model, custom_llm_provider, metadata or {})
) or reconstruct_model_name(raw_model, logged_provider, metadata or {})
failed_with_prompt_shaped_model: Final = (
_get_status_for_spend_log(metadata=metadata) == "failure"
and not _model_group

View file

@ -4049,6 +4049,61 @@ def test_get_logging_payload_router_rejected_request_without_router_leaves_provi
assert _router_rejected_failure_payload("openai-group", None)["custom_llm_provider"] == ""
@pytest.mark.parametrize(
"litellm_params,expected_provider",
[
({"model": "github_copilot/gpt-4o"}, "github_copilot"),
({"model": "gpt-5", "custom_llm_provider": "chatgpt"}, "chatgpt"),
],
)
def test_get_logging_payload_inferred_provider_never_resolves_declared_authenticating_providers(
monkeypatch, litellm_params: dict[str, str], expected_provider: str
):
resolution_attempts: list[str] = []
def _router_init_stub(model, custom_llm_provider=None, *args, **kwargs):
return model.split("/", 1)[-1], custom_llm_provider or model.split("/", 1)[0], None, None
def _oauth_tripwire(model, *args, **kwargs):
resolution_attempts.append(model)
raise AssertionError("get_llm_provider would run the OAuth device flow")
monkeypatch.setattr(litellm, "get_llm_provider", _router_init_stub)
llm_router = litellm.Router(model_list=[{"model_name": "oauth-group", "litellm_params": litellm_params}])
monkeypatch.setattr(litellm, "get_llm_provider", _oauth_tripwire)
payload = _router_rejected_failure_payload("oauth-group", llm_router)
assert payload["custom_llm_provider"] == expected_provider
assert resolution_attempts == []
def test_get_logging_payload_inferred_provider_does_not_rewrite_spend_log_model():
llm_router = litellm.Router(
model_list=[
{
"model_name": "bedrock-group",
"litellm_params": {
"model": "bedrock/anthropic.claude-3-5-haiku-20241022-v1:0",
"aws_region_name": "us-east-1",
},
},
{
"model_name": "bedrock-group",
"litellm_params": {
"model": "bedrock/anthropic.claude-3-5-haiku-20241022-v1:0",
"aws_region_name": "us-west-2",
},
},
]
)
payload = _router_rejected_failure_payload("bedrock-group", llm_router)
assert payload["custom_llm_provider"] == "bedrock"
assert payload["model"] == "bedrock-group"
def test_get_logging_payload_logged_provider_wins_over_model_group_provider():
payload = get_logging_payload(
kwargs={