test_databricks_embeddings

This commit is contained in:
Ishaan Jaffer 2025-10-25 11:28:21 -07:00
parent a3febef431
commit c06098c351

View file

@ -802,35 +802,83 @@ class TestDatabricksCompletion(BaseLLMChatTest, BaseAnthropicChatTest):
@pytest.mark.parametrize("sync_mode", [True, False])
@pytest.mark.asyncio
async def test_databricks_embeddings(sync_mode):
async def test_databricks_embeddings(sync_mode, monkeypatch):
"""
Test Databricks embeddings with instruction parameter in both sync and async modes using mocked HTTP responses.
"""
import openai
try:
litellm.set_verbose = True
litellm.drop_params = True
base_url = "https://my.workspace.cloud.databricks.com/serving-endpoints"
api_key = "dapimykey"
monkeypatch.setenv("DATABRICKS_API_BASE", base_url)
monkeypatch.setenv("DATABRICKS_API_KEY", api_key)
if sync_mode:
mock_response = Mock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.json.return_value = mock_embedding_response()
inputs = ["good morning from litellm"]
instruction = "Represent this sentence for searching relevant passages:"
litellm.set_verbose = True
litellm.drop_params = True
if sync_mode:
sync_handler = HTTPHandler()
with patch.object(HTTPHandler, "post", return_value=mock_response) as mock_post:
response = litellm.embedding(
model="databricks/databricks-bge-large-en",
input=["good morning from litellm"],
instruction="Represent this sentence for searching relevant passages:",
input=inputs,
instruction=instruction,
client=sync_handler,
)
else:
openai.types.CreateEmbeddingResponse.model_validate(
response.model_dump(), strict=True
)
mock_post.assert_called_once_with(
f"{base_url}/embeddings",
headers={
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
},
data=json.dumps(
{
"model": "databricks-bge-large-en",
"input": inputs,
"instruction": instruction,
}
),
)
else:
async_handler = AsyncHTTPHandler()
with patch.object(AsyncHTTPHandler, "post", return_value=mock_response) as mock_post:
response = await litellm.aembedding(
model="databricks/databricks-bge-large-en",
input=["good morning from litellm"],
instruction="Represent this sentence for searching relevant passages:",
input=inputs,
instruction=instruction,
client=async_handler,
)
print(f"response: {response}")
openai.types.CreateEmbeddingResponse.model_validate(
response.model_dump(), strict=True
)
openai.types.CreateEmbeddingResponse.model_validate(
response.model_dump(), strict=True
)
# stubbed endpoint is setup to return this
# assert response.data[0]["embedding"] == [0.1, 0.2, 0.3]
except Exception as e:
pytest.fail(f"Error occurred: {e}")
mock_post.assert_called_once_with(
f"{base_url}/embeddings",
headers={
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
},
data=json.dumps(
{
"model": "databricks-bge-large-en",
"input": inputs,
"instruction": instruction,
}
),
)
def test_completion_with_prompt_caching_claude_model(monkeypatch):