fix _add_deployment_model_to_endpoint_for_llm_passthrough_route

This commit is contained in:
Ishaan Jaffer 2025-10-16 15:10:54 -07:00
parent e5956ff0d4
commit cc5eac4965
3 changed files with 93 additions and 3 deletions

View file

@ -2,5 +2,8 @@ model_list:
- model_name: mistral/*
litellm_params:
model: mistral/*
- model_name: special-bedrock-model
litellm_params:
model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0
aws_region_name: us-west-2
custom_llm_provider: bedrock

View file

@ -2753,7 +2753,23 @@ class Router:
it should be actually sent as /model/us.anthropic.claude-3-5-sonnet-20240620-v1:0/invoke
"""
if "endpoint" in kwargs and kwargs["endpoint"]:
kwargs["endpoint"] = kwargs["endpoint"].replace(model, model_name)
# For provider-specific endpoints, strip the provider prefix from model_name
# e.g., "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0" -> "us.anthropic.claude-3-5-sonnet-20240620-v1:0"
from litellm import get_llm_provider
try:
# get_llm_provider returns (model_without_prefix, provider, api_key, api_base)
stripped_model_name, _, _, _ = get_llm_provider(
model=model_name,
custom_llm_provider=kwargs.get("custom_llm_provider"),
api_base=kwargs.get("api_base"),
)
replacement_model_name = stripped_model_name
except Exception:
# If get_llm_provider fails, fall back to using model_name as-is
replacement_model_name = model_name
kwargs["endpoint"] = kwargs["endpoint"].replace(model, replacement_model_name)
return kwargs
async def _ageneric_api_call_with_fallbacks_helper(

View file

@ -1548,3 +1548,74 @@ def test_get_deployment_model_info_base_model_merge_priority():
assert result["key"] == "gpt-4"
print("✓ Base model merge priority test passed!")
def test_add_deployment_model_to_endpoint_for_llm_passthrough_route():
"""
Test that _add_deployment_model_to_endpoint_for_llm_passthrough_route correctly strips bedrock provider prefix
"""
router = litellm.Router(
model_list=[
{
"model_name": "special-bedrock-model",
"litellm_params": {
"model": "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0",
},
}
],
)
# Test Case 1: Bedrock model with provider prefix - should strip "bedrock/" prefix
kwargs = {
"endpoint": "/model/special-bedrock-model/invoke",
"custom_llm_provider": "bedrock",
}
result = router._add_deployment_model_to_endpoint_for_llm_passthrough_route(
kwargs=kwargs,
model="special-bedrock-model",
model_name="bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0",
)
assert (
result["endpoint"] == "/model/us.anthropic.claude-3-5-sonnet-20240620-v1:0/invoke"
), f"Expected '/model/us.anthropic.claude-3-5-sonnet-20240620-v1:0/invoke', got '{result['endpoint']}'"
# Test Case 2: Bedrock invoke-with-response-stream endpoint
kwargs = {
"endpoint": "/model/special-bedrock-model/invoke-with-response-stream",
"custom_llm_provider": "bedrock",
}
result = router._add_deployment_model_to_endpoint_for_llm_passthrough_route(
kwargs=kwargs,
model="special-bedrock-model",
model_name="bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0",
)
assert (
result["endpoint"] == "/model/us.anthropic.claude-3-5-sonnet-20240620-v1:0/invoke-with-response-stream"
), f"Expected streaming endpoint with stripped prefix, got '{result['endpoint']}'"
# Test Case 3: Bedrock converse endpoint
kwargs = {
"endpoint": "/model/bedrock-model/converse",
"custom_llm_provider": "bedrock",
}
result = router._add_deployment_model_to_endpoint_for_llm_passthrough_route(
kwargs=kwargs,
model="bedrock-model",
model_name="bedrock/us.meta.llama3-8b-instruct-v1:0",
)
assert (
result["endpoint"] == "/model/us.meta.llama3-8b-instruct-v1:0/converse"
), f"Expected '/model/us.meta.llama3-8b-instruct-v1:0/converse', got '{result['endpoint']}'"
# Test Case 4: Bedrock provider prefix auto-detected from model_name
kwargs = {
"endpoint": "/model/router-model/invoke",
}
result = router._add_deployment_model_to_endpoint_for_llm_passthrough_route(
kwargs=kwargs,
model="router-model",
model_name="bedrock/us.meta.llama3-8b-instruct-v1:0",
)
assert (
result["endpoint"] == "/model/us.meta.llama3-8b-instruct-v1:0/invoke"
), f"Expected '/model/us.meta.llama3-8b-instruct-v1:0/invoke', got '{result['endpoint']}'"