fix(main.py): handle router custom azure model name for responses api bridge

This commit is contained in:
Krrish Dholakia 2025-07-01 16:48:57 -07:00
parent 31d5ebc5a6
commit dda4532f10
4 changed files with 82 additions and 17 deletions

View file

@ -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,

View 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

View file

@ -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]

View file

@ -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