mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-15 23:31:29 +00:00
fix(router): rewrite the passthrough model group as a whole path segment
This commit is contained in:
parent
bbbdccb82d
commit
3f695846b3
3 changed files with 43 additions and 3 deletions
|
|
@ -8,11 +8,27 @@ if TYPE_CHECKING:
|
|||
from httpx import URL, Headers, Response
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.utils import CostResponseTypes
|
||||
from litellm.types.utils import CostResponseTypes, StandardPassThroughResponseObject
|
||||
|
||||
from ..chat.transformation import BaseLLMException
|
||||
|
||||
|
||||
def strip_leading_model_segment(endpoint: str, model_names: tuple[str, ...]) -> str:
|
||||
path: Final = endpoint.lstrip("/")
|
||||
for model_name in model_names:
|
||||
if not model_name:
|
||||
continue
|
||||
if path == model_name:
|
||||
return ""
|
||||
if path.startswith(f"{model_name}/"):
|
||||
return path[len(model_name) + 1 :]
|
||||
return path
|
||||
|
||||
|
||||
def replace_path_segment(endpoint: str, segment: str, replacement: str) -> str:
|
||||
return "/".join(replacement if part == segment else part for part in endpoint.split("/"))
|
||||
|
||||
|
||||
class BasePassthroughConfig(BaseLLMModelInfo):
|
||||
@abstractmethod
|
||||
def is_streaming_request(self, endpoint: str, request_data: dict) -> bool:
|
||||
|
|
@ -104,7 +120,7 @@ class BasePassthroughConfig(BaseLLMModelInfo):
|
|||
request_data: dict,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
endpoint: str,
|
||||
) -> Optional["CostResponseTypes"]:
|
||||
) -> Optional["CostResponseTypes | StandardPassThroughResponseObject"]:
|
||||
pass
|
||||
|
||||
def handle_logging_collected_chunks(
|
||||
|
|
|
|||
|
|
@ -86,6 +86,7 @@ from litellm.litellm_core_utils.sensitive_data_masker import (
|
|||
mask_credentials_in_payload,
|
||||
mask_sensitive_structure,
|
||||
)
|
||||
from litellm.llms.base_llm.passthrough.transformation import replace_path_segment
|
||||
from litellm.llms.base_llm.vector_store.transformation import (
|
||||
RouterVectorStoreEmbeddingExecutor,
|
||||
vector_store_request_metadata,
|
||||
|
|
@ -5010,7 +5011,7 @@ class Router:
|
|||
# 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)
|
||||
kwargs["endpoint"] = replace_path_segment(kwargs["endpoint"], model, replacement_model_name)
|
||||
return kwargs
|
||||
|
||||
async def _ageneric_api_call_with_fallbacks_helper(self, model: str, original_generic_function: Callable, **kwargs):
|
||||
|
|
|
|||
|
|
@ -4469,6 +4469,29 @@ def test_get_deployment_model_info_base_model_merge_priority():
|
|||
print("✓ Base model merge priority test passed!")
|
||||
|
||||
|
||||
def test_add_deployment_model_to_endpoint_rewrites_whole_path_segments_only():
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt",
|
||||
"litellm_params": {
|
||||
"model": "azure_ai/gpt-5.4-mini",
|
||||
"api_base": "https://my-resource.services.ai.azure.com",
|
||||
"api_key": "key",
|
||||
},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
result = router._add_deployment_model_to_endpoint_for_llm_passthrough_route(
|
||||
kwargs={"endpoint": "gpt/openai/deployments/gpt-4o/chat/completions", "custom_llm_provider": "azure_ai"},
|
||||
model="gpt",
|
||||
model_name="azure_ai/gpt-5.4-mini",
|
||||
)
|
||||
|
||||
assert result["endpoint"] == "gpt-5.4-mini/openai/deployments/gpt-4o/chat/completions"
|
||||
|
||||
|
||||
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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue