mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge 481caa32de into 49affa7c01
This commit is contained in:
commit
75b685691b
2 changed files with 79 additions and 0 deletions
|
|
@ -7013,6 +7013,23 @@ def embedding(
|
|||
aembedding=aembedding,
|
||||
litellm_params={},
|
||||
)
|
||||
elif JSONProviderRegistry.exists(custom_llm_provider):
|
||||
if headers:
|
||||
optional_params["extra_headers"] = headers
|
||||
|
||||
response = openai_chat_completions.embedding(
|
||||
model=model,
|
||||
input=input,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
logging_obj=logging,
|
||||
timeout=timeout,
|
||||
model_response=EmbeddingResponse(),
|
||||
optional_params=optional_params,
|
||||
client=client,
|
||||
aembedding=aembedding,
|
||||
max_retries=max_retries,
|
||||
)
|
||||
elif custom_llm_provider in litellm._custom_providers:
|
||||
custom_handler: CustomLLM | None = None
|
||||
for item in litellm.custom_provider_map:
|
||||
|
|
|
|||
|
|
@ -245,6 +245,68 @@ class TestPinstripes:
|
|||
assert result["temperature"] == 0.7
|
||||
|
||||
|
||||
class TestJSONProviderEmbedding:
|
||||
"""Regression tests for https://github.com/BerriAI/litellm/issues/34503
|
||||
|
||||
JSON-configured providers are OpenAI-compatible, so embedding() must route them to the
|
||||
OpenAI embeddings handler instead of raising LiteLLMUnknownProvider.
|
||||
"""
|
||||
|
||||
@pytest.mark.respx()
|
||||
def test_scaleway_embedding_routed_to_openai_handler(self, respx_mock, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.setenv("SCW_SECRET_KEY", "fake-scaleway-key")
|
||||
|
||||
route = respx_mock.post("https://api.scaleway.ai/v1/embeddings").respond(
|
||||
json={
|
||||
"object": "list",
|
||||
"data": [
|
||||
{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}
|
||||
],
|
||||
"model": "BAAI/bge-multilingual-gemma2",
|
||||
"usage": {"prompt_tokens": 3, "total_tokens": 3},
|
||||
}
|
||||
)
|
||||
|
||||
response = litellm.embedding(
|
||||
model="scaleway/BAAI/bge-multilingual-gemma2",
|
||||
input=["hello world"],
|
||||
)
|
||||
|
||||
assert route.called
|
||||
request = route.calls[0].request
|
||||
assert request.headers["authorization"] == "Bearer fake-scaleway-key"
|
||||
assert json.loads(request.content)["model"] == "BAAI/bge-multilingual-gemma2"
|
||||
assert response.data[0]["embedding"] == [0.1, 0.2, 0.3]
|
||||
|
||||
@pytest.mark.respx()
|
||||
def test_json_provider_embedding_honors_custom_api_base_and_headers(self, respx_mock, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
|
||||
route = respx_mock.post("https://custom.publicai.local/v1/embeddings").respond(
|
||||
json={
|
||||
"object": "list",
|
||||
"data": [
|
||||
{"object": "embedding", "index": 0, "embedding": [0.4, 0.5]}
|
||||
],
|
||||
"model": "some-embedding-model",
|
||||
"usage": {"prompt_tokens": 2, "total_tokens": 2},
|
||||
}
|
||||
)
|
||||
|
||||
response = litellm.embedding(
|
||||
model="publicai/some-embedding-model",
|
||||
input=["hello world"],
|
||||
api_base="https://custom.publicai.local/v1",
|
||||
api_key="fake-publicai-key",
|
||||
extra_headers={"x-tenant-id": "tenant-123"},
|
||||
)
|
||||
|
||||
assert route.called
|
||||
assert route.calls[0].request.headers["x-tenant-id"] == "tenant-123"
|
||||
assert response.data[0]["embedding"] == [0.4, 0.5]
|
||||
|
||||
|
||||
class TestDarkbloom:
|
||||
def test_darkbloom_json_config_exists(self):
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue