From d812a5356ea00258d9d1f1c01ebd74570e2cc8d2 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Fri, 1 May 2026 11:44:33 -0700 Subject: [PATCH] fix managed container routing after staging merge --- .../container_endpoints/handler_factory.py | 45 +++++-------------- .../test_azure_container_transformation.py | 27 ++++++++++- .../test_container_proxy_ownership.py | 30 ++++++++----- 3 files changed, 57 insertions(+), 45 deletions(-) diff --git a/litellm/proxy/container_endpoints/handler_factory.py b/litellm/proxy/container_endpoints/handler_factory.py index 5a642e3fc9d..f4887f812cb 100644 --- a/litellm/proxy/container_endpoints/handler_factory.py +++ b/litellm/proxy/container_endpoints/handler_factory.py @@ -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: diff --git a/tests/test_litellm/containers/test_azure_container_transformation.py b/tests/test_litellm/containers/test_azure_container_transformation.py index 45fa23bcb6e..7623c5b0a17 100644 --- a/tests/test_litellm/containers/test_azure_container_transformation.py +++ b/tests/test_litellm/containers/test_azure_container_transformation.py @@ -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" diff --git a/tests/test_litellm/containers/test_container_proxy_ownership.py b/tests/test_litellm/containers/test_container_proxy_ownership.py index 19682649bb9..b696f8fc48f 100644 --- a/tests/test_litellm/containers/test_container_proxy_ownership.py +++ b/tests/test_litellm/containers/test_container_proxy_ownership.py @@ -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"