diff --git a/litellm/router.py b/litellm/router.py index 88b4087c1e..364e671930 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -386,9 +386,9 @@ class Router: ) # names of models under litellm_params. ex. azure/chatgpt-v-2 self.deployment_latency_map = {} ### CACHING ### - cache_type: Literal["local", "redis", "redis-semantic", "s3", "disk"] = ( - "local" # default to an in-memory cache - ) + cache_type: Literal[ + "local", "redis", "redis-semantic", "s3", "disk" + ] = "local" # default to an in-memory cache redis_cache = None cache_config: Dict[str, Any] = {} @@ -430,9 +430,9 @@ class Router: self.default_max_parallel_requests = default_max_parallel_requests self.provider_default_deployment_ids: List[str] = [] self.pattern_router = PatternMatchRouter() - self.team_pattern_routers: Dict[str, PatternMatchRouter] = ( - {} - ) # {"TEAM_ID": PatternMatchRouter} + self.team_pattern_routers: Dict[ + str, PatternMatchRouter + ] = {} # {"TEAM_ID": PatternMatchRouter} self.auto_routers: Dict[str, "AutoRouter"] = {} # Initialize model_group_alias early since it's used in set_model_list @@ -613,9 +613,9 @@ class Router: ) ) - self.model_group_retry_policy: Optional[Dict[str, RetryPolicy]] = ( - model_group_retry_policy - ) + self.model_group_retry_policy: Optional[ + Dict[str, RetryPolicy] + ] = model_group_retry_policy self.allowed_fails_policy: Optional[AllowedFailsPolicy] = None if allowed_fails_policy is not None: @@ -722,7 +722,10 @@ class Router: valid_strategy_strings = ["simple-shuffle"] + [s.value for s in RoutingStrategy] if routing_strategy is not None: - is_valid_string = isinstance(routing_strategy, str) and routing_strategy in valid_strategy_strings + is_valid_string = ( + isinstance(routing_strategy, str) + and routing_strategy in valid_strategy_strings + ) is_valid_enum = isinstance(routing_strategy, RoutingStrategy) if not is_valid_string and not is_valid_enum: raise ValueError( @@ -1071,7 +1074,7 @@ class Router: self.delete_container = self.factory_function( delete_container, call_type="delete_container" ) - + # Auto-register JSON-generated container file endpoints for name, func in container_file_endpoints.items(): setattr(self, name, self.factory_function(func, call_type=name)) # type: ignore[arg-type] @@ -1500,10 +1503,7 @@ class Router: async def _acompletion( self, model: str, messages: List[Dict[str, str]], **kwargs - ) -> Union[ - ModelResponse, - CustomStreamWrapper, - ]: + ) -> Union[ModelResponse, CustomStreamWrapper,]: """ - Get an available deployment - call it with a semaphore over the call @@ -3021,7 +3021,9 @@ class Router: kwargs["original_generic_function"] = original_function kwargs["original_function"] = self._aguardrail_helper self._update_kwargs_before_fallbacks( - model=guardrail_name, kwargs=kwargs, metadata_variable_name="litellm_metadata" + model=guardrail_name, + kwargs=kwargs, + metadata_variable_name="litellm_metadata", ) verbose_router_logger.debug( f"Inside aguardrail() - guardrail_name: {guardrail_name}; kwargs: {kwargs}" @@ -3314,8 +3316,7 @@ class Router: kwargs["model"] = model kwargs["input"] = input kwargs["original_function"] = self._embedding - kwargs["num_retries"] = kwargs.get("num_retries", self.num_retries) - kwargs.setdefault("metadata", {}).update({"model_group": model}) + self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs) response = self.function_with_fallbacks(**kwargs) return response except Exception as e: @@ -3617,9 +3618,9 @@ class Router: healthy_deployments=healthy_deployments, responses=responses ) returned_response = cast(OpenAIFileObject, responses[0]) - returned_response._hidden_params["model_file_id_mapping"] = ( - model_file_id_mapping - ) + returned_response._hidden_params[ + "model_file_id_mapping" + ] = model_file_id_mapping return returned_response except Exception as e: verbose_router_logger.exception( @@ -4366,11 +4367,11 @@ class Router: if isinstance(e, litellm.ContextWindowExceededError): if context_window_fallbacks is not None: - context_window_fallback_model_group: Optional[List[str]] = ( - self._get_fallback_model_group_from_fallbacks( - fallbacks=context_window_fallbacks, - model_group=model_group, - ) + context_window_fallback_model_group: Optional[ + List[str] + ] = self._get_fallback_model_group_from_fallbacks( + fallbacks=context_window_fallbacks, + model_group=model_group, ) if context_window_fallback_model_group is None: raise original_exception @@ -4402,11 +4403,11 @@ class Router: e.message += "\n{}".format(error_message) elif isinstance(e, litellm.ContentPolicyViolationError): if content_policy_fallbacks is not None: - content_policy_fallback_model_group: Optional[List[str]] = ( - self._get_fallback_model_group_from_fallbacks( - fallbacks=content_policy_fallbacks, - model_group=model_group, - ) + content_policy_fallback_model_group: Optional[ + List[str] + ] = self._get_fallback_model_group_from_fallbacks( + fallbacks=content_policy_fallbacks, + model_group=model_group, ) if content_policy_fallback_model_group is None: raise original_exception @@ -5681,26 +5682,26 @@ class Router: """ from litellm.router_strategy.auto_router.auto_router import AutoRouter - auto_router_config_path: Optional[str] = ( - deployment.litellm_params.auto_router_config_path - ) + auto_router_config_path: Optional[ + str + ] = deployment.litellm_params.auto_router_config_path auto_router_config: Optional[str] = deployment.litellm_params.auto_router_config if auto_router_config_path is None and auto_router_config is None: raise ValueError( "auto_router_config_path or auto_router_config is required for auto-router deployments. Please set it in the litellm_params" ) - default_model: Optional[str] = ( - deployment.litellm_params.auto_router_default_model - ) + default_model: Optional[ + str + ] = deployment.litellm_params.auto_router_default_model if default_model is None: raise ValueError( "auto_router_default_model is required for auto-router deployments. Please set it in the litellm_params" ) - embedding_model: Optional[str] = ( - deployment.litellm_params.auto_router_embedding_model - ) + embedding_model: Optional[ + str + ] = deployment.litellm_params.auto_router_embedding_model if embedding_model is None: raise ValueError( "auto_router_embedding_model is required for auto-router deployments. Please set it in the litellm_params" @@ -6247,9 +6248,9 @@ class Router: # Add custom_llm_provider if deployment.litellm_params.custom_llm_provider: - credentials["custom_llm_provider"] = ( - deployment.litellm_params.custom_llm_provider - ) + credentials[ + "custom_llm_provider" + ] = deployment.litellm_params.custom_llm_provider elif "/" in deployment.litellm_params.model: # Extract provider from "provider/model" format credentials["custom_llm_provider"] = deployment.litellm_params.model.split( @@ -6943,42 +6944,44 @@ class Router: """ return candidate_id in self.model_id_to_deployment_index_map - def resolve_model_name_from_model_id(self, model_id: Optional[str]) -> Optional[str]: + def resolve_model_name_from_model_id( + self, model_id: Optional[str] + ) -> Optional[str]: """ Resolve model_name from model_id. - + This method attempts to find the correct model_name to use with the router so that litellm_params can be automatically injected from the model config. - + Strategy: 1. First, check if model_id directly matches a model_name or deployment ID 2. If not, search through router's model_list to find a match by litellm_params.model 3. Return the model_name if found, None otherwise - + Args: model_id: The model_id extracted from decoded video_id (could be model_name or litellm_params.model value) - + Returns: model_name if found, None otherwise. If None, the request will fall through to normal flow using environment variables. """ if not model_id: return None - + # Strategy 1: Check if model_id directly matches a model_name or deployment ID if model_id in self.model_names or self.has_model_id(model_id): return model_id - + # Strategy 2: Search through router's model_list to find by litellm_params.model all_models = self.get_model_list(model_name=None) if not all_models: return None - + for deployment in all_models: litellm_params = deployment.get("litellm_params", {}) actual_model = litellm_params.get("model") - + # Match by exact match or by checking if actual_model ends with /model_id or :model_id # e.g., model_id="veo-2.0-generate-001" matches actual_model="vertex_ai/veo-2.0-generate-001" matches = ( @@ -6986,12 +6989,12 @@ class Router: or (actual_model and actual_model.endswith(f"/{model_id}")) or (actual_model and actual_model.endswith(f":{model_id}")) ) - + if matches: model_name = deployment.get("model_name") if model_name: return model_name - + # No match found return None @@ -7785,14 +7788,18 @@ class Router: request_kwargs=request_kwargs, ) - verbose_router_logger.debug(f"healthy_deployments after team filter: {healthy_deployments}") + verbose_router_logger.debug( + f"healthy_deployments after team filter: {healthy_deployments}" + ) healthy_deployments = filter_web_search_deployments( healthy_deployments=healthy_deployments, request_kwargs=request_kwargs, ) - verbose_router_logger.debug(f"healthy_deployments after web search filter: {healthy_deployments}") + verbose_router_logger.debug( + f"healthy_deployments after web search filter: {healthy_deployments}" + ) if isinstance(healthy_deployments, dict): return healthy_deployments diff --git a/tests/router_unit_tests/test_router_embedding_headers.py b/tests/router_unit_tests/test_router_embedding_headers.py new file mode 100644 index 0000000000..6d480792b7 --- /dev/null +++ b/tests/router_unit_tests/test_router_embedding_headers.py @@ -0,0 +1,372 @@ +""" +Test suite for router embedding method header propagation. + +This tests the fix for the issue where the embedding method was not +propagating proxy model configuration headers to the LLM API calls. + +The fix ensures that router.embedding() calls _update_kwargs_before_fallbacks() +just like router.completion() does, which properly sets up metadata and allows +default_litellm_params (including headers) to be propagated. +""" +import os +import sys +from unittest.mock import MagicMock, patch, AsyncMock + +import pytest + +sys.path.insert(0, os.path.abspath("../..")) + +from litellm import Router + + +class TestRouterEmbeddingHeaders: + """Test that embedding methods properly propagate headers from router configuration.""" + + def test_embedding_calls_update_kwargs_before_fallbacks(self): + """ + Test that router.embedding() calls _update_kwargs_before_fallbacks. + + This ensures that metadata is properly set up before the fallback mechanism, + which is necessary for header propagation to work correctly. + """ + model_list = [ + { + "model_name": "text-embedding-ada-002", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "fake-key", + }, + } + ] + + router = Router(model_list=model_list) + + # Mock the _update_kwargs_before_fallbacks method to verify it's called + with patch.object( + router, + "_update_kwargs_before_fallbacks", + wraps=router._update_kwargs_before_fallbacks, + ) as mock_update: + with patch("litellm.embedding") as mock_litellm_embedding: + mock_litellm_embedding.return_value = MagicMock( + data=[{"embedding": [0.1, 0.2, 0.3]}] + ) + + router.embedding(model="text-embedding-ada-002", input=["test input"]) + + # Verify _update_kwargs_before_fallbacks was called + mock_update.assert_called_once() + call_kwargs = mock_update.call_args[1] + assert call_kwargs["model"] == "text-embedding-ada-002" + assert "kwargs" in call_kwargs + + @pytest.mark.asyncio + async def test_aembedding_calls_update_kwargs_before_fallbacks(self): + """ + Test that router.aembedding() calls _update_kwargs_before_fallbacks. + + This ensures consistency between sync and async embedding methods. + """ + model_list = [ + { + "model_name": "text-embedding-ada-002", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "fake-key", + }, + } + ] + + router = Router(model_list=model_list) + + # Mock the _update_kwargs_before_fallbacks method to verify it's called + with patch.object( + router, + "_update_kwargs_before_fallbacks", + wraps=router._update_kwargs_before_fallbacks, + ) as mock_update: + with patch( + "litellm.aembedding", new_callable=AsyncMock + ) as mock_litellm_aembedding: + mock_litellm_aembedding.return_value = MagicMock( + data=[{"embedding": [0.1, 0.2, 0.3]}] + ) + + await router.aembedding( + model="text-embedding-ada-002", input=["test input"] + ) + + # Verify _update_kwargs_before_fallbacks was called + mock_update.assert_called_once() + call_kwargs = mock_update.call_args[1] + assert call_kwargs["model"] == "text-embedding-ada-002" + assert "kwargs" in call_kwargs + + def test_embedding_propagates_default_litellm_params(self): + """ + Test that embedding calls properly propagate default_litellm_params including headers. + + This is the main fix - ensuring that headers set in default_litellm_params + are included in the embedding request. + """ + custom_headers = {"X-Custom-Header": "test-value", "X-API-Version": "v2"} + + model_list = [ + { + "model_name": "text-embedding-ada-002", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "fake-key", + }, + } + ] + + # Create router with default_litellm_params containing headers + router = Router( + model_list=model_list, + default_litellm_params={ + "headers": custom_headers, + "metadata": {"test_key": "test_value"}, + }, + ) + + with patch("litellm.embedding") as mock_litellm_embedding: + mock_litellm_embedding.return_value = MagicMock( + data=[{"embedding": [0.1, 0.2, 0.3]}] + ) + + router.embedding(model="text-embedding-ada-002", input=["test input"]) + + # Verify that litellm.embedding was called with the headers + mock_litellm_embedding.assert_called_once() + call_kwargs = mock_litellm_embedding.call_args[1] + + # Check that headers were included + assert "headers" in call_kwargs + assert call_kwargs["headers"] == custom_headers + + # Check that metadata was properly set up + assert "metadata" in call_kwargs + assert "model_group" in call_kwargs["metadata"] + assert call_kwargs["metadata"]["model_group"] == "text-embedding-ada-002" + + @pytest.mark.asyncio + async def test_aembedding_propagates_default_litellm_params(self): + """ + Test that async embedding calls properly propagate default_litellm_params including headers. + """ + custom_headers = {"X-Custom-Header": "test-value", "X-API-Version": "v2"} + + model_list = [ + { + "model_name": "text-embedding-ada-002", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "fake-key", + }, + } + ] + + # Create router with default_litellm_params containing headers + router = Router( + model_list=model_list, + default_litellm_params={ + "headers": custom_headers, + "metadata": {"test_key": "test_value"}, + }, + ) + + with patch( + "litellm.aembedding", new_callable=AsyncMock + ) as mock_litellm_aembedding: + mock_litellm_aembedding.return_value = MagicMock( + data=[{"embedding": [0.1, 0.2, 0.3]}] + ) + + await router.aembedding( + model="text-embedding-ada-002", input=["test input"] + ) + + # Verify that litellm.aembedding was called with the headers + mock_litellm_aembedding.assert_called_once() + call_kwargs = mock_litellm_aembedding.call_args[1] + + # Check that headers were included + assert "headers" in call_kwargs + assert call_kwargs["headers"] == custom_headers + + # Check that metadata was properly set up + assert "metadata" in call_kwargs + assert "model_group" in call_kwargs["metadata"] + assert call_kwargs["metadata"]["model_group"] == "text-embedding-ada-002" + + def test_embedding_metadata_includes_model_group(self): + """ + Test that embedding calls include model_group in metadata. + + The _update_kwargs_before_fallbacks method should set this up. + """ + model_list = [ + { + "model_name": "test-embedding-model", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "fake-key", + }, + } + ] + + router = Router(model_list=model_list) + + with patch("litellm.embedding") as mock_litellm_embedding: + mock_litellm_embedding.return_value = MagicMock( + data=[{"embedding": [0.1, 0.2, 0.3]}] + ) + + router.embedding(model="test-embedding-model", input=["test input"]) + + call_kwargs = mock_litellm_embedding.call_args[1] + + # Verify metadata contains model_group + assert "metadata" in call_kwargs + assert "model_group" in call_kwargs["metadata"] + assert call_kwargs["metadata"]["model_group"] == "test-embedding-model" + + def test_embedding_sets_num_retries_from_router(self): + """ + Test that embedding calls inherit num_retries from router configuration. + + This is set by _update_kwargs_before_fallbacks. + """ + model_list = [ + { + "model_name": "text-embedding-ada-002", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "fake-key", + }, + } + ] + + # Create router with num_retries set + router = Router(model_list=model_list, num_retries=3) + + with patch("litellm.embedding") as mock_litellm_embedding: + mock_litellm_embedding.return_value = MagicMock( + data=[{"embedding": [0.1, 0.2, 0.3]}] + ) + + router.embedding(model="text-embedding-ada-002", input=["test input"]) + + # Verify num_retries was not set in the call (it's handled by function_with_fallbacks) + # The important thing is that it was set in kwargs before being passed to function_with_fallbacks + # We verify this indirectly by checking that _update_kwargs_before_fallbacks was called + mock_litellm_embedding.assert_called_once() + + def test_embedding_sets_litellm_trace_id(self): + """ + Test that embedding calls include a litellm_trace_id. + + This is generated and set by _update_kwargs_before_fallbacks. + """ + model_list = [ + { + "model_name": "text-embedding-ada-002", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "fake-key", + }, + } + ] + + router = Router(model_list=model_list) + + with patch("litellm.embedding") as mock_litellm_embedding: + mock_litellm_embedding.return_value = MagicMock( + data=[{"embedding": [0.1, 0.2, 0.3]}] + ) + + router.embedding(model="text-embedding-ada-002", input=["test input"]) + + call_kwargs = mock_litellm_embedding.call_args[1] + + # Verify litellm_trace_id was set + assert "litellm_trace_id" in call_kwargs + assert isinstance(call_kwargs["litellm_trace_id"], str) + assert len(call_kwargs["litellm_trace_id"]) > 0 + + def test_embedding_consistency_with_completion(self): + """ + Test that embedding and completion methods handle kwargs similarly. + + Both should call _update_kwargs_before_fallbacks to ensure consistent behavior. + """ + custom_headers = {"X-Test": "value"} + + model_list = [ + { + "model_name": "gpt-3.5-turbo", + "litellm_params": { + "model": "gpt-3.5-turbo", + "api_key": "fake-key", + }, + }, + { + "model_name": "text-embedding-ada-002", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "fake-key", + }, + }, + ] + + router = Router( + model_list=model_list, default_litellm_params={"headers": custom_headers} + ) + + # Test completion + with patch("litellm.completion") as mock_completion: + mock_completion.return_value = MagicMock() + + router.completion( + model="gpt-3.5-turbo", messages=[{"role": "user", "content": "test"}] + ) + + completion_kwargs = mock_completion.call_args[1] + + # Test embedding + with patch("litellm.embedding") as mock_embedding: + mock_embedding.return_value = MagicMock( + data=[{"embedding": [0.1, 0.2, 0.3]}] + ) + + router.embedding(model="text-embedding-ada-002", input=["test input"]) + + embedding_kwargs = mock_embedding.call_args[1] + + # Both should have headers from default_litellm_params + assert "headers" in completion_kwargs + assert "headers" in embedding_kwargs + assert completion_kwargs["headers"] == custom_headers + assert embedding_kwargs["headers"] == custom_headers + + # Both should have metadata with model_group + assert "metadata" in completion_kwargs + assert "metadata" in embedding_kwargs + assert "model_group" in completion_kwargs["metadata"] + assert "model_group" in embedding_kwargs["metadata"] + + # Both should have litellm_trace_id + assert "litellm_trace_id" in completion_kwargs + assert "litellm_trace_id" in embedding_kwargs + + +if __name__ == "__main__": + # Run a simple test + test = TestRouterEmbeddingHeaders() + test.test_embedding_calls_update_kwargs_before_fallbacks() + test.test_embedding_propagates_default_litellm_params() + test.test_embedding_metadata_includes_model_group() + test.test_embedding_sets_litellm_trace_id() + test.test_embedding_consistency_with_completion() + print("All tests passed!") # noqa: T201 diff --git a/tests/router_unit_tests/test_router_embedding_integration.py b/tests/router_unit_tests/test_router_embedding_integration.py new file mode 100644 index 0000000000..ab2071714a --- /dev/null +++ b/tests/router_unit_tests/test_router_embedding_integration.py @@ -0,0 +1,355 @@ +""" +Integration tests for router embedding method with various configurations. + +These tests simulate real-world scenarios where headers and configuration +need to be properly propagated through the router to the LLM API. +""" +import os +import sys +from unittest.mock import MagicMock, patch, AsyncMock + +import pytest + +sys.path.insert(0, os.path.abspath("../..")) + +from litellm import Router + + +class TestRouterEmbeddingIntegration: + """Integration tests for embedding with router configuration.""" + + def test_embedding_with_deployment_specific_headers(self): + """ + Test that deployment-specific headers are propagated. + + This simulates a scenario where different deployments have + different header requirements (e.g., different API versions). + """ + model_list = [ + { + "model_name": "embedding-deployment-1", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "key-1", + "headers": {"X-Deployment": "deployment-1"}, + }, + }, + { + "model_name": "embedding-deployment-2", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "key-2", + "headers": {"X-Deployment": "deployment-2"}, + }, + }, + ] + + router = Router(model_list=model_list) + + # Test first deployment + with patch("litellm.embedding") as mock_embedding: + mock_embedding.return_value = MagicMock(data=[{"embedding": [0.1, 0.2]}]) + + router.embedding(model="embedding-deployment-1", input=["test"]) + + call_kwargs = mock_embedding.call_args[1] + assert call_kwargs["api_key"] == "key-1" + + # Test second deployment + with patch("litellm.embedding") as mock_embedding: + mock_embedding.return_value = MagicMock(data=[{"embedding": [0.1, 0.2]}]) + + router.embedding(model="embedding-deployment-2", input=["test"]) + + call_kwargs = mock_embedding.call_args[1] + assert call_kwargs["api_key"] == "key-2" + + def test_embedding_with_router_and_deployment_headers_merge(self): + """ + Test that router-level headers are propagated. + + When no request headers are provided, router default headers should be used. + """ + model_list = [ + { + "model_name": "test-embedding", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "test-key", + }, + } + ] + + router = Router( + model_list=model_list, + default_litellm_params={ + "headers": { + "X-Router-Header": "router-value", + "X-Common-Header": "router-common", + } + }, + ) + + # Test: No request headers - router headers should be used + with patch("litellm.embedding") as mock_embedding: + mock_embedding.return_value = MagicMock(data=[{"embedding": [0.1, 0.2]}]) + + router.embedding( + model="test-embedding", + input=["test"], + ) + + call_kwargs = mock_embedding.call_args[1] + + # Router headers should be present + assert "headers" in call_kwargs + assert call_kwargs["headers"]["X-Router-Header"] == "router-value" + assert call_kwargs["headers"]["X-Common-Header"] == "router-common" + + def test_embedding_metadata_propagation(self): + """ + Test that metadata is properly set up and propagated. + + This is important for logging, tracking, and debugging. + """ + model_list = [ + { + "model_name": "test-embedding", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "test-key", + }, + } + ] + + router = Router( + model_list=model_list, + default_litellm_params={ + "metadata": {"environment": "test", "service": "embedding-service"} + }, + ) + + with patch("litellm.embedding") as mock_embedding: + mock_embedding.return_value = MagicMock(data=[{"embedding": [0.1, 0.2]}]) + + router.embedding( + model="test-embedding", + input=["test"], + metadata={"request_id": "req-123"}, # Additional metadata from request + ) + + call_kwargs = mock_embedding.call_args[1] + + # Check metadata contains all expected fields + assert "metadata" in call_kwargs + metadata = call_kwargs["metadata"] + + # From _update_kwargs_before_fallbacks + assert "model_group" in metadata + assert metadata["model_group"] == "test-embedding" + + # From default_litellm_params + assert "environment" in metadata + assert metadata["environment"] == "test" + assert "service" in metadata + assert metadata["service"] == "embedding-service" + + # From request + assert "request_id" in metadata + assert metadata["request_id"] == "req-123" + + @pytest.mark.asyncio + async def test_async_embedding_with_multiple_retries(self): + """ + Test that async embedding properly uses num_retries from router config. + + This ensures the fix works with the retry mechanism. + """ + model_list = [ + { + "model_name": "test-embedding", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "test-key", + }, + } + ] + + router = Router(model_list=model_list, num_retries=2) + + with patch("litellm.aembedding", new_callable=AsyncMock) as mock_aembedding: + mock_aembedding.return_value = MagicMock(data=[{"embedding": [0.1, 0.2]}]) + + await router.aembedding(model="test-embedding", input=["test"]) + + # The call should succeed + mock_aembedding.assert_called_once() + + def test_embedding_with_timeout_from_router(self): + """ + Test that timeout settings from router config are propagated. + """ + model_list = [ + { + "model_name": "test-embedding", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "test-key", + }, + } + ] + + router = Router(model_list=model_list, timeout=30.0) + + with patch("litellm.embedding") as mock_embedding: + mock_embedding.return_value = MagicMock(data=[{"embedding": [0.1, 0.2]}]) + + router.embedding(model="test-embedding", input=["test"]) + + call_kwargs = mock_embedding.call_args[1] + + # Timeout should be set from router config + assert "timeout" in call_kwargs + assert call_kwargs["timeout"] == 30.0 + + def test_embedding_with_multiple_deployments_load_balancing(self): + """ + Test that headers are correctly propagated when router load balances + between multiple deployments. + """ + model_list = [ + { + "model_name": "shared-embedding-model", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "key-1", + }, + }, + { + "model_name": "shared-embedding-model", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "key-2", + }, + }, + ] + + router = Router( + model_list=model_list, + default_litellm_params={"headers": {"X-Shared-Header": "shared-value"}}, + ) + + # Make multiple calls and verify headers are always present + for i in range(5): + with patch("litellm.embedding") as mock_embedding: + mock_embedding.return_value = MagicMock( + data=[{"embedding": [0.1, 0.2]}] + ) + + router.embedding(model="shared-embedding-model", input=[f"test {i}"]) + + call_kwargs = mock_embedding.call_args[1] + + # Headers should always be present regardless of which deployment is chosen + assert "headers" in call_kwargs + assert call_kwargs["headers"]["X-Shared-Header"] == "shared-value" + + @pytest.mark.asyncio + async def test_embedding_with_fallback_configuration(self): + """ + Test that headers are propagated correctly when using fallback models. + """ + model_list = [ + { + "model_name": "primary-embedding", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "primary-key", + }, + }, + { + "model_name": "fallback-embedding", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "fallback-key", + }, + }, + ] + + router = Router( + model_list=model_list, + fallbacks=[{"primary-embedding": ["fallback-embedding"]}], + default_litellm_params={"headers": {"X-Fallback-Test": "test-value"}}, + ) + + # Simulate primary failing, fallback succeeding + with patch("litellm.aembedding", new_callable=AsyncMock) as mock_aembedding: + call_count = 0 + + async def side_effect(*args, **kwargs): + nonlocal call_count + call_count += 1 + if call_count == 1: + # First call (primary) fails + raise Exception("Primary failed") + else: + # Second call (fallback) succeeds + return MagicMock(data=[{"embedding": [0.1, 0.2]}]) + + mock_aembedding.side_effect = side_effect + + await router.aembedding(model="primary-embedding", input=["test"]) + + # Both calls should have headers + assert mock_aembedding.call_count == 2 + + # Check that both calls had headers + for call_obj in mock_aembedding.call_args_list: + call_kwargs = call_obj[1] + assert "headers" in call_kwargs + assert call_kwargs["headers"]["X-Fallback-Test"] == "test-value" + + def test_embedding_with_custom_provider_headers(self): + """ + Test that provider-specific headers are correctly propagated. + + Some providers require specific headers for API versioning, features, etc. + """ + model_list = [ + { + "model_name": "azure-embedding", + "litellm_params": { + "model": "azure/text-embedding-ada-002", + "api_key": "azure-key", + "api_base": "https://example.openai.azure.com", + "api_version": "2024-02-01", + }, + } + ] + + router = Router( + model_list=model_list, + default_litellm_params={ + "headers": {"X-Custom-Azure-Header": "azure-value"} + }, + ) + + with patch("litellm.embedding") as mock_embedding: + mock_embedding.return_value = MagicMock(data=[{"embedding": [0.1, 0.2]}]) + + router.embedding(model="azure-embedding", input=["test"]) + + call_kwargs = mock_embedding.call_args[1] + + # Verify Azure-specific params are present + assert call_kwargs["api_base"] == "https://example.openai.azure.com" + assert call_kwargs["api_version"] == "2024-02-01" + + # Verify custom headers are present + assert "headers" in call_kwargs + assert call_kwargs["headers"]["X-Custom-Azure-Header"] == "azure-value" + + +if __name__ == "__main__": + # Run tests + pytest.main([__file__, "-v"])