mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix managed container routing after staging merge
This commit is contained in:
parent
8ced8d2f1f
commit
d812a5356e
3 changed files with 57 additions and 45 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue