mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Add watsonx passthrough support (#28352)
* Add support for watsonx passthrough route Signed-off-by: T K Chandra Hasan <t.k.chandra.hasan@ibm.com> * Add UT for watsonx/passthrough/transformation & watsonx api route Vibe coded with IBM Bob Signed-off-by: T K Chandra Hasan <t.k.chandra.hasan@ibm.com> * Address review comments Signed-off-by: T K Chandra Hasan <t.k.chandra.hasan@ibm.com> * Lint & UT fixes Signed-off-by: T K Chandra Hasan <t.k.chandra.hasan@ibm.com> --------- Signed-off-by: T K Chandra Hasan <t.k.chandra.hasan@ibm.com>
This commit is contained in:
parent
d4a89bad1b
commit
c87343ca5f
8 changed files with 889 additions and 0 deletions
|
|
@ -106,6 +106,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
|
|||
# Health & ops
|
||||
"/health",
|
||||
"/metrics",
|
||||
"/watsonx"
|
||||
)
|
||||
|
||||
GATEWAY_EXACT_PATHS: frozenset[str] = frozenset(
|
||||
|
|
|
|||
0
litellm/llms/watsonx/passthrough/__init__.py
Normal file
0
litellm/llms/watsonx/passthrough/__init__.py
Normal file
69
litellm/llms/watsonx/passthrough/transformation.py
Normal file
69
litellm/llms/watsonx/passthrough/transformation.py
Normal file
|
|
@ -0,0 +1,69 @@
|
|||
from typing import TYPE_CHECKING, List, Optional, Tuple
|
||||
|
||||
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
|
||||
from litellm.llms.watsonx.common_utils import IBMWatsonXMixin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from httpx import URL
|
||||
|
||||
|
||||
class WatsonxPassthroughConfig(IBMWatsonXMixin, BasePassthroughConfig):
|
||||
"""
|
||||
Watsonx-specific passthrough configuration.
|
||||
"""
|
||||
|
||||
def is_streaming_request(self, endpoint: str, request_data: dict) -> bool:
|
||||
"""Check if request should be streamed"""
|
||||
return request_data.get("stream", False)
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
endpoint: str,
|
||||
request_query_params: Optional[dict],
|
||||
litellm_params: dict,
|
||||
) -> Tuple["URL", str]:
|
||||
"""
|
||||
Construct complete Watsonx URL with version parameter.
|
||||
|
||||
This ensures the version parameter is ALWAYS included in the URL,
|
||||
solving the query parameter issue.
|
||||
"""
|
||||
base_target_url = str(self.get_api_base(api_base))
|
||||
|
||||
# Use the format_url helper to construct URL with query params
|
||||
complete_url = self.format_url(
|
||||
endpoint=endpoint,
|
||||
base_target_url=base_target_url,
|
||||
request_query_params=request_query_params,
|
||||
)
|
||||
|
||||
return (complete_url, base_target_url)
|
||||
|
||||
@staticmethod
|
||||
def get_api_base(
|
||||
api_base: Optional[str] = None,
|
||||
) -> Optional[str]:
|
||||
return api_base or IBMWatsonXMixin()._get_base_url(api_base=api_base)
|
||||
|
||||
@staticmethod
|
||||
def get_api_key(
|
||||
api_key: Optional[str] = None,
|
||||
) -> Optional[str]:
|
||||
return (
|
||||
api_key
|
||||
or IBMWatsonXMixin.get_watsonx_credentials(
|
||||
optional_params=dict(), api_base=None, api_key=api_key
|
||||
)["api_key"]
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def get_base_model(model: str) -> Optional[str]:
|
||||
return model
|
||||
|
||||
def get_models(
|
||||
self, api_key: Optional[str] = None, api_base: Optional[str] = None
|
||||
) -> List[str]:
|
||||
return super().get_models(api_key, api_base)
|
||||
|
|
@ -419,6 +419,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/vllm",
|
||||
"/mistral",
|
||||
"/milvus",
|
||||
"/watsonx",
|
||||
]
|
||||
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -2433,3 +2433,89 @@ def create_generic_websocket_passthrough_endpoint(
|
|||
_forward_headers=forward_headers,
|
||||
cost_per_request=cost_per_request,
|
||||
)
|
||||
|
||||
|
||||
@router.api_route(
|
||||
"/watsonx/{endpoint:path}",
|
||||
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],
|
||||
tags=["Watsonx Pass-through", "pass-through"],
|
||||
)
|
||||
async def watsonx_proxy_route(
|
||||
endpoint: str,
|
||||
request: Request,
|
||||
fastapi_response: Response,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Watsonx pass-through endpoint.
|
||||
Allows using Watsonx APIs with automatic IAM token management and version parameter injection.
|
||||
|
||||
Example:
|
||||
POST /watsonx/ml/v1/text/tokenization
|
||||
POST /watsonx/ml/v1/text/generation
|
||||
"""
|
||||
# Direct passthrough with WatsonxPassthroughConfig
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
provider_config = ProviderConfigManager.get_provider_passthrough_config(
|
||||
provider=LlmProviders.WATSONX,
|
||||
model="",
|
||||
)
|
||||
|
||||
if provider_config is None:
|
||||
raise HTTPException(
|
||||
status_code=404, detail="Watsonx passthrough config not found"
|
||||
)
|
||||
|
||||
# Get complete URL with version parameter
|
||||
complete_url, _ = provider_config.get_complete_url(
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
model="",
|
||||
endpoint=endpoint,
|
||||
request_query_params=None,
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
# Get auth headers with IAM token
|
||||
auth_headers = provider_config.validate_environment(
|
||||
headers={},
|
||||
model="",
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
)
|
||||
|
||||
# Check for streaming
|
||||
is_streaming_request = False
|
||||
if request.method == "POST":
|
||||
if "multipart/form-data" not in request.headers.get("content-type", ""):
|
||||
_request_body = await request.json()
|
||||
else:
|
||||
_request_body = await get_form_data(request)
|
||||
|
||||
if _request_body.get("stream"):
|
||||
is_streaming_request = True
|
||||
|
||||
request_query_params = dict(request.query_params)
|
||||
if request_query_params.get("version") is None:
|
||||
request_query_params["version"] = litellm.WATSONX_DEFAULT_API_VERSION
|
||||
|
||||
# Create pass-through endpoint
|
||||
endpoint_func = create_pass_through_route(
|
||||
endpoint=endpoint,
|
||||
target=str(complete_url),
|
||||
custom_headers=auth_headers,
|
||||
is_streaming_request=is_streaming_request,
|
||||
custom_llm_provider="watsonx",
|
||||
query_params=request_query_params,
|
||||
)
|
||||
|
||||
return await endpoint_func(
|
||||
request,
|
||||
fastapi_response,
|
||||
user_api_key_dict,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -8936,6 +8936,12 @@ class ProviderConfigManager:
|
|||
)
|
||||
|
||||
return AzurePassthroughConfig()
|
||||
elif LlmProviders.WATSONX == provider:
|
||||
from litellm.llms.watsonx.passthrough.transformation import (
|
||||
WatsonxPassthroughConfig,
|
||||
)
|
||||
|
||||
return WatsonxPassthroughConfig()
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -0,0 +1,282 @@
|
|||
"""
|
||||
Unit tests for WatsonxPassthroughConfig transformation.
|
||||
|
||||
Tests the Watsonx-specific passthrough configuration including URL construction,
|
||||
streaming detection, and authentication handling.
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../../../.."))
|
||||
|
||||
import litellm
|
||||
from litellm.llms.watsonx.passthrough.transformation import WatsonxPassthroughConfig
|
||||
|
||||
|
||||
class TestWatsonxPassthroughConfig:
|
||||
"""Tests for WatsonxPassthroughConfig class."""
|
||||
|
||||
def test_is_streaming_request_true(self):
|
||||
"""Test that streaming is detected when stream=True in request data."""
|
||||
config = WatsonxPassthroughConfig()
|
||||
request_data = {"stream": True, "input": "test"}
|
||||
|
||||
result = config.is_streaming_request(
|
||||
endpoint="ml/v1/text/generation", request_data=request_data
|
||||
)
|
||||
|
||||
assert result is True
|
||||
|
||||
def test_is_streaming_request_false(self):
|
||||
"""Test that streaming is not detected when stream=False in request data."""
|
||||
config = WatsonxPassthroughConfig()
|
||||
request_data = {"stream": False, "input": "test"}
|
||||
|
||||
result = config.is_streaming_request(
|
||||
endpoint="ml/v1/text/generation", request_data=request_data
|
||||
)
|
||||
|
||||
assert result is False
|
||||
|
||||
def test_is_streaming_request_missing_stream_key(self):
|
||||
"""Test that streaming defaults to False when stream key is missing."""
|
||||
config = WatsonxPassthroughConfig()
|
||||
request_data = {"input": "test"}
|
||||
|
||||
result = config.is_streaming_request(
|
||||
endpoint="ml/v1/text/generation", request_data=request_data
|
||||
)
|
||||
|
||||
assert result is False
|
||||
|
||||
def test_get_complete_url_with_api_base(self):
|
||||
"""Test URL construction with explicit api_base."""
|
||||
config = WatsonxPassthroughConfig()
|
||||
api_base = "https://us-south.ml.cloud.ibm.com"
|
||||
endpoint = "ml/v1/text/generation"
|
||||
request_query_params = {"version": "2024-03-19"}
|
||||
|
||||
complete_url, base_target_url = config.get_complete_url(
|
||||
api_base=api_base,
|
||||
api_key=None,
|
||||
model="ibm/granite-13b-chat-v2",
|
||||
endpoint=endpoint,
|
||||
request_query_params=request_query_params,
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert isinstance(complete_url, httpx.URL)
|
||||
assert str(complete_url).startswith(api_base)
|
||||
assert endpoint in str(complete_url)
|
||||
assert "version=2024-03-19" in str(complete_url)
|
||||
assert base_target_url == api_base
|
||||
|
||||
@patch("litellm.llms.watsonx.common_utils.get_secret_str")
|
||||
def test_get_complete_url_with_env_api_base(self, mock_get_secret):
|
||||
"""Test URL construction with api_base from environment."""
|
||||
config = WatsonxPassthroughConfig()
|
||||
env_api_base = "https://eu-de.ml.cloud.ibm.com"
|
||||
mock_get_secret.return_value = env_api_base
|
||||
|
||||
endpoint = "ml/v1/text/tokenization"
|
||||
request_query_params = {"version": "2024-03-19"}
|
||||
|
||||
complete_url, base_target_url = config.get_complete_url(
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
model="ibm/granite-13b-chat-v2",
|
||||
endpoint=endpoint,
|
||||
request_query_params=request_query_params,
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert isinstance(complete_url, httpx.URL)
|
||||
assert str(complete_url).startswith(env_api_base)
|
||||
assert endpoint in str(complete_url)
|
||||
assert base_target_url == env_api_base
|
||||
|
||||
def test_get_complete_url_with_query_params(self):
|
||||
"""Test that query parameters are correctly added to URL."""
|
||||
config = WatsonxPassthroughConfig()
|
||||
api_base = "https://us-south.ml.cloud.ibm.com"
|
||||
endpoint = "ml/v1/text/generation"
|
||||
request_query_params = {
|
||||
"version": "2024-03-19",
|
||||
}
|
||||
|
||||
complete_url, _ = config.get_complete_url(
|
||||
api_base=api_base,
|
||||
api_key=None,
|
||||
model="ibm/granite-13b-chat-v2",
|
||||
endpoint=endpoint,
|
||||
request_query_params=request_query_params,
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
url_str = str(complete_url)
|
||||
assert "version=2024-03-19" in url_str
|
||||
|
||||
def test_get_complete_url_without_query_params(self):
|
||||
"""Test URL construction without query parameters."""
|
||||
config = WatsonxPassthroughConfig()
|
||||
api_base = "https://us-south.ml.cloud.ibm.com"
|
||||
endpoint = "ml/v1/models"
|
||||
|
||||
complete_url, base_target_url = config.get_complete_url(
|
||||
api_base=api_base,
|
||||
api_key=None,
|
||||
model="",
|
||||
endpoint=endpoint,
|
||||
request_query_params=None,
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert isinstance(complete_url, httpx.URL)
|
||||
assert str(complete_url) == f"{api_base}/{endpoint}"
|
||||
assert base_target_url == api_base
|
||||
assert "version=2024-03-19" not in str(complete_url)
|
||||
|
||||
@patch("litellm.llms.watsonx.common_utils.get_secret_str")
|
||||
def test_get_api_base_with_explicit_value(self, mock_get_secret):
|
||||
"""Test get_api_base returns explicit value when provided."""
|
||||
explicit_base = "https://custom.watsonx.com"
|
||||
|
||||
result = WatsonxPassthroughConfig.get_api_base(api_base=explicit_base)
|
||||
|
||||
assert result == explicit_base
|
||||
mock_get_secret.assert_not_called()
|
||||
|
||||
@patch("litellm.llms.watsonx.common_utils.get_secret_str")
|
||||
def test_get_api_base_from_environment(self, mock_get_secret):
|
||||
"""Test get_api_base retrieves from environment when not provided."""
|
||||
env_base = "https://env.watsonx.com"
|
||||
mock_get_secret.return_value = env_base
|
||||
|
||||
result = WatsonxPassthroughConfig.get_api_base(api_base=None)
|
||||
|
||||
assert result == env_base
|
||||
mock_get_secret.assert_called_once_with("WATSONX_API_BASE")
|
||||
|
||||
@patch("litellm.llms.watsonx.common_utils.get_secret_str")
|
||||
def test_get_api_key_with_explicit_value(self, mock_get_secret):
|
||||
"""Test get_api_key returns explicit value when provided."""
|
||||
explicit_key = "test-api-key-123"
|
||||
|
||||
result = WatsonxPassthroughConfig.get_api_key(api_key=explicit_key)
|
||||
|
||||
assert result == explicit_key
|
||||
mock_get_secret.assert_not_called()
|
||||
|
||||
@patch("litellm.llms.watsonx.common_utils.get_secret_str")
|
||||
def test_get_api_key_from_environment(self, mock_get_secret):
|
||||
"""Test get_api_key retrieves from environment when not provided."""
|
||||
env_key = "env-api-key-456"
|
||||
mock_get_secret.return_value = env_key
|
||||
|
||||
result = WatsonxPassthroughConfig.get_api_key(api_key=None)
|
||||
|
||||
assert result == env_key
|
||||
mock_get_secret.assert_any_call("WATSONX_APIKEY")
|
||||
|
||||
def test_get_base_model_returns_model(self):
|
||||
"""Test get_base_model returns the model as-is."""
|
||||
model = "ibm/granite-13b-chat-v2"
|
||||
|
||||
result = WatsonxPassthroughConfig.get_base_model(model)
|
||||
|
||||
assert result == model
|
||||
|
||||
def test_get_base_model_with_deployment(self):
|
||||
"""Test get_base_model with deployment model."""
|
||||
model = "deployment/test-deployment-id"
|
||||
|
||||
result = WatsonxPassthroughConfig.get_base_model(model)
|
||||
|
||||
assert result == model
|
||||
|
||||
def test_get_complete_url_with_different_endpoints(self):
|
||||
"""Test URL construction with various endpoint paths."""
|
||||
config = WatsonxPassthroughConfig()
|
||||
api_base = "https://us-south.ml.cloud.ibm.com"
|
||||
|
||||
endpoints = [
|
||||
"ml/v1/text/generation",
|
||||
"ml/v1/text/tokenization",
|
||||
"ml/v1/deployments/test-id/text/generation",
|
||||
"ml/v1/models",
|
||||
"ml/v1/foundation_model_specs",
|
||||
]
|
||||
|
||||
for endpoint in endpoints:
|
||||
complete_url, base_target_url = config.get_complete_url(
|
||||
api_base=api_base,
|
||||
api_key=None,
|
||||
model="",
|
||||
endpoint=endpoint,
|
||||
request_query_params={"version": "2024-03-19"},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert isinstance(complete_url, httpx.URL)
|
||||
assert endpoint in str(complete_url)
|
||||
assert base_target_url == api_base
|
||||
|
||||
def test_get_complete_url_preserves_query_param_order(self):
|
||||
"""Test that query parameters maintain their values correctly."""
|
||||
config = WatsonxPassthroughConfig()
|
||||
api_base = "https://us-south.ml.cloud.ibm.com"
|
||||
endpoint = "ml/v1/text/generation"
|
||||
request_query_params = {
|
||||
"version": "2024-03-19",
|
||||
"project_id": "abc-123",
|
||||
"space_id": "xyz-789",
|
||||
}
|
||||
|
||||
complete_url, _ = config.get_complete_url(
|
||||
api_base=api_base,
|
||||
api_key=None,
|
||||
model="",
|
||||
endpoint=endpoint,
|
||||
request_query_params=request_query_params,
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
url_str = str(complete_url)
|
||||
# Verify all params are present
|
||||
assert "version=2024-03-19" in url_str
|
||||
assert "project_id=abc-123" in url_str
|
||||
assert "space_id=xyz-789" in url_str
|
||||
|
||||
def test_is_streaming_request_with_various_stream_values(self):
|
||||
"""Test streaming detection with different stream value types."""
|
||||
config = WatsonxPassthroughConfig()
|
||||
|
||||
# Test with boolean True
|
||||
assert config.is_streaming_request("endpoint", {"stream": True}) is True
|
||||
|
||||
# Test with boolean False
|
||||
assert config.is_streaming_request("endpoint", {"stream": False}) is False
|
||||
|
||||
# Test with string "true" (truthy string)
|
||||
result = config.is_streaming_request("endpoint", {"stream": "true"})
|
||||
assert result == "true" # Returns the value as-is from .get()
|
||||
|
||||
# Test with integer 1 (truthy)
|
||||
result = config.is_streaming_request("endpoint", {"stream": 1})
|
||||
assert result == 1
|
||||
|
||||
# Test with integer 0 (falsy)
|
||||
result = config.is_streaming_request("endpoint", {"stream": 0})
|
||||
assert result == 0
|
||||
|
||||
# Test with None
|
||||
result = config.is_streaming_request("endpoint", {"stream": None})
|
||||
assert result is None
|
||||
|
||||
# Test with empty dict (defaults to False)
|
||||
assert config.is_streaming_request("endpoint", {}) is False
|
||||
|
|
@ -0,0 +1,444 @@
|
|||
"""
|
||||
Unit tests for watsonx_proxy_route endpoint.
|
||||
|
||||
Tests the Watsonx pass-through endpoint that handles automatic IAM token management
|
||||
and version parameter injection.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, MagicMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException, Request, Response
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../../../..")
|
||||
) # Adds the parent directory to the system path
|
||||
|
||||
import litellm
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
|
||||
watsonx_proxy_route,
|
||||
)
|
||||
|
||||
|
||||
class TestWatsonxProxyRoute:
|
||||
"""Tests for the Watsonx pass-through route."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_watsonx_proxy_route_success_non_streaming(self):
|
||||
"""Test successful non-streaming request through Watsonx proxy route."""
|
||||
# Setup mocks
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.query_params = {}
|
||||
mock_request.headers = {"content-type": "application/json"}
|
||||
mock_request.json = AsyncMock(return_value={"stream": False, "input": "test"})
|
||||
mock_response = MagicMock(spec=Response)
|
||||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||||
|
||||
# Mock provider config
|
||||
mock_provider_config = MagicMock()
|
||||
mock_provider_config.get_complete_url.return_value = (
|
||||
"https://us-south.ml.cloud.ibm.com/ml/v1/text/generation",
|
||||
{},
|
||||
)
|
||||
mock_provider_config.validate_environment.return_value = {
|
||||
"Authorization": "Bearer test-iam-token"
|
||||
}
|
||||
|
||||
# Mock endpoint function
|
||||
mock_endpoint_func = AsyncMock(
|
||||
return_value={"model_id": "ibm/granite-13b-chat-v2", "results": []}
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_passthrough_config",
|
||||
return_value=mock_provider_config,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route",
|
||||
return_value=mock_endpoint_func,
|
||||
) as mock_create_route,
|
||||
):
|
||||
result = await watsonx_proxy_route(
|
||||
endpoint="ml/v1/text/generation",
|
||||
request=mock_request,
|
||||
fastapi_response=mock_response,
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
)
|
||||
|
||||
# Verify provider config was called correctly
|
||||
mock_provider_config.get_complete_url.assert_called_once()
|
||||
mock_provider_config.validate_environment.assert_called_once()
|
||||
|
||||
# Verify create_pass_through_route was called with correct parameters
|
||||
mock_create_route.assert_called_once()
|
||||
call_args = mock_create_route.call_args[1]
|
||||
assert call_args["endpoint"] == "ml/v1/text/generation"
|
||||
assert (
|
||||
call_args["target"]
|
||||
== "https://us-south.ml.cloud.ibm.com/ml/v1/text/generation"
|
||||
)
|
||||
assert (
|
||||
call_args["custom_headers"]["Authorization"] == "Bearer test-iam-token"
|
||||
)
|
||||
assert call_args["is_streaming_request"] is False
|
||||
assert call_args["custom_llm_provider"] == "watsonx"
|
||||
assert (
|
||||
call_args["query_params"]["version"]
|
||||
== litellm.WATSONX_DEFAULT_API_VERSION
|
||||
)
|
||||
|
||||
# Verify endpoint function was called
|
||||
mock_endpoint_func.assert_called_once_with(
|
||||
mock_request, mock_response, mock_user_api_key_dict
|
||||
)
|
||||
|
||||
assert result == {"model_id": "ibm/granite-13b-chat-v2", "results": []}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_watsonx_proxy_route_success_streaming(self):
|
||||
"""Test successful streaming request through Watsonx proxy route."""
|
||||
# Setup mocks
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.query_params = {}
|
||||
mock_request.headers = {"content-type": "application/json"}
|
||||
mock_request.json = AsyncMock(return_value={"stream": True, "input": "test"})
|
||||
mock_response = MagicMock(spec=Response)
|
||||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||||
|
||||
# Mock provider config
|
||||
mock_provider_config = MagicMock()
|
||||
mock_provider_config.get_complete_url.return_value = (
|
||||
"https://us-south.ml.cloud.ibm.com/ml/v1/text/generation_stream",
|
||||
{},
|
||||
)
|
||||
mock_provider_config.validate_environment.return_value = {
|
||||
"Authorization": "Bearer test-iam-token"
|
||||
}
|
||||
|
||||
# Mock endpoint function
|
||||
mock_endpoint_func = AsyncMock(return_value="streaming_response")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_passthrough_config",
|
||||
return_value=mock_provider_config,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route",
|
||||
return_value=mock_endpoint_func,
|
||||
) as mock_create_route,
|
||||
):
|
||||
result = await watsonx_proxy_route(
|
||||
endpoint="ml/v1/text/generation_stream",
|
||||
request=mock_request,
|
||||
fastapi_response=mock_response,
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
)
|
||||
|
||||
# Verify create_pass_through_route was called with streaming enabled
|
||||
mock_create_route.assert_called_once()
|
||||
call_args = mock_create_route.call_args[1]
|
||||
assert call_args["is_streaming_request"] is True
|
||||
|
||||
assert result == "streaming_response"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_watsonx_proxy_route_get_request(self):
|
||||
"""Test GET request through Watsonx proxy route."""
|
||||
# Setup mocks
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "GET"
|
||||
mock_request.query_params = {"project_id": "test-project"}
|
||||
mock_request.headers = {}
|
||||
mock_response = MagicMock(spec=Response)
|
||||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||||
|
||||
# Mock provider config
|
||||
mock_provider_config = MagicMock()
|
||||
mock_provider_config.get_complete_url.return_value = (
|
||||
"https://us-south.ml.cloud.ibm.com/ml/v1/models",
|
||||
{},
|
||||
)
|
||||
mock_provider_config.validate_environment.return_value = {
|
||||
"Authorization": "Bearer test-iam-token"
|
||||
}
|
||||
|
||||
# Mock endpoint function
|
||||
mock_endpoint_func = AsyncMock(return_value={"resources": []})
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_passthrough_config",
|
||||
return_value=mock_provider_config,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route",
|
||||
return_value=mock_endpoint_func,
|
||||
) as mock_create_route,
|
||||
):
|
||||
result = await watsonx_proxy_route(
|
||||
endpoint="ml/v1/models",
|
||||
request=mock_request,
|
||||
fastapi_response=mock_response,
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
)
|
||||
|
||||
# Verify is_streaming_request is False for GET requests
|
||||
mock_create_route.assert_called_once()
|
||||
call_args = mock_create_route.call_args[1]
|
||||
assert call_args["is_streaming_request"] is False
|
||||
|
||||
assert result == {"resources": []}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_watsonx_proxy_route_multipart_form_data(self):
|
||||
"""Test multipart/form-data request through Watsonx proxy route."""
|
||||
# Setup mocks
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.query_params = {}
|
||||
mock_request.headers = {"content-type": "multipart/form-data; boundary=----"}
|
||||
mock_response = MagicMock(spec=Response)
|
||||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||||
|
||||
# Mock form data
|
||||
mock_form_data = {"file": "test_file", "stream": False}
|
||||
|
||||
# Mock provider config
|
||||
mock_provider_config = MagicMock()
|
||||
mock_provider_config.get_complete_url.return_value = (
|
||||
"https://us-south.ml.cloud.ibm.com/ml/v1/text/tokenization",
|
||||
{},
|
||||
)
|
||||
mock_provider_config.validate_environment.return_value = {
|
||||
"Authorization": "Bearer test-iam-token"
|
||||
}
|
||||
|
||||
# Mock endpoint function
|
||||
mock_endpoint_func = AsyncMock(return_value={"token_count": 10})
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_passthrough_config",
|
||||
return_value=mock_provider_config,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_form_data",
|
||||
return_value=mock_form_data,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route",
|
||||
return_value=mock_endpoint_func,
|
||||
) as mock_create_route,
|
||||
):
|
||||
result = await watsonx_proxy_route(
|
||||
endpoint="ml/v1/text/tokenization",
|
||||
request=mock_request,
|
||||
fastapi_response=mock_response,
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
)
|
||||
|
||||
# Verify is_streaming_request is False for non-streaming form data
|
||||
mock_create_route.assert_called_once()
|
||||
call_args = mock_create_route.call_args[1]
|
||||
assert call_args["is_streaming_request"] is False
|
||||
|
||||
assert result == {"token_count": 10}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_watsonx_proxy_route_no_provider_config(self):
|
||||
"""Test that HTTPException is raised when provider config is not found."""
|
||||
# Setup mocks
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.query_params = {}
|
||||
mock_request.headers = {"content-type": "application/json"}
|
||||
mock_response = MagicMock(spec=Response)
|
||||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_passthrough_config",
|
||||
return_value=None,
|
||||
),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await watsonx_proxy_route(
|
||||
endpoint="ml/v1/text/generation",
|
||||
request=mock_request,
|
||||
fastapi_response=mock_response,
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 404
|
||||
assert exc_info.value.detail == "Watsonx passthrough config not found"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_watsonx_proxy_route_version_parameter_injection(self):
|
||||
"""Test that version parameter is correctly injected into query params."""
|
||||
# Setup mocks
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.query_params = {}
|
||||
mock_request.headers = {"content-type": "application/json"}
|
||||
mock_request.json = AsyncMock(return_value={"input": "test"})
|
||||
mock_response = MagicMock(spec=Response)
|
||||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||||
|
||||
# Mock provider config
|
||||
mock_provider_config = MagicMock()
|
||||
mock_provider_config.get_complete_url.return_value = (
|
||||
"https://us-south.ml.cloud.ibm.com/ml/v1/text/generation",
|
||||
{},
|
||||
)
|
||||
mock_provider_config.validate_environment.return_value = {
|
||||
"Authorization": "Bearer test-iam-token"
|
||||
}
|
||||
|
||||
# Mock endpoint function
|
||||
mock_endpoint_func = AsyncMock(return_value={})
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_passthrough_config",
|
||||
return_value=mock_provider_config,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route",
|
||||
return_value=mock_endpoint_func,
|
||||
) as mock_create_route,
|
||||
):
|
||||
await watsonx_proxy_route(
|
||||
endpoint="ml/v1/text/generation",
|
||||
request=mock_request,
|
||||
fastapi_response=mock_response,
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
)
|
||||
|
||||
# Verify version parameter is injected
|
||||
mock_create_route.assert_called_once()
|
||||
call_args = mock_create_route.call_args[1]
|
||||
assert "query_params" in call_args
|
||||
assert "version" in call_args["query_params"]
|
||||
assert (
|
||||
call_args["query_params"]["version"]
|
||||
== litellm.WATSONX_DEFAULT_API_VERSION
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_watsonx_proxy_route_custom_headers_from_validate_environment(self):
|
||||
"""Test that custom headers from validate_environment are passed through."""
|
||||
# Setup mocks
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.query_params = {}
|
||||
mock_request.headers = {"content-type": "application/json"}
|
||||
mock_request.json = AsyncMock(return_value={"input": "test"})
|
||||
mock_response = MagicMock(spec=Response)
|
||||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||||
|
||||
# Mock provider config with custom headers
|
||||
mock_provider_config = MagicMock()
|
||||
mock_provider_config.get_complete_url.return_value = (
|
||||
"https://us-south.ml.cloud.ibm.com/ml/v1/text/generation",
|
||||
{},
|
||||
)
|
||||
mock_provider_config.validate_environment.return_value = {
|
||||
"Authorization": "Bearer test-iam-token",
|
||||
"X-Custom-Header": "custom-value",
|
||||
}
|
||||
|
||||
# Mock endpoint function
|
||||
mock_endpoint_func = AsyncMock(return_value={})
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_passthrough_config",
|
||||
return_value=mock_provider_config,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route",
|
||||
return_value=mock_endpoint_func,
|
||||
) as mock_create_route,
|
||||
):
|
||||
await watsonx_proxy_route(
|
||||
endpoint="ml/v1/text/generation",
|
||||
request=mock_request,
|
||||
fastapi_response=mock_response,
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
)
|
||||
|
||||
# Verify custom headers are passed through
|
||||
mock_create_route.assert_called_once()
|
||||
call_args = mock_create_route.call_args[1]
|
||||
assert "custom_headers" in call_args
|
||||
assert (
|
||||
call_args["custom_headers"]["Authorization"] == "Bearer test-iam-token"
|
||||
)
|
||||
assert call_args["custom_headers"]["X-Custom-Header"] == "custom-value"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_watsonx_proxy_route_different_endpoints(self):
|
||||
"""Test various Watsonx endpoint paths."""
|
||||
endpoints = [
|
||||
"ml/v1/text/generation",
|
||||
"ml/v1/text/tokenization",
|
||||
"ml/v1/deployments/test-deployment/text/generation",
|
||||
"ml/v1/models",
|
||||
]
|
||||
|
||||
for endpoint_path in endpoints:
|
||||
# Setup mocks
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.query_params = {}
|
||||
mock_request.headers = {"content-type": "application/json"}
|
||||
mock_request.json = AsyncMock(return_value={"input": "test"})
|
||||
mock_response = MagicMock(spec=Response)
|
||||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||||
|
||||
# Mock provider config
|
||||
mock_provider_config = MagicMock()
|
||||
mock_provider_config.get_complete_url.return_value = (
|
||||
f"https://us-south.ml.cloud.ibm.com/{endpoint_path}",
|
||||
{},
|
||||
)
|
||||
mock_provider_config.validate_environment.return_value = {
|
||||
"Authorization": "Bearer test-iam-token"
|
||||
}
|
||||
|
||||
# Mock endpoint function
|
||||
mock_endpoint_func = AsyncMock(return_value={})
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_passthrough_config",
|
||||
return_value=mock_provider_config,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route",
|
||||
return_value=mock_endpoint_func,
|
||||
) as mock_create_route,
|
||||
):
|
||||
await watsonx_proxy_route(
|
||||
endpoint=endpoint_path,
|
||||
request=mock_request,
|
||||
fastapi_response=mock_response,
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
)
|
||||
|
||||
# Verify endpoint is passed correctly
|
||||
mock_create_route.assert_called_once()
|
||||
call_args = mock_create_route.call_args[1]
|
||||
assert call_args["endpoint"] == endpoint_path
|
||||
assert (
|
||||
call_args["target"]
|
||||
== f"https://us-south.ml.cloud.ibm.com/{endpoint_path}"
|
||||
)
|
||||
Loading…
Add table
Reference in a new issue