mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
fix(proxy): relay Azure passthrough body model groups through the router (#43896)
Co-authored-by: yassin <yassin@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
a6f6c64b6e
commit
fc8f3a26bb
4 changed files with 284 additions and 2 deletions
|
|
@ -64,6 +64,16 @@ def logged_responses_stream(all_chunks: Sequence[str], logging_obj: Logging) ->
|
|||
|
||||
|
||||
AZURE_DEPLOYMENT_SEGMENT: Final = re.compile(r"(?<![^/])openai/deployments/([^/]+)")
|
||||
AZURE_BODY_MODEL_INFERENCE_ENDPOINTS: Final = frozenset(
|
||||
{"responses", "chat/completions", "completions", "embeddings", "images/generations", "audio/speech"}
|
||||
)
|
||||
|
||||
|
||||
def is_azure_body_model_inference_endpoint(endpoint: str) -> bool:
|
||||
if AZURE_DEPLOYMENT_SEGMENT.search(endpoint) is not None:
|
||||
return False
|
||||
path: Final = endpoint.strip("/")
|
||||
return any(path == name or path.endswith(f"/{name}") for name in AZURE_BODY_MODEL_INFERENCE_ENDPOINTS)
|
||||
|
||||
|
||||
def azure_router_model_in_endpoint(endpoint: str, router_models: Collection[str]) -> str | None:
|
||||
|
|
|
|||
|
|
@ -44,7 +44,10 @@ from litellm.constants import (
|
|||
)
|
||||
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
|
||||
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
|
||||
from litellm.llms.azure.passthrough.transformation import foreign_azure_deployment
|
||||
from litellm.llms.azure.passthrough.transformation import (
|
||||
foreign_azure_deployment,
|
||||
is_azure_body_model_inference_endpoint,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.llms.deepgram.common_utils import (
|
||||
deepgram_listen_callback_params,
|
||||
|
|
@ -2128,6 +2131,35 @@ async def relay_nvidia_nim_request(
|
|||
)
|
||||
|
||||
|
||||
async def _relay_azure_body_model_group(
|
||||
llm_router: litellm.Router | None,
|
||||
endpoint: str,
|
||||
request: Request,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> Response | None:
|
||||
if llm_router is None or not is_azure_body_model_inference_endpoint(endpoint):
|
||||
return None
|
||||
if not is_json_content_type(request.headers.get("content-type", "")):
|
||||
return None
|
||||
request_body: Final = await get_request_body(request)
|
||||
model: Final = _optional_str(request_body.get("model"))
|
||||
if model is None or not is_passthrough_request_using_router_model(request_body, llm_router):
|
||||
return None
|
||||
is_streaming_request: Final = is_passthrough_request_streaming(request_body)
|
||||
return await open_sse_before_first_byte(
|
||||
_relay_azure_router_model(
|
||||
llm_router=llm_router,
|
||||
model=model,
|
||||
endpoint=endpoint,
|
||||
request=request,
|
||||
request_body=request_body,
|
||||
is_streaming_request=is_streaming_request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
),
|
||||
ping_interval_seconds=(litellm.sse_keepalive_ping_interval_seconds if is_streaming_request else None),
|
||||
)
|
||||
|
||||
|
||||
@router.api_route(
|
||||
"/azure_ai/{endpoint:path}",
|
||||
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],
|
||||
|
|
@ -2248,6 +2280,12 @@ async def azure_proxy_route(
|
|||
extra_headers=cast(dict, extra_headers),
|
||||
)
|
||||
|
||||
body_model_group_relay: Final = await _relay_azure_body_model_group(
|
||||
llm_router=llm_router, endpoint=endpoint, request=request, user_api_key_dict=user_api_key_dict
|
||||
)
|
||||
if body_model_group_relay is not None:
|
||||
return body_model_group_relay
|
||||
|
||||
base_target_url = get_secret_str(secret_name="AZURE_API_BASE")
|
||||
if base_target_url is None:
|
||||
raise Exception("Required 'AZURE_API_BASE' in environment to make pass-through calls to Azure.")
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ import json
|
|||
import logging
|
||||
import os
|
||||
import traceback
|
||||
from collections.abc import Iterator, Mapping
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable, Iterator, Mapping
|
||||
from types import MappingProxyType, SimpleNamespace
|
||||
from typing import Final
|
||||
from unittest import mock
|
||||
|
|
@ -6440,6 +6440,218 @@ class TestAzureRelayDeploymentSegment:
|
|||
assert [call["model"] for call in captured] == ["gpt", "gpt"]
|
||||
|
||||
|
||||
_AzureRelayUpstream = Callable[[], Awaitable[httpx.Response | AsyncIterator[bytes]]]
|
||||
|
||||
|
||||
async def _azure_relay_json_upstream() -> httpx.Response:
|
||||
return httpx.Response(200, json={"id": "resp_1", "model": "gpt-5.4-fallback"}, headers={"x-request-id": "r-1"})
|
||||
|
||||
|
||||
class _AzureBodyModelGroupRouter:
|
||||
def __init__(self, captured: list[dict], upstream: _AzureRelayUpstream = _azure_relay_json_upstream) -> None:
|
||||
self.captured = captured
|
||||
self.upstream = upstream
|
||||
|
||||
def get_model_names(self, team_id=None):
|
||||
return ["gpt-5.4", "azure-gpt-5.4"]
|
||||
|
||||
def get_model_list(self, model_name=None, team_id=None):
|
||||
rows = [
|
||||
{"model_name": "gpt-5.4", "litellm_params": {"model": "azure/gpt-5.4-primary", "api_key": "k"}},
|
||||
{"model_name": "azure-gpt-5.4", "litellm_params": {"model": "azure/gpt-5.4-fallback", "api_key": "k"}},
|
||||
]
|
||||
return [row for row in rows if model_name is None or row["model_name"] == model_name]
|
||||
|
||||
async def allm_passthrough_route(self, **kwargs):
|
||||
self.captured.append(kwargs)
|
||||
return await self.upstream()
|
||||
|
||||
|
||||
class TestAzureBodyModelGroupRelay:
|
||||
def _install(
|
||||
self,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
body: dict,
|
||||
upstream: _AzureRelayUpstream = _azure_relay_json_upstream,
|
||||
) -> list[dict]:
|
||||
import litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints as ep
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
captured: list[dict] = []
|
||||
|
||||
async def fake_get_request_body(_request: Request) -> dict:
|
||||
return body
|
||||
|
||||
monkeypatch.setattr(proxy_server, "llm_router", _AzureBodyModelGroupRouter(captured, upstream))
|
||||
monkeypatch.setattr(ep, "get_request_body", fake_get_request_body)
|
||||
monkeypatch.delenv("AZURE_API_BASE", raising=False)
|
||||
return captured
|
||||
|
||||
def _request(self, content_type: str = "application/json") -> Request:
|
||||
request = MagicMock(spec=Request)
|
||||
request.method = "POST"
|
||||
request.headers = {"content-type": content_type}
|
||||
request.query_params = {"api-version": "2025-03-01-preview"}
|
||||
return request
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_body_naming_a_model_group_is_relayed_through_the_router(self, monkeypatch):
|
||||
body = {"model": "gpt-5.4", "input": "ping", "max_output_tokens": 16}
|
||||
captured = self._install(monkeypatch, body)
|
||||
|
||||
result = await azure_proxy_route(
|
||||
endpoint="openai/v1/responses",
|
||||
request=self._request(),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-token"),
|
||||
)
|
||||
|
||||
assert result.status_code == 200
|
||||
assert json.loads(result.body) == {"id": "resp_1", "model": "gpt-5.4-fallback"}
|
||||
assert result.headers["x-request-id"] == "r-1"
|
||||
(relay,) = captured
|
||||
assert relay["model"] == "gpt-5.4"
|
||||
assert relay["endpoint"] == "openai/v1/responses"
|
||||
assert relay["json"] == body
|
||||
assert relay["request_query_params"] == {"api-version": "2025-03-01-preview"}
|
||||
assert relay["stream"] is False
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_responses_body_naming_a_model_group_is_relayed_as_a_stream(self, monkeypatch):
|
||||
async def upstream_events() -> AsyncIterator[bytes]:
|
||||
yield b"event: response.created\ndata: {}\n\n"
|
||||
yield b"event: response.completed\ndata: {}\n\n"
|
||||
|
||||
async def streaming_upstream() -> AsyncIterator[bytes]:
|
||||
return upstream_events()
|
||||
|
||||
captured = self._install(monkeypatch, {"model": "gpt-5.4", "input": "ping", "stream": True}, streaming_upstream)
|
||||
|
||||
result = await azure_proxy_route(
|
||||
endpoint="openai/v1/responses",
|
||||
request=self._request(),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-token"),
|
||||
)
|
||||
|
||||
assert isinstance(result, StreamingResponse)
|
||||
streamed = b"".join([chunk async for chunk in result.body_iterator])
|
||||
assert streamed == b"event: response.created\ndata: {}\n\nevent: response.completed\ndata: {}\n\n"
|
||||
(relay,) = captured
|
||||
assert relay["model"] == "gpt-5.4"
|
||||
assert relay["stream"] is True
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_body_naming_no_model_group_still_goes_to_the_operator_azure_endpoint(self, monkeypatch):
|
||||
import litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints as ep
|
||||
|
||||
captured = self._install(monkeypatch, {"model": "gpt-5.4-raw-deployment", "input": "ping"})
|
||||
monkeypatch.setenv("AZURE_API_BASE", "https://operator.openai.azure.com")
|
||||
monkeypatch.setenv("AZURE_API_KEY", "operator-key")
|
||||
routes: list[dict] = []
|
||||
|
||||
def fake_create_pass_through_route(**kwargs):
|
||||
routes.append(kwargs)
|
||||
return AsyncMock(return_value=Response(content=b"{}", status_code=200))
|
||||
|
||||
monkeypatch.setattr(ep, "create_pass_through_route", fake_create_pass_through_route)
|
||||
|
||||
result = await azure_proxy_route(
|
||||
endpoint="openai/v1/responses",
|
||||
request=self._request(),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-token"),
|
||||
)
|
||||
|
||||
assert result.status_code == 200
|
||||
assert captured == []
|
||||
(route,) = routes
|
||||
assert route["target"] == "https://operator.openai.azure.com/openai/v1/responses"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deployment_path_keeps_its_direct_route_even_when_the_body_names_a_model_group(self, monkeypatch):
|
||||
import litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints as ep
|
||||
|
||||
captured = self._install(monkeypatch, {"model": "gpt-5.4", "messages": [{"role": "user", "content": "ping"}]})
|
||||
monkeypatch.setenv("AZURE_API_BASE", "https://operator.openai.azure.com")
|
||||
monkeypatch.setenv("AZURE_API_KEY", "operator-key")
|
||||
routes: list[dict] = []
|
||||
|
||||
def fake_create_pass_through_route(**kwargs):
|
||||
routes.append(kwargs)
|
||||
return AsyncMock(return_value=Response(content=b"{}", status_code=200))
|
||||
|
||||
monkeypatch.setattr(ep, "create_pass_through_route", fake_create_pass_through_route)
|
||||
|
||||
result = await azure_proxy_route(
|
||||
endpoint="openai/deployments/gpt-5.4-raw-deployment/chat/completions",
|
||||
request=self._request(),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-token"),
|
||||
)
|
||||
|
||||
assert result.status_code == 200
|
||||
assert captured == []
|
||||
(route,) = routes
|
||||
assert route["target"] == (
|
||||
"https://operator.openai.azure.com/openai/deployments/gpt-5.4-raw-deployment/chat/completions"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_json_body_is_not_parsed_for_a_model_group(self, monkeypatch):
|
||||
import litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints as ep
|
||||
|
||||
captured = self._install(monkeypatch, {"model": "gpt-5.4", "input": "ping"})
|
||||
monkeypatch.setenv("AZURE_API_BASE", "https://operator.openai.azure.com")
|
||||
monkeypatch.setenv("AZURE_API_KEY", "operator-key")
|
||||
routes: list[dict] = []
|
||||
|
||||
def fake_create_pass_through_route(**kwargs):
|
||||
routes.append(kwargs)
|
||||
return AsyncMock(return_value=Response(content=b"{}", status_code=200))
|
||||
|
||||
monkeypatch.setattr(ep, "create_pass_through_route", fake_create_pass_through_route)
|
||||
|
||||
result = await azure_proxy_route(
|
||||
endpoint="openai/v1/responses",
|
||||
request=self._request(content_type="text/plain"),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-token"),
|
||||
)
|
||||
|
||||
assert result.status_code == 200
|
||||
assert captured == []
|
||||
(route,) = routes
|
||||
assert route["target"] == "https://operator.openai.azure.com/openai/v1/responses"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resource_endpoint_body_naming_a_model_group_keeps_the_operator_account(self, monkeypatch):
|
||||
import litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints as ep
|
||||
|
||||
captured = self._install(monkeypatch, {"model": "gpt-5.4", "training_file": "file-abc123"})
|
||||
monkeypatch.setenv("AZURE_API_BASE", "https://operator.openai.azure.com")
|
||||
monkeypatch.setenv("AZURE_API_KEY", "operator-key")
|
||||
routes: list[dict] = []
|
||||
|
||||
def fake_create_pass_through_route(**kwargs):
|
||||
routes.append(kwargs)
|
||||
return AsyncMock(return_value=Response(content=b"{}", status_code=200))
|
||||
|
||||
monkeypatch.setattr(ep, "create_pass_through_route", fake_create_pass_through_route)
|
||||
|
||||
result = await azure_proxy_route(
|
||||
endpoint="openai/v1/fine_tuning/jobs",
|
||||
request=self._request(),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-token"),
|
||||
)
|
||||
|
||||
assert result.status_code == 200
|
||||
assert captured == []
|
||||
(route,) = routes
|
||||
assert route["target"] == "https://operator.openai.azure.com/openai/v1/fine_tuning/jobs"
|
||||
|
||||
|
||||
AZURE_SPEECH_SHORT_AUDIO_ENDPOINT: Final = "/speech/recognition/conversation/cognitiveservices/v1"
|
||||
AZURE_SPEECH_BATCH_ENDPOINT: Final = "/speechtotext/v3.2/transcriptions"
|
||||
AZURE_SPEECH_FAST_ENDPOINT: Final = "/speechtotext/transcriptions:transcribe"
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from litellm.llms.azure.passthrough.transformation import (
|
|||
AzurePassthroughConfig,
|
||||
azure_router_model_in_endpoint,
|
||||
foreign_azure_deployment,
|
||||
is_azure_body_model_inference_endpoint,
|
||||
)
|
||||
from litellm.types.llms.openai import ResponseCompletedEvent, ResponsesAPIResponse
|
||||
from litellm.types.utils import EmbeddingResponse, ModelResponse
|
||||
|
|
@ -487,3 +488,24 @@ def test_foreign_azure_deployment_skips_the_router_when_the_segment_is_the_group
|
|||
)
|
||||
def test_azure_router_model_in_endpoint_picks_the_first_router_model_segment(endpoint, expected):
|
||||
assert azure_router_model_in_endpoint(endpoint, frozenset({"gpt", "other-group"})) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"endpoint, expected",
|
||||
[
|
||||
("openai/v1/responses", True),
|
||||
("openai/responses", True),
|
||||
("/openai/v1/chat/completions/", True),
|
||||
("openai/v1/embeddings", True),
|
||||
("models/chat/completions", True),
|
||||
("openai/v1/audio/speech", True),
|
||||
("openai/deployments/gpt-5.4/chat/completions", False),
|
||||
("openai/deployments/gpt-5.4/responses", False),
|
||||
("openai/v1/fine_tuning/jobs", False),
|
||||
("openai/v1/assistants", False),
|
||||
("openai/v1/responses/resp_123", False),
|
||||
("openai/v1/batches", False),
|
||||
],
|
||||
)
|
||||
def test_is_azure_body_model_inference_endpoint_admits_only_deployment_less_inference_paths(endpoint, expected):
|
||||
assert is_azure_body_model_inference_endpoint(endpoint) is expected
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue