mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(main.py): handle router custom azure model name for responses api bridge
This commit is contained in:
parent
31d5ebc5a6
commit
dda4532f10
4 changed files with 82 additions and 17 deletions
|
|
@ -822,6 +822,31 @@ def mock_completion(
|
|||
raise Exception("Mock completion response failed - {}".format(e))
|
||||
|
||||
|
||||
def responses_api_bridge_check(
|
||||
model: str,
|
||||
custom_llm_provider: str,
|
||||
) -> dict:
|
||||
model_info = {}
|
||||
try:
|
||||
model_info = _get_model_info_helper(
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
if model_info.get("mode") is None and model.startswith("responses/"):
|
||||
model = model.split("/")[1]
|
||||
mode = "responses"
|
||||
model_info["mode"] = mode
|
||||
except Exception as e:
|
||||
verbose_logger.debug("Error getting model info: {}".format(e))
|
||||
|
||||
if model.startswith(
|
||||
"responses/"
|
||||
): # handle azure models - `azure/responses/<deployment-name>`
|
||||
model = model.split("/")[1]
|
||||
mode = "responses"
|
||||
model_info["mode"] = mode
|
||||
return cast(dict, model_info)
|
||||
|
||||
|
||||
@tracer.wrap()
|
||||
@client
|
||||
def completion( # type: ignore # noqa: PLR0915
|
||||
|
|
@ -1290,19 +1315,9 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
)
|
||||
|
||||
## RESPONSES API BRIDGE LOGIC ## - check if model has 'mode: responses' in litellm.model_cost map
|
||||
try:
|
||||
model_info = _get_model_info_helper(
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.debug("Error getting model info: {}".format(e))
|
||||
model_info = {}
|
||||
if model.startswith(
|
||||
"responses/"
|
||||
): # handle azure models - `azure/responses/<deployment-name>`
|
||||
model = model.split("/")[1]
|
||||
mode = "responses"
|
||||
model_info["mode"] = mode
|
||||
model_info = responses_api_bridge_check(
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
|
||||
if model_info.get("mode") == "responses":
|
||||
from litellm.completion_extras import responses_api_bridge
|
||||
|
|
@ -4937,7 +4952,10 @@ def transcription(
|
|||
provider_config=provider_config,
|
||||
litellm_params=litellm_params_dict,
|
||||
)
|
||||
elif custom_llm_provider in [LlmProviders.DEEPGRAM.value, LlmProviders.ELEVENLABS.value]:
|
||||
elif custom_llm_provider in [
|
||||
LlmProviders.DEEPGRAM.value,
|
||||
LlmProviders.ELEVENLABS.value,
|
||||
]:
|
||||
response = base_llm_http_handler.audio_transcriptions(
|
||||
model=model,
|
||||
audio_file=file,
|
||||
|
|
|
|||
|
|
@ -3,5 +3,13 @@ model_list:
|
|||
litellm_params:
|
||||
model: "gemini/*"
|
||||
api_key: os.environ/GEMINI_API_KEY
|
||||
litellm_settings:
|
||||
check_provider_endpoint: true
|
||||
- model_name: "[IP-approved] o3-pro"
|
||||
litellm_params:
|
||||
model: azure/responses/o_series/webinterface-o3-pro
|
||||
api_base: "https://webhook.site/fba79dae-220a-4bb7-9a3a-8caa49604e55"
|
||||
api_key: "sk-1234567890"
|
||||
api_version: "preview"
|
||||
stream: True
|
||||
model_info:
|
||||
input_cost_per_token: 0.00002 # $20 per 1M tokens
|
||||
output_cost_per_token: 0.00008 # $80 per 1M tokens
|
||||
|
|
@ -162,7 +162,12 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
|
|||
litellm_provider: Required[str]
|
||||
mode: Required[
|
||||
Literal[
|
||||
"completion", "embedding", "image_generation", "chat", "audio_transcription"
|
||||
"completion",
|
||||
"embedding",
|
||||
"image_generation",
|
||||
"chat",
|
||||
"audio_transcription",
|
||||
"responses",
|
||||
]
|
||||
]
|
||||
tpm: Optional[int]
|
||||
|
|
|
|||
|
|
@ -602,3 +602,37 @@ def test_router_should_include_deployment():
|
|||
assert (
|
||||
result is True
|
||||
), "Should return True when matching model with exact model_name"
|
||||
|
||||
|
||||
def test_router_responses_api_bridge():
|
||||
"""
|
||||
Test that router.responses_api_bridge returns the correct response
|
||||
"""
|
||||
import respx
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "[IP-approved] o3-pro",
|
||||
"litellm_params": {
|
||||
"model": "azure/responses/o_series/webinterface-o3-pro",
|
||||
"api_base": "https://webhook.site/fba79dae-220a-4bb7-9a3a-8caa49604e55",
|
||||
"api_key": "sk-1234567890",
|
||||
"api_version": "preview",
|
||||
"stream": True,
|
||||
},
|
||||
"model_info": {
|
||||
"input_cost_per_token": 0.00002,
|
||||
"output_cost_per_token": 0.00008,
|
||||
},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
## CONFIRM BRIDGE IS CALLED
|
||||
with patch.object(litellm, "responses", return_value=AsyncMock()) as mock_responses:
|
||||
result = router.completion(
|
||||
model="[IP-approved] o3-pro",
|
||||
messages=[{"role": "user", "content": "Hello, world!"}],
|
||||
)
|
||||
assert mock_responses.call_count == 1
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue