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