test_vertex_ai_bge_with_endpoint_id_pattern

This commit is contained in:
Ishaan Jaffer 2025-10-28 18:48:19 -07:00 • committed by Sameer Kankute
parent 8ea8c674e1
commit 075a80b747

View file

@ -99,3 +99,85 @@ def test_vertex_ai_bge_embedding_with_custom_api_base():
assert len(response.data) == 2
assert "embedding" in response.data[0]
def test_vertex_ai_bge_with_endpoint_id_pattern():
"""
Test BGE with vertex_ai/bge/endpoint_id pattern.
This test verifies that the pattern vertex_ai/bge/204379420394258432
correctly triggers BGE transformations and routes to the endpoint.
"""
client = HTTPHandler()
def mock_auth_token(*args, **kwargs):
return "fake-token", "fake-project"
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],
[0.6, 0.7, 0.8, 0.9, 1.0]
],
"deployedModelId": "204379420394258432",
"model": "projects/1060139831167/locations/europe-west4/models/baai_bge-base-en",
"modelDisplayName": "baai_bge-base-en",
"modelVersionId": "1"
}
mock_post.return_value = mock_response
response = litellm.embedding(
model="vertex_ai/bge/204379420394258432",
input=["Hello", "World"],
vertex_project="1060139831167",
vertex_location="europe-west4",
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"
# Vertex AI may use 'json' or 'data' parameter
if "json" in kwargs:
request_data = kwargs["json"]
elif "data" in kwargs:
request_data = json.loads(kwargs["data"])
else:
request_data = {}
print("\n" + "="*50)
print("BGE Endpoint Pattern Test:")
print("="*50)
print(f"Model: vertex_ai/bge/204379420394258432")
print(f"API URL: {api_url_called}")
print("Request Body:")
print(json.dumps(request_data, indent=2))
print("="*50 + "\n")
# Verify URL contains the endpoint ID and uses endpoints/ path
assert "204379420394258432" in api_url_called, f"Endpoint ID not in URL: {api_url_called}"
assert "endpoints" in api_url_called, f"Expected 'endpoints' in URL, got: {api_url_called}"
# Verify BGE-specific request format (uses "prompt" not "content")
assert "instances" in request_data
assert "prompt" in request_data["instances"][0]
assert request_data["instances"][0]["prompt"] == "Hello"
# Verify response
assert isinstance(response.data, list)
assert len(response.data) == 2