Add tests to validate the retrieve path

This commit is contained in:
Marcos Griselli 2026-01-23 23:24:40 -08:00
parent 1920e3c9dc
commit f35f549748
2 changed files with 85 additions and 0 deletions

View file

@ -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

View file

@ -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