diff --git a/tests/llm_translation/test_databricks.py b/tests/llm_translation/test_databricks.py index 77d803c061f..e9e6e67ec1f 100644 --- a/tests/llm_translation/test_databricks.py +++ b/tests/llm_translation/test_databricks.py @@ -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):