mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
acd4f0eb04
commit
9af014d75a
2 changed files with 63 additions and 3 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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={
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue