mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
test_vertex_ai_bge_with_endpoint_id_pattern
This commit is contained in:
parent
8ea8c674e1
commit
075a80b747
1 changed files with 82 additions and 0 deletions
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue