diff --git a/litellm/router.py b/litellm/router.py index 43b53d14d79..19273660795 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -727,15 +727,11 @@ class Router: startup_nodes = cache_config.get("startup_nodes") if not startup_nodes: _env_cluster_nodes = get_secret("REDIS_CLUSTER_NODES") - if _env_cluster_nodes is not None and isinstance( - _env_cluster_nodes, str - ): + if _env_cluster_nodes is not None and isinstance(_env_cluster_nodes, str): startup_nodes = json.loads(_env_cluster_nodes) if startup_nodes: - return RedisClusterCache( - **{**cache_config, "startup_nodes": startup_nodes} - ) + return RedisClusterCache(**{**cache_config, "startup_nodes": startup_nodes}) else: return RedisCache(**cache_config) @@ -1473,6 +1469,9 @@ class Router: silent_kwargs.pop("standard_logging_object", None) silent_kwargs.pop("proxy_server_request", None) + silent_kwargs["stream"] = False + silent_kwargs["num_retries"] = 0 + return silent_kwargs def _silent_experiment_completion( @@ -1493,13 +1492,26 @@ class Router: ) silent_kwargs = self._get_silent_experiment_kwargs(**kwargs) + if ( + "metadata" in silent_kwargs + and "model_group" in silent_kwargs["metadata"] + ): + silent_kwargs["metadata"]["model_group"] = silent_model # Trigger the silent request - self.completion( - model=silent_model, - messages=cast(List[Dict[str, str]], messages), - **silent_kwargs, - ) + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + try: + loop.run_until_complete( + self.acompletion( + model=silent_model, + messages=cast(List[AllMessageValues], messages), + **silent_kwargs, + ) + ) + finally: + loop.close() + asyncio.set_event_loop(None) except Exception as e: verbose_router_logger.error( f"Silent experiment failed for model {silent_model}: {str(e)}" @@ -1705,7 +1717,9 @@ class Router: and isinstance(fallback_item, ModelResponseStream) and hasattr(fallback_item, "usage") ): - self._combine_fallback_usage(fallback_item, complete_response_object_usage) + self._combine_fallback_usage( + fallback_item, complete_response_object_usage + ) yield fallback_item else: # If fallback returns a non-streaming response, yield None @@ -1825,13 +1839,11 @@ class Router: router_self._update_kwargs_before_fallbacks( model=model_group, kwargs=initial_kwargs ) - fallback_response = ( - router_self.function_with_fallbacks( - **initial_kwargs, - fallbacks=fallbacks, - context_window_fallbacks=context_window_fallbacks, - content_policy_fallbacks=content_policy_fallbacks, - ) + fallback_response = router_self.function_with_fallbacks( + **initial_kwargs, + fallbacks=fallbacks, + context_window_fallbacks=context_window_fallbacks, + content_policy_fallbacks=content_policy_fallbacks, ) if hasattr(fallback_response, "__iter__"): @@ -1841,7 +1853,9 @@ class Router: and isinstance(fallback_item, ModelResponseStream) and hasattr(fallback_item, "usage") ): - router_self._combine_fallback_usage(fallback_item, complete_response_object_usage) + router_self._combine_fallback_usage( + fallback_item, complete_response_object_usage + ) yield fallback_item else: yield None @@ -1891,6 +1905,11 @@ class Router: ) silent_kwargs = self._get_silent_experiment_kwargs(**kwargs) + if ( + "metadata" in silent_kwargs + and "model_group" in silent_kwargs["metadata"] + ): + silent_kwargs["metadata"]["model_group"] = silent_model # Trigger the silent request await self.acompletion( @@ -2753,10 +2772,9 @@ class Router: litellm_model = data.get("model", None) # litellm_agent/ prefix only strips the model name, no prompt_id needed - is_litellm_agent_model = ( - isinstance(litellm_model, str) - and litellm_model.startswith("litellm_agent/") - ) + is_litellm_agent_model = isinstance( + litellm_model, str + ) and litellm_model.startswith("litellm_agent/") prompt_id = kwargs.get("prompt_id") or prompt_management_deployment[ "litellm_params" @@ -6538,7 +6556,7 @@ class Router: tiers = complexity_router_config.get("tiers", {}) # Use MEDIUM tier as fallback default default_model = tiers.get("MEDIUM") or tiers.get("SIMPLE") - + if default_model is None: raise ValueError( "complexity_router_default_model is required for complexity-router deployments, " @@ -6771,7 +6789,9 @@ class Router: ######################################################### # Check if this is a complexity-router deployment ######################################################### - if self._is_complexity_router_deployment(litellm_params=deployment.litellm_params): + if self._is_complexity_router_deployment( + litellm_params=deployment.litellm_params + ): self.init_complexity_router_deployment(deployment=deployment) return deployment @@ -6863,9 +6883,7 @@ class Router: # zero-cost models, causing budget checks to block free models. _model_id = deployment.model_info.id if _model_id is not None: - _model_info_dict: dict = deployment.model_info.model_dump( - exclude_none=True - ) + _model_info_dict: dict = deployment.model_info.model_dump(exclude_none=True) for field in CustomPricingLiteLLMParams.model_fields.keys(): field_value = deployment.litellm_params.get(field) if field_value is not None: @@ -7157,7 +7175,10 @@ class Router: @overload def get_router_model_info( - self, deployment: Union[dict, "Deployment"], received_model_name: str, id: None = None + self, + deployment: Union[dict, "Deployment"], + received_model_name: str, + id: None = None, ) -> ModelMapInfo: pass @@ -7197,7 +7218,9 @@ class Router: ## GET BASE MODEL base_model = (deployment.get("model_info") or {}).get("base_model", None) if base_model is None: - base_model = (deployment.get("litellm_params") or {}).get("base_model", None) + base_model = (deployment.get("litellm_params") or {}).get( + "base_model", None + ) model = base_model @@ -7232,12 +7255,12 @@ class Router: if potential_models is not None: for potential_model in potential_models: try: - if (potential_model.get("model_info") or {}).get( - "id" - ) == (deployment.get("model_info") or {}).get("id"): - model = (potential_model.get("litellm_params") or {}).get( - "model" - ) + if (potential_model.get("model_info") or {}).get("id") == ( + deployment.get("model_info") or {} + ).get("id"): + model = ( + potential_model.get("litellm_params") or {} + ).get("model") break except Exception: pass @@ -8160,7 +8183,9 @@ class Router: - team_id: Optional[str] - the team id, to resolve team-specific models """ # Check if this is the no-args hot path (cacheable) - _use_cache = model_name is None and model_access_group is None and team_id is None + _use_cache = ( + model_name is None and model_access_group is None and team_id is None + ) # Return cached result for the no-args hot path if _use_cache and self._access_groups_cache is not None: diff --git a/tests/test_litellm/test_router_silent_experiment.py b/tests/test_litellm/test_router_silent_experiment.py index a23ea80f7ce..5d5620600f5 100644 --- a/tests/test_litellm/test_router_silent_experiment.py +++ b/tests/test_litellm/test_router_silent_experiment.py @@ -173,12 +173,15 @@ def test_router_silent_experiment_completion(): router = Router(model_list=model_list) - # Mock litellm.completion + # Mock litellm.completion for primary call and litellm.acompletion for shadow call mock_response = litellm.ModelResponse(choices=[{"message": {"content": "hello"}}]) mock_completion = MagicMock(return_value=mock_response) + mock_acompletion = AsyncMock(return_value=mock_response) # Patch at the litellm module level - with patch.object(litellm, "completion", mock_completion): + with patch.object(litellm, "completion", mock_completion), patch.object( + litellm, "acompletion", mock_acompletion + ): response = router.completion( model="primary-model", messages=[{"role": "user", "content": "hi"}], @@ -191,24 +194,124 @@ def test_router_silent_experiment_completion(): time.sleep(0.5) - # Should have 2 calls - assert mock_completion.call_count == 2 + # Should have 1 call to completion (primary) and 1 call to acompletion (shadow) + assert mock_completion.call_count == 1 + assert mock_acompletion.call_count == 1 - call_args_list = mock_completion.call_args_list + primary_args, primary_kwargs = mock_completion.call_args + silent_args, silent_kwargs = mock_acompletion.call_args # Verify no silent_model in any call - for call in call_args_list: - args, kwargs = call - assert "silent_model" not in kwargs + assert "silent_model" not in primary_kwargs + assert "silent_model" not in silent_kwargs + + assert silent_kwargs.get("metadata", {}).get("is_silent_experiment") is True + assert silent_kwargs["model"] == "openai/gpt-4" + assert silent_kwargs.get("stream") is False + + assert ( + primary_kwargs.get("metadata", {}).get("is_silent_experiment") is not True + ) + assert primary_kwargs["model"] == "openai/gpt-3.5-turbo" + + +def test_silent_experiment_forces_stream_false(): + """Verify that _get_silent_experiment_kwargs() sets stream=False even if stream=True in kwargs.""" + model_list = [ + { + "model_name": "gpt-3.5-turbo", + "litellm_params": {"model": "gpt-3.5-turbo", "api_key": "fake-key"}, + }, + ] + router = Router(model_list=model_list) + kwargs = {"stream": True, "metadata": {}} + result = router._get_silent_experiment_kwargs(**kwargs) + assert result["stream"] is False + + +def test_silent_experiment_sets_zero_retries(): + """Verify num_retries=0 in silent kwargs.""" + model_list = [ + { + "model_name": "gpt-3.5-turbo", + "litellm_params": {"model": "gpt-3.5-turbo", "api_key": "fake-key"}, + }, + ] + router = Router(model_list=model_list) + kwargs = {"num_retries": 3, "metadata": {}} + result = router._get_silent_experiment_kwargs(**kwargs) + assert result["num_retries"] == 0 + + +@pytest.mark.asyncio +async def test_silent_experiment_streaming_primary_triggers_shadow(): + """ + Mock a streaming primary request and verify the silent model call is made with stream=False + and both primary + silent calls complete. + """ + model_list = [ + { + "model_name": "primary-model", + "litellm_params": { + "model": "openai/gpt-3.5-turbo", + "api_key": "fake-key", + "silent_model": "silent-model", + }, + }, + { + "model_name": "silent-model", + "litellm_params": { + "model": "openai/gpt-4", + "api_key": "fake-key", + }, + }, + ] + + router = Router(model_list=model_list) + + mock_response = litellm.ModelResponse(choices=[{"message": {"content": "hello"}}]) + mock_acompletion = AsyncMock(return_value=mock_response) + + with patch.object(litellm, "acompletion", mock_acompletion): + await router.acompletion( + model="primary-model", + messages=[{"role": "user", "content": "hi"}], + stream=True, + ) + + await asyncio.sleep(0.1) + + assert mock_acompletion.call_count == 2 + calls = mock_acompletion.call_args_list - # Find the silent call silent_call = next( ( c - for c in call_args_list + for c in calls if c[1].get("metadata", {}).get("is_silent_experiment") is True ), None, ) assert silent_call is not None - assert silent_call[1]["model"] == "openai/gpt-4" + assert silent_call[1].get("stream") is False + + +def test_silent_experiment_sync_uses_async_path(): + """Verify the sync _silent_experiment_completion calls acompletion (not completion).""" + model_list = [ + { + "model_name": "gpt-3.5-turbo", + "litellm_params": {"model": "gpt-3.5-turbo", "api_key": "fake-key"}, + }, + ] + router = Router(model_list=model_list) + messages = [{"role": "user", "content": "hi"}] + + with patch.object( + router, "acompletion", new_callable=AsyncMock, return_value=None + ) as mock_acompletion: + router._silent_experiment_completion( + silent_model="gpt-3.5-turbo", + messages=messages, + ) + mock_acompletion.assert_called_once()