From a8a38778a3c6e257fc9fa20c1c94dd55b258e2d7 Mon Sep 17 00:00:00 2001 From: Xianzong Xie Date: Thu, 4 Dec 2025 17:47:30 -0800 Subject: [PATCH] fix: resolve provider from router for polling_via_cache - Fix bug where model names without slash (e.g., 'gpt-5') couldn't match providers in polling_via_cache list - Look up model in llm_router.model_name_to_deployment_indices - Check ALL deployments for matching provider (supports load balancing) - Check custom_llm_provider first, then extract from model string - Add comprehensive tests for provider resolution logic Committed-By-Agent: cursor --- .../proxy/response_api_endpoints/endpoints.py | 41 +++- .../test_response_polling_handler.py | 210 ++++++++++++++++++ 2 files changed, 246 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index d435f0a34cd..3956d081f4b 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -89,12 +89,43 @@ async def responses_api( # Enable for all models/providers should_use_polling = True elif isinstance(polling_via_cache_enabled, list): - # Check if provider is in the list (e.g., ["openai", "anthropic"]) + # Check if provider is in the list (e.g., ["openai", "bedrock"]) model = data.get("model", "") - # Extract provider from model (e.g., "openai/gpt-4" -> "openai") - provider = model.split("/")[0] if "/" in model else model - if provider in polling_via_cache_enabled: - should_use_polling = True + + # First, try to get provider from model string format "provider/model" + if "/" in model: + provider = model.split("/")[0] + if provider in polling_via_cache_enabled: + should_use_polling = True + # Otherwise, check ALL deployments for this model_name in router + elif llm_router is not None: + try: + # Get all deployment indices for this model name + indices = llm_router.model_name_to_deployment_indices.get(model, []) + for idx in indices: + deployment_dict = llm_router.model_list[idx] + litellm_params = deployment_dict.get("litellm_params", {}) + + # Check custom_llm_provider first + dep_provider = litellm_params.get("custom_llm_provider") + + # Then try to extract from model (e.g., "openai/gpt-5") + if not dep_provider: + dep_model = litellm_params.get("model", "") + if "/" in dep_model: + dep_provider = dep_model.split("/")[0] + + # If ANY deployment's provider matches, enable polling + if dep_provider and dep_provider in polling_via_cache_enabled: + should_use_polling = True + verbose_proxy_logger.debug( + f"Polling enabled for model={model}, provider={dep_provider}" + ) + break + except Exception as e: + verbose_proxy_logger.debug( + f"Could not resolve provider for model {model}: {e}" + ) # If all conditions are met, use polling mode if should_use_polling: diff --git a/tests/proxy_unit_tests/test_response_polling_handler.py b/tests/proxy_unit_tests/test_response_polling_handler.py index b47888dc4f7..545fc385a36 100644 --- a/tests/proxy_unit_tests/test_response_polling_handler.py +++ b/tests/proxy_unit_tests/test_response_polling_handler.py @@ -649,3 +649,213 @@ class TestBackgroundStreamingModule: assert asyncio.iscoroutinefunction(background_streaming_task) + +class TestProviderResolutionForPolling: + """ + Test cases for provider resolution logic used to determine + if polling_via_cache should be enabled for a given model. + + This tests the logic in endpoints.py that resolves model names + to their providers using the router's deployment configuration. + """ + + def test_provider_from_model_string_with_slash(self): + """Test extracting provider from 'provider/model' format""" + model = "openai/gpt-4o" + + # Direct extraction when model has slash + if "/" in model: + provider = model.split("/")[0] + else: + provider = None + + assert provider == "openai" + + def test_provider_from_model_string_without_slash(self): + """Test that model without slash doesn't extract provider directly""" + model = "gpt-5" + + # No slash means we can't extract provider directly + if "/" in model: + provider = model.split("/")[0] + else: + provider = None + + assert provider is None + + def test_provider_resolution_from_router_single_deployment(self): + """Test resolving provider from router with single deployment""" + # Simulate router's model_name_to_deployment_indices + model_name_to_deployment_indices = { + "gpt-5": [0], # Single deployment at index 0 + } + model_list = [ + { + "model_name": "gpt-5", + "litellm_params": { + "model": "openai/gpt-5", + "api_key": "sk-test", + } + } + ] + + model = "gpt-5" + polling_via_cache_enabled = ["openai"] + should_use_polling = False + + # Simulate the resolution logic + indices = model_name_to_deployment_indices.get(model, []) + for idx in indices: + deployment_dict = model_list[idx] + litellm_params = deployment_dict.get("litellm_params", {}) + + dep_provider = litellm_params.get("custom_llm_provider") + if not dep_provider: + dep_model = litellm_params.get("model", "") + if "/" in dep_model: + dep_provider = dep_model.split("/")[0] + + if dep_provider and dep_provider in polling_via_cache_enabled: + should_use_polling = True + break + + assert should_use_polling is True + + def test_provider_resolution_from_router_multiple_deployments_match(self): + """Test resolving provider when multiple deployments exist and one matches""" + model_name_to_deployment_indices = { + "gpt-4o": [0, 1], # Two deployments + } + model_list = [ + { + "model_name": "gpt-4o", + "litellm_params": { + "model": "openai/gpt-4o", + } + }, + { + "model_name": "gpt-4o", + "litellm_params": { + "model": "azure/gpt-4o-deployment", + } + } + ] + + model = "gpt-4o" + polling_via_cache_enabled = ["openai"] # Only openai in list + should_use_polling = False + + indices = model_name_to_deployment_indices.get(model, []) + for idx in indices: + deployment_dict = model_list[idx] + litellm_params = deployment_dict.get("litellm_params", {}) + + dep_provider = litellm_params.get("custom_llm_provider") + if not dep_provider: + dep_model = litellm_params.get("model", "") + if "/" in dep_model: + dep_provider = dep_model.split("/")[0] + + if dep_provider and dep_provider in polling_via_cache_enabled: + should_use_polling = True + break + + # Should be True because first deployment is openai + assert should_use_polling is True + + def test_provider_resolution_from_router_no_match(self): + """Test that polling is disabled when no deployment provider matches""" + model_name_to_deployment_indices = { + "claude-3": [0], + } + model_list = [ + { + "model_name": "claude-3", + "litellm_params": { + "model": "anthropic/claude-3-sonnet", + } + } + ] + + model = "claude-3" + polling_via_cache_enabled = ["openai", "bedrock"] # anthropic not in list + should_use_polling = False + + indices = model_name_to_deployment_indices.get(model, []) + for idx in indices: + deployment_dict = model_list[idx] + litellm_params = deployment_dict.get("litellm_params", {}) + + dep_provider = litellm_params.get("custom_llm_provider") + if not dep_provider: + dep_model = litellm_params.get("model", "") + if "/" in dep_model: + dep_provider = dep_model.split("/")[0] + + if dep_provider and dep_provider in polling_via_cache_enabled: + should_use_polling = True + break + + assert should_use_polling is False + + def test_provider_resolution_with_custom_llm_provider(self): + """Test that custom_llm_provider takes precedence over model string""" + model_name_to_deployment_indices = { + "my-model": [0], + } + model_list = [ + { + "model_name": "my-model", + "litellm_params": { + "model": "some-custom-model", + "custom_llm_provider": "openai", # Explicit provider + } + } + ] + + model = "my-model" + polling_via_cache_enabled = ["openai"] + should_use_polling = False + + indices = model_name_to_deployment_indices.get(model, []) + for idx in indices: + deployment_dict = model_list[idx] + litellm_params = deployment_dict.get("litellm_params", {}) + + # custom_llm_provider should be checked first + dep_provider = litellm_params.get("custom_llm_provider") + if not dep_provider: + dep_model = litellm_params.get("model", "") + if "/" in dep_model: + dep_provider = dep_model.split("/")[0] + + if dep_provider and dep_provider in polling_via_cache_enabled: + should_use_polling = True + break + + assert should_use_polling is True + + def test_provider_resolution_model_not_in_router(self): + """Test that unknown model doesn't enable polling""" + model_name_to_deployment_indices = { + "gpt-5": [0], + } + model_list = [ + { + "model_name": "gpt-5", + "litellm_params": {"model": "openai/gpt-5"} + } + ] + + model = "unknown-model" # Not in router + polling_via_cache_enabled = ["openai"] + should_use_polling = False + + indices = model_name_to_deployment_indices.get(model, []) # Empty list + for idx in indices: + # This loop won't execute + pass + + assert should_use_polling is False + assert len(indices) == 0 +