test_vertex_ai_bge_psc_endpoint_url_construction

This commit is contained in:
Ishaan Jaffer 2025-10-29 10:01:06 -07:00 • committed by Sameer Kankute
parent bfa7f12d4c
commit fe03833d3b

View file

@ -181,3 +181,71 @@ def test_vertex_ai_bge_with_endpoint_id_pattern():
assert len(response.data) == 2
def test_vertex_ai_bge_psc_endpoint_url_construction():
"""
Test that BGE models with PSC endpoints construct correct URL without bge/ prefix.
Verifies that vertex_ai/bge/378943383978115072 with api_base http://10.128.16.2
constructs URL: http://10.128.16.2/v1/projects/{project}/locations/{location}/endpoints/378943383978115072:predict
The bge/ prefix should be stripped from the endpoint URL.
"""
client = HTTPHandler()
def mock_auth_token(*args, **kwargs):
return "fake-token", "gen-lang-client-0682925754"
with patch.object(client, "post") as mock_post, patch(
"litellm.llms.vertex_ai.vertex_embeddings.embedding_handler.VertexEmbedding._ensure_access_token",
side_effect=mock_auth_token
):
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = {
"predictions": [
[0.1, 0.2, 0.3, 0.4, 0.5]
]
}
mock_post.return_value = mock_response
response = litellm.embedding(
model="vertex_ai/bge/378943383978115072",
input=["The food was delicious and the waiter.."],
api_base="http://10.128.16.2",
vertex_project="gen-lang-client-0682925754",
vertex_location="us-central1",
client=client
)
mock_post.assert_called_once()
call_args = mock_post.call_args
kwargs = call_args.kwargs if hasattr(call_args, 'kwargs') else call_args[1]
if "url" in kwargs:
api_url_called = kwargs["url"]
elif len(call_args[0]) > 0:
api_url_called = call_args[0][0]
else:
api_url_called = "Unknown"
print("\n" + "="*50)
print("PSC Endpoint URL Construction Test:")
print("="*50)
print(f"Model: vertex_ai/bge/378943383978115072")
print(f"API Base: http://10.128.16.2")
print(f"Constructed URL: {api_url_called}")
print("="*50 + "\n")
# Verify the URL is constructed correctly
expected_url = "http://10.128.16.2/v1/projects/gen-lang-client-0682925754/locations/us-central1/endpoints/378943383978115072:predict"
assert api_url_called == expected_url, f"Expected URL: {expected_url}, Got: {api_url_called}"
# Verify bge/ prefix is NOT in the URL
assert "bge/" not in api_url_called, f"URL should not contain 'bge/' prefix: {api_url_called}"
# Verify response works
assert isinstance(response.data, list)
assert len(response.data) == 1