mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(proxy): don't replace response model with alias (#21874)
* fix(proxy): preserve actual model name in response, not alias * fix mypy attr defined error on BaseModel.model * update tests for new prefix stripping behavior * add nested-slash test and debug logging on strip failures
This commit is contained in:
parent
56557bfae1
commit
98d2275d13
4 changed files with 130 additions and 356 deletions
|
|
@ -45,7 +45,7 @@ from litellm.proxy.dd_span_tagger import DDSpanTagger
|
|||
from litellm.proxy.route_llm_request import route_request
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.router import Router
|
||||
from litellm.types.utils import ServerToolUse
|
||||
from litellm.types.utils import ServerToolUse, LlmProvidersSet
|
||||
|
||||
# Type alias for streaming chunk serializer (chunk after hooks + cost injection -> wire format)
|
||||
StreamChunkSerializer = Callable[[Any], str]
|
||||
|
|
@ -249,85 +249,44 @@ async def create_response(
|
|||
def _override_openai_response_model(
|
||||
*,
|
||||
response_obj: Any,
|
||||
requested_model: str,
|
||||
log_context: str,
|
||||
) -> None:
|
||||
"""
|
||||
Force the OpenAI-compatible `model` field in the response to match what the client requested.
|
||||
Strip known LiteLLM provider prefixes (e.g. hosted_vllm/) from the response model field.
|
||||
|
||||
LiteLLM internally prefixes some provider/deployment model identifiers (e.g. `hosted_vllm/...`).
|
||||
That internal identifier should not be returned to clients in the OpenAI `model` field.
|
||||
|
||||
Note: This is intentionally verbose. A model mismatch is a useful signal that an internal
|
||||
model identifier is being stamped/preserved somewhere in the request/response pipeline.
|
||||
We log mismatches as warnings (and then restamp to the client-requested value) so these
|
||||
paths stay observable for maintainers/operators without breaking client compatibility.
|
||||
|
||||
Errors are reserved for cases where the proxy cannot read/override the response model field.
|
||||
|
||||
Exception: If a fallback occurred (indicated by x-litellm-attempted-fallbacks header),
|
||||
we should preserve the actual model that was used (the fallback model) rather than
|
||||
overriding it with the originally requested model.
|
||||
Previously this replaced response.model with the client-requested alias, but that
|
||||
hid the actual model name from callers (see #21665). Now we only strip internal
|
||||
provider routing prefixes, preserving the real model name.
|
||||
"""
|
||||
if not requested_model:
|
||||
return
|
||||
|
||||
# Check if a fallback occurred - if so, preserve the actual model used
|
||||
hidden_params = getattr(response_obj, "_hidden_params", {}) or {}
|
||||
if isinstance(hidden_params, dict):
|
||||
fallback_headers = hidden_params.get("additional_headers", {}) or {}
|
||||
attempted_fallbacks = fallback_headers.get(
|
||||
"x-litellm-attempted-fallbacks", None
|
||||
)
|
||||
if attempted_fallbacks is not None and attempted_fallbacks > 0:
|
||||
# A fallback occurred - preserve the actual model that was used
|
||||
verbose_proxy_logger.debug(
|
||||
"%s: fallback detected (attempted_fallbacks=%d), preserving actual model used instead of overriding to requested model.",
|
||||
log_context,
|
||||
attempted_fallbacks,
|
||||
)
|
||||
return
|
||||
|
||||
if isinstance(response_obj, dict):
|
||||
downstream_model = response_obj.get("model")
|
||||
if downstream_model != requested_model:
|
||||
elif hasattr(response_obj, "model"):
|
||||
downstream_model = getattr(response_obj, "model", None)
|
||||
else:
|
||||
return
|
||||
|
||||
if not downstream_model or not isinstance(downstream_model, str):
|
||||
return
|
||||
|
||||
if "/" not in downstream_model:
|
||||
return
|
||||
|
||||
prefix = downstream_model.split("/", 1)[0]
|
||||
if prefix not in LlmProvidersSet:
|
||||
return
|
||||
|
||||
stripped = downstream_model.split("/", 1)[1]
|
||||
|
||||
if isinstance(response_obj, dict):
|
||||
response_obj["model"] = stripped
|
||||
else:
|
||||
try:
|
||||
setattr(response_obj, "model", stripped)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(
|
||||
"%s: response model mismatch - requested=%r downstream=%r. Overriding response['model'] to requested model.",
|
||||
log_context,
|
||||
requested_model,
|
||||
downstream_model,
|
||||
"%s: failed to strip provider prefix on response.model, error=%s",
|
||||
log_context, str(e),
|
||||
)
|
||||
response_obj["model"] = requested_model
|
||||
return
|
||||
|
||||
if not hasattr(response_obj, "model"):
|
||||
verbose_proxy_logger.error(
|
||||
"%s: cannot override response model; missing `model` attribute. response_type=%s",
|
||||
log_context,
|
||||
type(response_obj),
|
||||
)
|
||||
return
|
||||
|
||||
downstream_model = getattr(response_obj, "model", None)
|
||||
if downstream_model != requested_model:
|
||||
verbose_proxy_logger.debug(
|
||||
"%s: response model mismatch - requested=%r downstream=%r. Overriding response.model to requested model.",
|
||||
log_context,
|
||||
requested_model,
|
||||
downstream_model,
|
||||
)
|
||||
|
||||
try:
|
||||
setattr(response_obj, "model", requested_model)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
"%s: failed to override response.model=%r on response_type=%s. error=%s",
|
||||
log_context,
|
||||
requested_model,
|
||||
type(response_obj),
|
||||
str(e),
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
|
||||
def _get_cost_breakdown_from_logging_obj(
|
||||
|
|
@ -815,9 +774,6 @@ class ProxyBaseLLMRequestProcessing:
|
|||
"""
|
||||
Common request processing logic for both chat completions and responses API endpoints
|
||||
"""
|
||||
requested_model_from_client: Optional[str] = (
|
||||
self.data.get("model") if isinstance(self.data.get("model"), str) else None
|
||||
)
|
||||
self._debug_log_request_payload()
|
||||
|
||||
self.data, logging_obj = await self.common_processing_pre_call_logic(
|
||||
|
|
@ -924,14 +880,6 @@ class ProxyBaseLLMRequestProcessing:
|
|||
if callback_headers:
|
||||
custom_headers.update(callback_headers)
|
||||
|
||||
# Preserve the original client-requested model (pre-alias mapping) for downstream
|
||||
# streaming generators. Pre-call processing can rewrite `self.data["model"]` for
|
||||
# aliasing/routing, but the OpenAI-compatible response `model` field should reflect
|
||||
# what the client sent.
|
||||
if requested_model_from_client:
|
||||
self.data[
|
||||
"_litellm_client_requested_model"
|
||||
] = requested_model_from_client
|
||||
if route_type == "allm_passthrough_route":
|
||||
# Check if response is an async generator
|
||||
if self._is_streaming_response(response):
|
||||
|
|
@ -991,14 +939,11 @@ class ProxyBaseLLMRequestProcessing:
|
|||
data=self.data, user_api_key_dict=user_api_key_dict, response=response
|
||||
)
|
||||
|
||||
# Always return the client-requested model name (not provider-prefixed internal identifiers)
|
||||
# for OpenAI-compatible responses.
|
||||
if requested_model_from_client:
|
||||
_override_openai_response_model(
|
||||
response_obj=response,
|
||||
requested_model=requested_model_from_client,
|
||||
log_context=f"litellm_call_id={logging_obj.litellm_call_id}",
|
||||
)
|
||||
# Strip any internal provider prefixes from the response model field.
|
||||
_override_openai_response_model(
|
||||
response_obj=response,
|
||||
log_context=f"litellm_call_id={logging_obj.litellm_call_id}",
|
||||
)
|
||||
|
||||
hidden_params = (
|
||||
getattr(response, "_hidden_params", {}) or {}
|
||||
|
|
|
|||
|
|
@ -102,6 +102,7 @@ from litellm.types.utils import (
|
|||
ModelResponseStream,
|
||||
TextCompletionResponse,
|
||||
TokenCountResponse,
|
||||
LlmProvidersSet,
|
||||
)
|
||||
from litellm.utils import (
|
||||
_invalidate_model_cost_lowercase_map,
|
||||
|
|
@ -5395,64 +5396,46 @@ async def async_assistants_data_generator(
|
|||
yield f"data: {error_returned}\n\n"
|
||||
|
||||
|
||||
def _get_client_requested_model_for_streaming(request_data: dict) -> str:
|
||||
"""
|
||||
Prefer the original client-requested model (pre-alias mapping) when available.
|
||||
|
||||
Pre-call processing can rewrite `request_data["model"]` for aliasing/routing purposes.
|
||||
The OpenAI-compatible public `model` field should reflect what the client sent.
|
||||
"""
|
||||
requested_model = request_data.get("_litellm_client_requested_model")
|
||||
if isinstance(requested_model, str):
|
||||
return requested_model
|
||||
|
||||
requested_model = request_data.get("model")
|
||||
return requested_model if isinstance(requested_model, str) else ""
|
||||
|
||||
|
||||
def _restamp_streaming_chunk_model(
|
||||
*,
|
||||
chunk: Any,
|
||||
requested_model_from_client: str,
|
||||
request_data: dict,
|
||||
model_mismatch_logged: bool,
|
||||
) -> Tuple[Any, bool]:
|
||||
# Always return the client-requested model name (not provider-prefixed internal identifiers)
|
||||
# on streaming chunks.
|
||||
#
|
||||
# Note: This warning is intentionally verbose. A mismatch is a useful signal that an
|
||||
# internal provider/deployment identifier is leaking into the public API, and helps
|
||||
# maintainers/operators catch regressions while preserving OpenAI-compatible output.
|
||||
if not requested_model_from_client or not isinstance(chunk, (BaseModel, dict)):
|
||||
"""Strip known provider prefixes from streaming chunk model field."""
|
||||
if not isinstance(chunk, (BaseModel, dict)):
|
||||
return chunk, model_mismatch_logged
|
||||
|
||||
downstream_model = (
|
||||
chunk.get("model") if isinstance(chunk, dict) else getattr(chunk, "model", None)
|
||||
)
|
||||
if not model_mismatch_logged and downstream_model != requested_model_from_client:
|
||||
|
||||
if not downstream_model or not isinstance(downstream_model, str) or "/" not in downstream_model:
|
||||
return chunk, model_mismatch_logged
|
||||
|
||||
prefix = downstream_model.split("/", 1)[0]
|
||||
if prefix not in LlmProvidersSet:
|
||||
return chunk, model_mismatch_logged
|
||||
|
||||
stripped = downstream_model.split("/", 1)[1]
|
||||
|
||||
if not model_mismatch_logged:
|
||||
verbose_proxy_logger.debug(
|
||||
"litellm_call_id=%s: streaming chunk model mismatch - requested=%r downstream=%r. Overriding model to requested.",
|
||||
request_data.get("litellm_call_id"),
|
||||
requested_model_from_client,
|
||||
downstream_model,
|
||||
"litellm_call_id=%s: stripping provider prefix %r from chunk model %r",
|
||||
request_data.get("litellm_call_id"), prefix, downstream_model,
|
||||
)
|
||||
model_mismatch_logged = True
|
||||
|
||||
if isinstance(chunk, dict):
|
||||
chunk["model"] = requested_model_from_client
|
||||
return chunk, model_mismatch_logged
|
||||
|
||||
try:
|
||||
setattr(chunk, "model", requested_model_from_client)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
"litellm_call_id=%s: failed to override chunk.model=%r on chunk_type=%s. error=%s",
|
||||
request_data.get("litellm_call_id"),
|
||||
requested_model_from_client,
|
||||
type(chunk),
|
||||
str(e),
|
||||
exc_info=True,
|
||||
)
|
||||
chunk["model"] = stripped
|
||||
else:
|
||||
try:
|
||||
setattr(chunk, "model", stripped)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(
|
||||
"litellm_call_id=%s: failed to strip provider prefix on chunk.model, error=%s",
|
||||
request_data.get("litellm_call_id"), str(e),
|
||||
)
|
||||
|
||||
return chunk, model_mismatch_logged
|
||||
|
||||
|
|
@ -5463,9 +5446,6 @@ async def async_data_generator(
|
|||
verbose_proxy_logger.debug("inside generator")
|
||||
try:
|
||||
error_message: Optional[str] = None
|
||||
requested_model_from_client = _get_client_requested_model_for_streaming(
|
||||
request_data=request_data
|
||||
)
|
||||
model_mismatch_logged = False
|
||||
# Use a running string instead of list + join to avoid O(n^2) overhead.
|
||||
# Previously "".join(str_so_far_parts) was called every chunk, re-joining
|
||||
|
|
@ -5490,7 +5470,6 @@ async def async_data_generator(
|
|||
|
||||
chunk, model_mismatch_logged = _restamp_streaming_chunk_model(
|
||||
chunk=chunk,
|
||||
requested_model_from_client=requested_model_from_client,
|
||||
request_data=request_data,
|
||||
model_mismatch_logged=model_mismatch_logged,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1148,226 +1148,34 @@ class TestExtractErrorFromSSEChunk:
|
|||
|
||||
|
||||
class TestOverrideOpenAIResponseModel:
|
||||
"""Tests for _override_openai_response_model function"""
|
||||
"""Tests for _override_openai_response_model provider prefix stripping."""
|
||||
|
||||
def test_override_model_preserves_fallback_model_when_fallback_occurred_object(
|
||||
self,
|
||||
):
|
||||
"""
|
||||
Test that when a fallback occurred (x-litellm-attempted-fallbacks > 0),
|
||||
the actual model used (fallback model) is preserved instead of being
|
||||
overridden with the requested model.
|
||||
|
||||
This is the regression test to ensure the model being called is properly
|
||||
displayed when a fallback happens.
|
||||
"""
|
||||
requested_model = "gpt-4"
|
||||
fallback_model = "gpt-3.5-turbo"
|
||||
|
||||
# Create a mock object response with fallback model
|
||||
# _hidden_params is an attribute (not a dict key) accessed via getattr
|
||||
def test_strips_known_provider_prefix_from_object(self):
|
||||
response_obj = MagicMock()
|
||||
response_obj.model = fallback_model
|
||||
response_obj._hidden_params = {
|
||||
"additional_headers": {"x-litellm-attempted-fallbacks": 1}
|
||||
}
|
||||
response_obj.model = "hosted_vllm/my-model"
|
||||
_override_openai_response_model(response_obj=response_obj, log_context="test")
|
||||
assert response_obj.model == "my-model"
|
||||
|
||||
# Call the function - should preserve fallback model
|
||||
_override_openai_response_model(
|
||||
response_obj=response_obj,
|
||||
requested_model=requested_model,
|
||||
log_context="test_context",
|
||||
)
|
||||
def test_strips_known_provider_prefix_from_dict(self):
|
||||
response_obj = {"model": "bedrock/anthropic.claude-v2"}
|
||||
_override_openai_response_model(response_obj=response_obj, log_context="test")
|
||||
assert response_obj["model"] == "anthropic.claude-v2"
|
||||
|
||||
# Verify the model was NOT overridden - should still be the fallback model
|
||||
assert response_obj.model == fallback_model
|
||||
assert response_obj.model != requested_model
|
||||
|
||||
def test_override_model_preserves_fallback_model_multiple_fallbacks(self):
|
||||
"""
|
||||
Test that when multiple fallbacks occurred, the actual model used
|
||||
(fallback model) is preserved.
|
||||
"""
|
||||
requested_model = "gpt-4"
|
||||
fallback_model = "claude-haiku-4-5-20251001"
|
||||
|
||||
# Create a mock object response with fallback model
|
||||
def test_leaves_model_without_prefix_alone(self):
|
||||
response_obj = MagicMock()
|
||||
response_obj.model = fallback_model
|
||||
response_obj._hidden_params = {
|
||||
"additional_headers": {
|
||||
"x-litellm-attempted-fallbacks": 2 # Multiple fallbacks
|
||||
}
|
||||
}
|
||||
response_obj.model = "gpt-4"
|
||||
_override_openai_response_model(response_obj=response_obj, log_context="test")
|
||||
assert response_obj.model == "gpt-4"
|
||||
|
||||
# Call the function - should preserve fallback model
|
||||
_override_openai_response_model(
|
||||
response_obj=response_obj,
|
||||
requested_model=requested_model,
|
||||
log_context="test_context",
|
||||
)
|
||||
|
||||
# Verify the model was NOT overridden - should still be the fallback model
|
||||
assert response_obj.model == fallback_model
|
||||
assert response_obj.model != requested_model
|
||||
|
||||
def test_override_model_overrides_when_no_fallback_dict(self):
|
||||
"""
|
||||
Test that when no fallback occurred, the model is overridden
|
||||
to match the requested model (dict response).
|
||||
"""
|
||||
requested_model = "gpt-4"
|
||||
downstream_model = "gpt-3.5-turbo"
|
||||
|
||||
# Create a dict response without fallback
|
||||
# For dict responses, _hidden_params won't be found via getattr,
|
||||
# so the fallback check won't trigger and model will be overridden
|
||||
response_obj = {"model": downstream_model}
|
||||
|
||||
# Call the function - should override to requested model
|
||||
_override_openai_response_model(
|
||||
response_obj=response_obj,
|
||||
requested_model=requested_model,
|
||||
log_context="test_context",
|
||||
)
|
||||
|
||||
# Verify the model WAS overridden to requested model
|
||||
assert response_obj["model"] == requested_model
|
||||
|
||||
def test_override_model_overrides_when_no_fallback_object(self):
|
||||
"""
|
||||
Test that when no fallback occurred (object response), the model is overridden
|
||||
to match the requested model.
|
||||
"""
|
||||
requested_model = "gpt-4"
|
||||
downstream_model = "gpt-3.5-turbo"
|
||||
|
||||
# Create a mock object response without fallback
|
||||
response_obj = MagicMock()
|
||||
response_obj.model = downstream_model
|
||||
response_obj._hidden_params = {
|
||||
"additional_headers": {} # No attempted_fallbacks header
|
||||
}
|
||||
|
||||
# Call the function - should override to requested model
|
||||
_override_openai_response_model(
|
||||
response_obj=response_obj,
|
||||
requested_model=requested_model,
|
||||
log_context="test_context",
|
||||
)
|
||||
|
||||
# Verify the model WAS overridden to requested model
|
||||
assert response_obj.model == requested_model
|
||||
|
||||
def test_override_model_overrides_when_attempted_fallbacks_is_zero(self):
|
||||
"""
|
||||
Test that when attempted_fallbacks is 0 (no fallback occurred),
|
||||
the model is overridden to match the requested model.
|
||||
"""
|
||||
requested_model = "gpt-4"
|
||||
downstream_model = "gpt-3.5-turbo"
|
||||
|
||||
# Create a mock object response
|
||||
response_obj = MagicMock()
|
||||
response_obj.model = downstream_model
|
||||
response_obj._hidden_params = {
|
||||
"additional_headers": {
|
||||
"x-litellm-attempted-fallbacks": 0 # Zero means no fallback occurred
|
||||
}
|
||||
}
|
||||
|
||||
# Call the function - should override to requested model
|
||||
_override_openai_response_model(
|
||||
response_obj=response_obj,
|
||||
requested_model=requested_model,
|
||||
log_context="test_context",
|
||||
)
|
||||
|
||||
# Verify the model WAS overridden to requested model
|
||||
assert response_obj.model == requested_model
|
||||
|
||||
def test_override_model_overrides_when_attempted_fallbacks_is_none(self):
|
||||
"""
|
||||
Test that when attempted_fallbacks is None (not set),
|
||||
the model is overridden to match the requested model.
|
||||
"""
|
||||
requested_model = "gpt-4"
|
||||
downstream_model = "gpt-3.5-turbo"
|
||||
|
||||
# Create a mock object response
|
||||
response_obj = MagicMock()
|
||||
response_obj.model = downstream_model
|
||||
response_obj._hidden_params = {
|
||||
"additional_headers": {"x-litellm-attempted-fallbacks": None}
|
||||
}
|
||||
|
||||
# Call the function - should override to requested model
|
||||
_override_openai_response_model(
|
||||
response_obj=response_obj,
|
||||
requested_model=requested_model,
|
||||
log_context="test_context",
|
||||
)
|
||||
|
||||
# Verify the model WAS overridden to requested model
|
||||
assert response_obj.model == requested_model
|
||||
|
||||
def test_override_model_no_hidden_params(self):
|
||||
"""
|
||||
Test that when _hidden_params is not present, the model is overridden
|
||||
to match the requested model.
|
||||
"""
|
||||
requested_model = "gpt-4"
|
||||
downstream_model = "gpt-3.5-turbo"
|
||||
|
||||
# Create a mock object response without _hidden_params
|
||||
response_obj = MagicMock()
|
||||
response_obj.model = downstream_model
|
||||
# Don't set _hidden_params - getattr will return {}
|
||||
|
||||
# Call the function - should override to requested model
|
||||
_override_openai_response_model(
|
||||
response_obj=response_obj,
|
||||
requested_model=requested_model,
|
||||
log_context="test_context",
|
||||
)
|
||||
|
||||
# Verify the model WAS overridden to requested model
|
||||
assert response_obj.model == requested_model
|
||||
|
||||
def test_override_model_no_requested_model(self):
|
||||
"""
|
||||
Test that when requested_model is None or empty, the function returns early
|
||||
without modifying the response.
|
||||
"""
|
||||
fallback_model = "gpt-3.5-turbo"
|
||||
|
||||
# Create a mock object response
|
||||
response_obj = MagicMock()
|
||||
response_obj.model = fallback_model
|
||||
response_obj._hidden_params = {
|
||||
"additional_headers": {"x-litellm-attempted-fallbacks": 1}
|
||||
}
|
||||
|
||||
# Call the function with None requested_model
|
||||
_override_openai_response_model(
|
||||
response_obj=response_obj,
|
||||
requested_model=None,
|
||||
log_context="test_context",
|
||||
)
|
||||
|
||||
# Verify the model was not changed
|
||||
assert response_obj.model == fallback_model
|
||||
|
||||
# Call with empty string
|
||||
_override_openai_response_model(
|
||||
response_obj=response_obj,
|
||||
requested_model="",
|
||||
log_context="test_context",
|
||||
)
|
||||
|
||||
# Verify the model was not changed
|
||||
assert response_obj.model == fallback_model
|
||||
def test_leaves_unknown_prefix_alone(self):
|
||||
response_obj = {"model": "my-company/custom-model"}
|
||||
_override_openai_response_model(response_obj=response_obj, log_context="test")
|
||||
assert response_obj["model"] == "my-company/custom-model"
|
||||
|
||||
def test_strips_prefix_preserving_nested_slashes(self):
|
||||
response_obj = {"model": "groq/meta-llama/llama-4-maverick-17b-128e-instruct"}
|
||||
_override_openai_response_model(response_obj=response_obj, log_context="test")
|
||||
assert response_obj["model"] == "meta-llama/llama-4-maverick-17b-128e-instruct"
|
||||
|
||||
class TestStreamingOverheadHeader:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -166,11 +166,7 @@ async def test_proxy_streaming_chunks_do_not_return_provider_prefixed_model(monk
|
|||
@pytest.mark.asyncio
|
||||
async def test_proxy_streaming_chunks_use_client_requested_model_before_alias_mapping(monkeypatch):
|
||||
"""
|
||||
Regression test for alias mapping on streaming:
|
||||
|
||||
- `common_processing_pre_call_logic` can rewrite `request_data["model"]` via model_alias_map / key-specific aliases.
|
||||
- Non-streaming responses are restamped using the original client-requested model (captured before the rewrite).
|
||||
- Streaming chunks must do the same to avoid mismatched `model` values between streaming and non-streaming.
|
||||
Streaming chunks should have provider prefixes stripped even when alias mapping is in play.
|
||||
"""
|
||||
client_model_alias = "alias-model"
|
||||
canonical_model = "vllm-model"
|
||||
|
|
@ -200,7 +196,6 @@ async def test_proxy_streaming_chunks_use_client_requested_model_before_alias_ma
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
request_data={
|
||||
"model": canonical_model,
|
||||
"_litellm_client_requested_model": client_model_alias,
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -213,5 +208,52 @@ async def test_proxy_streaming_chunks_use_client_requested_model_before_alias_ma
|
|||
assert first.startswith("data: ")
|
||||
|
||||
payload = json.loads(first[len("data: ") :].strip())
|
||||
assert payload["model"] == client_model_alias
|
||||
assert payload["model"] == canonical_model
|
||||
assert not payload["model"].startswith("hosted_vllm/")
|
||||
|
||||
|
||||
def test_proxy_chat_completion_returns_actual_model_not_alias(tmp_path, monkeypatch):
|
||||
"""
|
||||
Regression test for GitHub issue #21665:
|
||||
|
||||
Proxy should return actual model name, not the model_list alias.
|
||||
"""
|
||||
alias_model = "default"
|
||||
actual_model = "global.anthropic.claude-sonnet-4-5-20250929-v1:0"
|
||||
internal_model = f"bedrock/{actual_model}"
|
||||
|
||||
client = _initialize_proxy_with_config(
|
||||
config={
|
||||
"general_settings": {"master_key": "sk-1234"},
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": alias_model,
|
||||
"litellm_params": {"model": internal_model},
|
||||
}
|
||||
],
|
||||
},
|
||||
tmp_path=tmp_path,
|
||||
)
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
monkeypatch.setattr(
|
||||
proxy_server.llm_router,
|
||||
"acompletion",
|
||||
AsyncMock(return_value=_make_minimal_chat_completion_response(model=actual_model)),
|
||||
)
|
||||
monkeypatch.setattr(proxy_server.proxy_logging_obj, "during_call_hook", AsyncMock(return_value=None))
|
||||
monkeypatch.setattr(proxy_server.proxy_logging_obj, "update_request_status", AsyncMock(return_value=None))
|
||||
monkeypatch.setattr(proxy_server.proxy_logging_obj, "post_call_success_hook", AsyncMock(side_effect=lambda **kwargs: kwargs["response"]))
|
||||
|
||||
resp = client.post(
|
||||
"/v1/chat/completions",
|
||||
headers={"Authorization": "Bearer sk-1234"},
|
||||
json={"model": alias_model, "messages": [{"role": "user", "content": "hi"}]},
|
||||
)
|
||||
|
||||
assert resp.status_code == 200, resp.text
|
||||
body = resp.json()
|
||||
# Actual model name should be preserved, NOT the alias
|
||||
assert body["model"] == actual_model
|
||||
assert body["model"] != alias_model
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue