diff --git a/tests/pass_through_unit_tests/test_pass_through_unit_tests.py b/tests/pass_through_unit_tests/test_pass_through_unit_tests.py index 0e681cb1e02..b56a598aaee 100644 --- a/tests/pass_through_unit_tests/test_pass_through_unit_tests.py +++ b/tests/pass_through_unit_tests/test_pass_through_unit_tests.py @@ -420,6 +420,28 @@ def test_is_bedrock_agent_runtime_route(): assert _is_bedrock_agent_runtime_route("/some/random/endpoint") is False +def test_is_bedrock_kb_retrieve_route(): + """ + Test that _is_bedrock_kb_retrieve_route correctly identifies KB retrieve endpoints + """ + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + _is_bedrock_kb_retrieve_route, + ) + + # Test KB retrieve endpoints (should return True) + assert _is_bedrock_kb_retrieve_route("/knowledgebases/kb-123/retrieve") is True + assert _is_bedrock_kb_retrieve_route("knowledgebases/kb-123/retrieve") is True + assert _is_bedrock_kb_retrieve_route("/knowledgebases/kb-123/retrieve/") is True + + # Test non-retrieve KB endpoints (should return False) + assert _is_bedrock_kb_retrieve_route("/knowledgebases/kb-123/query") is False + assert _is_bedrock_kb_retrieve_route("/knowledgebases/kb-123") is False + + # Test non-KB endpoints (should return False) + assert _is_bedrock_kb_retrieve_route("/model/test/converse") is False + assert _is_bedrock_kb_retrieve_route("/some/random/endpoint") is False + + def test_init_kwargs_filters_pricing_params(mock_request, mock_user_api_key_dict): """ Test that pricing parameters are properly filtered out from the request body diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 368e2103590..23a55c0ca5f 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -1334,6 +1334,69 @@ class TestBedrockLLMProxyRoute: # and they're available in the router's deployment assert mock_process.called + @pytest.mark.asyncio + async def test_bedrock_kb_retrieve_passthrough_uses_correct_url(self): + """ + Test that _bedrock_kb_retrieve_passthrough constructs the correct URL + using bedrock-agent-runtime host and signs the request properly. + """ + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + _bedrock_kb_retrieve_passthrough, + ) + + mock_request = MagicMock(spec=Request) + mock_request.method = "POST" + + mock_credentials = MagicMock() + mock_credentials.access_key = "test-access-key" + mock_credentials.secret_key = "test-secret-key" + mock_credentials.token = None + + mock_bedrock_llm = MagicMock() + mock_bedrock_llm.get_aws_region_name_for_non_llm_api_calls.return_value = ( + "us-east-1" + ) + + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.content = b'{"retrievalResults": []}' + mock_response.headers = {"content-type": "application/json"} + mock_response.raise_for_status = MagicMock() + + captured_url = None + + async def capture_request(**kwargs): + nonlocal captured_url + captured_url = kwargs.get("url") + return mock_response + + mock_client = MagicMock() + mock_client.request = AsyncMock(side_effect=capture_request) + + with patch( + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client" + ) as mock_get_client, patch( + "botocore.auth.SigV4Auth" + ): + mock_get_client.return_value.client = mock_client + + response = await _bedrock_kb_retrieve_passthrough( + request=mock_request, + aws_region_name="us-east-1", + encoded_endpoint="/knowledgebases/KB123/retrieve", + data={"retrievalQuery": {"text": "test query"}}, + credentials=mock_credentials, + bedrock_llm=mock_bedrock_llm, + ) + + # Verify the URL uses bedrock-agent-runtime host + assert captured_url is not None + assert "bedrock-agent-runtime.us-east-1.amazonaws.com" in captured_url + assert "/knowledgebases/KB123/retrieve" in captured_url + + # Verify response is returned + assert response.status_code == 200 + class TestLLMPassthroughFactoryProxyRoute: @pytest.mark.asyncio