fix managed container routing after staging merge

This commit is contained in:
user 2026-05-01 11:44:33 -07:00
parent 8ced8d2f1f
commit d812a5356e
3 changed files with 57 additions and 45 deletions

View file

@ -19,10 +19,7 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
get_custom_llm_provider_from_request_headers,
get_custom_llm_provider_from_request_query,
)
from litellm.proxy.container_endpoints.ownership import (
assert_user_can_access_container,
get_container_forwarding_params,
)
from litellm.proxy.container_endpoints.ownership import assert_user_can_access_container
def _load_endpoints_config() -> Dict:
@ -188,18 +185,15 @@ async def _process_binary_request(
or "openai"
)
original_container_id, custom_llm_provider = await assert_user_can_access_container(
await assert_user_can_access_container(
container_id=container_id,
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
data: Dict[str, Any] = {
"container_id": container_id,
"file_id": file_id,
**get_container_forwarding_params(
container_id,
original_container_id,
custom_llm_provider,
),
"custom_llm_provider": custom_llm_provider,
}
processor = ProxyBaseLLMRequestProcessing(data=data)
@ -308,19 +302,14 @@ async def _process_multipart_upload_request(
or "openai"
)
original_container_id, custom_llm_provider = await assert_user_can_access_container(
await assert_user_can_access_container(
container_id=container_id,
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
data.update(
get_container_forwarding_params(
container_id,
original_container_id,
custom_llm_provider,
)
)
data["container_id"] = container_id
data["custom_llm_provider"] = custom_llm_provider
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
@ -387,22 +376,12 @@ async def _process_request(
# Validate container_id ownership if present in path_params.
if "container_id" in path_params:
original_container_id, custom_llm_provider = (
await assert_user_can_access_container(
container_id=path_params["container_id"],
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
await assert_user_can_access_container(
container_id=path_params["container_id"],
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
data.update(
get_container_forwarding_params(
path_params["container_id"],
original_container_id,
custom_llm_provider,
)
)
else:
data["custom_llm_provider"] = custom_llm_provider
data["custom_llm_provider"] = custom_llm_provider
processor = ProxyBaseLLMRequestProcessing(data=data)
try:

View file

@ -1,6 +1,6 @@
import os
import sys
from unittest.mock import MagicMock
from unittest.mock import AsyncMock, MagicMock
from urllib.parse import parse_qs, urlparse
import httpx
@ -13,7 +13,6 @@ from litellm.llms.azure.containers.transformation import AzureContainerConfig
from litellm.llms.base_llm.containers.transformation import BaseContainerConfig
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.containers.main import (
ContainerFileListResponse,
ContainerListResponse,
ContainerObject,
DeleteContainerResult,
@ -556,6 +555,12 @@ class TestAzureContainerKnownFailureRegressions:
"base_process_llm_request",
_mock_base_process_llm_request,
)
access_check = AsyncMock(return_value=("cntr_123", "azure"))
monkeypatch.setattr(
handler_factory,
"assert_user_can_access_container",
access_check,
)
request = Request(
{
@ -576,6 +581,8 @@ class TestAzureContainerKnownFailureRegressions:
path_params={"container_id": encoded_id},
)
access_check.assert_awaited_once()
assert access_check.await_args.kwargs["container_id"] == encoded_id
assert captured["route_type"] == "alist_container_files"
assert captured["data"]["container_id"] == encoded_id
assert captured["data"]["custom_llm_provider"] == "openai"
@ -620,6 +627,12 @@ class TestAzureContainerKnownFailureRegressions:
"base_process_llm_request",
_mock_base_process_llm_request,
)
access_check = AsyncMock(return_value=("cntr_123", "azure"))
monkeypatch.setattr(
handler_factory,
"assert_user_can_access_container",
access_check,
)
request = Request(
{
@ -640,6 +653,8 @@ class TestAzureContainerKnownFailureRegressions:
user_api_key_dict=MagicMock(),
)
access_check.assert_awaited_once()
assert access_check.await_args.kwargs["container_id"] == encoded_id
assert captured["route_type"] == "aretrieve_container_file_content"
assert captured["data"]["container_id"] == encoded_id
assert captured["data"]["file_id"] == "cfile_abc"
@ -700,6 +715,12 @@ class TestAzureContainerKnownFailureRegressions:
"base_process_llm_request",
_mock_base_process_llm_request,
)
access_check = AsyncMock(return_value=("cntr_123", "azure"))
monkeypatch.setattr(
handler_factory,
"assert_user_can_access_container",
access_check,
)
request = Request(
{
@ -719,6 +740,8 @@ class TestAzureContainerKnownFailureRegressions:
container_id=encoded_id,
)
access_check.assert_awaited_once()
assert access_check.await_args.kwargs["container_id"] == encoded_id
assert captured["route_type"] == "aupload_container_file"
assert captured["data"]["container_id"] == encoded_id
assert captured["data"]["custom_llm_provider"] == "openai"

View file

@ -630,7 +630,9 @@ async def test_should_include_memory_container_list_when_db_recovers_without_row
@pytest.mark.asyncio
async def test_should_forward_decoded_container_id_for_proxy_forwarding(monkeypatch):
async def test_should_validate_owner_and_preserve_managed_id_for_proxy_forwarding(
monkeypatch,
):
from litellm.proxy.container_endpoints import handler_factory
proxy_server_stub = SimpleNamespace(
@ -665,10 +667,11 @@ async def test_should_forward_decoded_container_id_for_proxy_forwarding(monkeypa
"ProxyBaseLLMRequestProcessing",
FakeProcessor,
)
access_check = AsyncMock(return_value=("cntr_provider", "azure"))
monkeypatch.setattr(
handler_factory,
"assert_user_can_access_container",
AsyncMock(return_value=("cntr_provider", "azure")),
access_check,
)
encoded_id = ResponsesAPIRequestUtils._build_container_id(
custom_llm_provider="azure",
@ -684,13 +687,17 @@ async def test_should_forward_decoded_container_id_for_proxy_forwarding(monkeypa
path_params={"container_id": encoded_id},
)
assert result["container_id"] == "cntr_provider"
assert result["custom_llm_provider"] == "azure"
assert result["model_id"] == "router-gpt"
access_check.assert_awaited_once()
assert access_check.await_args.kwargs["container_id"] == encoded_id
assert result["container_id"] == encoded_id
assert result["custom_llm_provider"] == "openai"
assert "model_id" not in result
@pytest.mark.asyncio
async def test_should_forward_decoded_container_id_for_multipart_upload(monkeypatch):
async def test_should_validate_owner_and_preserve_managed_id_for_multipart_upload(
monkeypatch,
):
from litellm.proxy.common_utils import http_parsing_utils
from litellm.proxy.container_endpoints import handler_factory
@ -726,10 +733,11 @@ async def test_should_forward_decoded_container_id_for_multipart_upload(monkeypa
"ProxyBaseLLMRequestProcessing",
FakeProcessor,
)
access_check = AsyncMock(return_value=("cntr_provider", "azure"))
monkeypatch.setattr(
handler_factory,
"assert_user_can_access_container",
AsyncMock(return_value=("cntr_provider", "azure")),
access_check,
)
monkeypatch.setattr(
http_parsing_utils,
@ -755,9 +763,11 @@ async def test_should_forward_decoded_container_id_for_multipart_upload(monkeypa
container_id=encoded_id,
)
assert result["container_id"] == "cntr_provider"
assert result["custom_llm_provider"] == "azure"
assert result["model_id"] == "router-gpt"
access_check.assert_awaited_once()
assert access_check.await_args.kwargs["container_id"] == encoded_id
assert result["container_id"] == encoded_id
assert result["custom_llm_provider"] == "openai"
assert "model_id" not in result
assert result["file"] == "file-data"