From 4852641ad5690f5e7f34effa144d79837c2e271e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=86=AF=E5=9F=BA=E9=AD=81?= <1412414664@qq.com> Date: Mon, 8 Jun 2026 13:14:08 +0800 Subject: [PATCH] fix(ui): preserve configured model in connection test --- .../health_endpoints/_health_endpoints.py | 2 + .../health_endpoints/test_health_endpoints.py | 71 +++++++++++++++++-- 2 files changed, 69 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index c109f374993..bf7e3eba048 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -1932,6 +1932,8 @@ async def test_model_connection( # Merge: config params (from proxy config) as base, request params override # This allows users to override specific params while using config for credentials litellm_params = {**config_litellm_params, **request_litellm_params} + if config_litellm_params.get("model") is not None: + litellm_params["model"] = config_litellm_params["model"] ## Auth check await ModelManagementAuthChecks.can_user_make_model_call( diff --git a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py index 80a4804956c..59893972045 100644 --- a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py @@ -450,22 +450,85 @@ async def test_test_model_connection_loads_config_from_router(): model_params = ahealth_check_call_args.kwargs.get("model_params", {}) # Verify that config params were loaded and merged - # Note: request params override config params, so model from request is used assert model_params.get("api_key") == "resolved-api-key-from-env" assert ( model_params.get("api_base") == "https://resolved-endpoint.openai.azure.com/" ) assert model_params.get("api_version") == "2024-10-21" - assert ( - model_params.get("model") == "gpt-4o" - ) # Request param overrides config param + assert model_params.get("model") == "azure/gpt-4o" # Verify result assert result["status"] == "success" assert "result" in result +@pytest.mark.asyncio +async def test_test_model_connection_preserves_configured_provider_model(): + from litellm.types.router import Deployment, LiteLLM_Params + + mock_router = MagicMock() + mock_router.get_deployment.return_value = Deployment( + model_name="claude-haiku-4.5", + litellm_params=LiteLLM_Params( + model="bedrock/converse/eu.anthropic.claude-haiku-4-5-20251001-v1:0", + aws_region_name="eu-central-1", + drop_params=True, + additional_drop_params=["extra_headers"], + ), + model_info={"id": "bedrock-deployment-id", "mode": "chat"}, + ) + + mock_ahealth_check = AsyncMock(return_value={"status": "healthy"}) + mock_run_with_timeout = AsyncMock(return_value={"status": "healthy"}) + + def mock_update_params(model_info, litellm_params): + params = litellm_params.copy() + params["messages"] = [{"role": "user", "content": "test"}] + return params + + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.llm_router", mock_router), + patch("litellm.proxy.proxy_server.premium_user", False), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", + AsyncMock(), + ), + patch( + "litellm.proxy.health_endpoints._health_endpoints.litellm.ahealth_check", + mock_ahealth_check, + ), + patch( + "litellm.proxy.health_endpoints._health_endpoints.run_with_timeout", + mock_run_with_timeout, + ), + patch( + "litellm.proxy.health_endpoints._health_endpoints._update_litellm_params_for_health_check", + mock_update_params, + ), + patch( + "litellm.proxy.health_endpoints._health_endpoints._reject_os_environ_references", + lambda params: None, + ), + ): + result = await health_test_model_connection( + request=MagicMock(), + mode="chat", + litellm_params={"model": "claude-haiku-4.5"}, + model_info={"id": "bedrock-deployment-id", "mode": "chat"}, + user_api_key_dict=MagicMock(), + ) + + assert result["status"] == "success" + model_params = mock_ahealth_check.call_args.kwargs["model_params"] + assert ( + model_params["model"] + == "bedrock/converse/eu.anthropic.claude-haiku-4-5-20251001-v1:0" + ) + assert model_params["aws_region_name"] == "eu-central-1" + + @pytest.mark.asyncio async def test_test_model_connection_uses_model_info_id_to_disambiguate_duplicate_model_names(): """