mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
test(predibase): add targeted branch coverage for migration
Add focused Predibase tests to cover remaining transformation and handler branches highlighted by Codecov, including env URL fallback and async/sync delegation paths. Made-with: Cursor
This commit is contained in:
parent
b3fec9ebcc
commit
1bcdd957f0
1 changed files with 190 additions and 0 deletions
|
|
@ -225,6 +225,114 @@ def test_predibase_transform_response_missing_generated_text():
|
|||
)
|
||||
|
||||
|
||||
def test_predibase_transform_response_non_dict_payload():
|
||||
config = PredibaseConfig()
|
||||
raw_response = Mock()
|
||||
raw_response.text = "[]"
|
||||
raw_response.status_code = 200
|
||||
raw_response.headers = {}
|
||||
raw_response.json.return_value = []
|
||||
|
||||
with pytest.raises(PredibaseError, match="'completion_response' is not a dictionary"):
|
||||
config.transform_response(
|
||||
model="predibase-model",
|
||||
raw_response=raw_response,
|
||||
model_response=_build_model_response(),
|
||||
logging_obj=Mock(),
|
||||
request_data={},
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding=Mock(),
|
||||
api_key="test-key",
|
||||
)
|
||||
|
||||
|
||||
def test_predibase_transform_response_best_of_with_empty_generated_text(monkeypatch):
|
||||
config = PredibaseConfig()
|
||||
logging_obj = Mock()
|
||||
encoding = Mock()
|
||||
encoding.encode.return_value = [1]
|
||||
monkeypatch.setattr("litellm.token_counter", lambda messages: 1)
|
||||
|
||||
raw_response = httpx.Response(
|
||||
status_code=200,
|
||||
json={
|
||||
"generated_text": "primary-output",
|
||||
"details": {
|
||||
"finish_reason": "stop",
|
||||
"tokens": [],
|
||||
"best_of_sequences": [
|
||||
{
|
||||
"generated_text": "",
|
||||
"finish_reason": "length",
|
||||
"tokens": [],
|
||||
}
|
||||
],
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
result = config.transform_response(
|
||||
model="predibase-model",
|
||||
raw_response=raw_response,
|
||||
model_response=_build_model_response(),
|
||||
logging_obj=logging_obj,
|
||||
request_data={"inputs": "hello", "parameters": {}},
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
optional_params={"best_of": 2},
|
||||
litellm_params={},
|
||||
encoding=encoding,
|
||||
api_key="test-key",
|
||||
)
|
||||
|
||||
assert len(result.choices) == 2
|
||||
assert result.choices[1].message.content is None
|
||||
|
||||
|
||||
def test_predibase_transform_response_empty_output_sets_completion_tokens_zero(monkeypatch):
|
||||
config = PredibaseConfig()
|
||||
logging_obj = Mock()
|
||||
encoding = Mock()
|
||||
monkeypatch.setattr("litellm.token_counter", lambda messages: 3)
|
||||
|
||||
raw_response = httpx.Response(
|
||||
status_code=200,
|
||||
json={"generated_text": "", "details": {"tokens": [], "finish_reason": "stop"}},
|
||||
)
|
||||
|
||||
result = config.transform_response(
|
||||
model="predibase-model",
|
||||
raw_response=raw_response,
|
||||
model_response=_build_model_response(),
|
||||
logging_obj=logging_obj,
|
||||
request_data={"inputs": "hello", "parameters": {}},
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding=encoding,
|
||||
api_key="test-key",
|
||||
)
|
||||
|
||||
assert result.usage.prompt_tokens == 3
|
||||
assert result.usage.completion_tokens == 0
|
||||
|
||||
|
||||
def test_predibase_get_complete_url_uses_env_base_url(monkeypatch):
|
||||
config = PredibaseConfig()
|
||||
monkeypatch.setenv("PREDIBASE_API_BASE", "https://env.predibase.com")
|
||||
|
||||
url = config.get_complete_url(
|
||||
api_base=None,
|
||||
api_key="test-key",
|
||||
model="predibase-model",
|
||||
optional_params={},
|
||||
litellm_params={"predibase_tenant_id": "tenant-123"},
|
||||
)
|
||||
|
||||
assert url.startswith("https://env.predibase.com/tenant-123/")
|
||||
|
||||
|
||||
def test_predibase_transform_response_usage_fallbacks(monkeypatch):
|
||||
config = PredibaseConfig()
|
||||
logging_obj = Mock()
|
||||
|
|
@ -256,6 +364,43 @@ def test_predibase_transform_response_usage_fallbacks(monkeypatch):
|
|||
assert result.usage.completion_tokens == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_predibase_async_completion_uses_default_config_when_none(monkeypatch):
|
||||
handler = PredibaseChatCompletion()
|
||||
mock_response = httpx.Response(status_code=200, json={"generated_text": "ok"})
|
||||
|
||||
async_handler = Mock()
|
||||
async_handler.post = AsyncMock(return_value=mock_response)
|
||||
monkeypatch.setattr(
|
||||
"litellm.llms.predibase.chat.handler.get_async_httpx_client",
|
||||
lambda **kwargs: async_handler,
|
||||
)
|
||||
|
||||
default_config = Mock()
|
||||
default_config.transform_response.return_value = _build_model_response()
|
||||
monkeypatch.setattr("litellm.PredibaseConfig", lambda: default_config)
|
||||
|
||||
result = await handler.async_completion(
|
||||
model="predibase-model",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
api_base="https://serving.example.com/x/generate",
|
||||
model_response=_build_model_response(),
|
||||
print_verbose=Mock(),
|
||||
encoding=Mock(),
|
||||
api_key="test-key",
|
||||
logging_obj=Mock(),
|
||||
stream=False,
|
||||
data={"inputs": "hello", "parameters": {}},
|
||||
optional_params={},
|
||||
timeout=10,
|
||||
litellm_params={},
|
||||
headers={"Authorization": "Bearer test"},
|
||||
)
|
||||
|
||||
assert result is default_config.transform_response.return_value
|
||||
default_config.transform_response.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_predibase_async_completion_uses_passed_config(monkeypatch):
|
||||
handler = PredibaseChatCompletion()
|
||||
|
|
@ -293,6 +438,51 @@ async def test_predibase_async_completion_uses_passed_config(monkeypatch):
|
|||
passed_config.transform_response.assert_called_once()
|
||||
|
||||
|
||||
def test_predibase_completion_sync_returns_transform_response(monkeypatch):
|
||||
handler = PredibaseChatCompletion()
|
||||
expected = _build_model_response()
|
||||
|
||||
def fake_validate_environment(self, **kwargs):
|
||||
return {"Authorization": "Bearer test"}
|
||||
|
||||
def fake_get_complete_url(self, **kwargs):
|
||||
return "https://serving.example.com/tenant/deployments/v2/llms/model/generate"
|
||||
|
||||
def fake_transform_request(self, **kwargs):
|
||||
return {"inputs": "hello", "parameters": {}}
|
||||
|
||||
def fake_transform_response(self, **kwargs):
|
||||
return expected
|
||||
|
||||
monkeypatch.setattr(PredibaseConfig, "validate_environment", fake_validate_environment)
|
||||
monkeypatch.setattr(PredibaseConfig, "get_complete_url", fake_get_complete_url)
|
||||
monkeypatch.setattr(PredibaseConfig, "transform_request", fake_transform_request)
|
||||
monkeypatch.setattr(PredibaseConfig, "transform_response", fake_transform_response)
|
||||
monkeypatch.setattr(
|
||||
"litellm.module_level_client.post",
|
||||
lambda *args, **kwargs: httpx.Response(status_code=200, json={"generated_text": "ok"}),
|
||||
)
|
||||
|
||||
result = handler.completion(
|
||||
model="predibase-model",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
api_base="https://serving.example.com",
|
||||
custom_prompt_dict={},
|
||||
model_response=_build_model_response(),
|
||||
print_verbose=Mock(),
|
||||
encoding=Mock(),
|
||||
api_key="test-key",
|
||||
logging_obj=Mock(),
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
tenant_id="tenant-123",
|
||||
timeout=10,
|
||||
acompletion=False,
|
||||
)
|
||||
|
||||
assert result is expected
|
||||
|
||||
|
||||
def test_predibase_completion_passes_existing_config_to_async_completion(monkeypatch):
|
||||
handler = PredibaseChatCompletion()
|
||||
captured = {}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue