From db4183715a27c66e20a45cd1e6028fff4d595ada Mon Sep 17 00:00:00 2001 From: Tornike Gurgenidze Date: Fri, 23 May 2025 09:55:46 +0400 Subject: [PATCH] feat: add embeddings to CustomLLM (#10980) * feat: add embeddings to CustomLLM * feat: add aembedding to custom llm --- litellm/llms/custom_llm.py | 26 +++++++- litellm/main.py | 24 ++++++++ tests/local_testing/test_custom_llm.py | 82 +++++++++++++++++++++++++- 3 files changed, 130 insertions(+), 2 deletions(-) diff --git a/litellm/llms/custom_llm.py b/litellm/llms/custom_llm.py index a2d04b1838d..390258e4e82 100644 --- a/litellm/llms/custom_llm.py +++ b/litellm/llms/custom_llm.py @@ -14,7 +14,7 @@ import httpx from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.types.utils import GenericStreamingChunk -from litellm.utils import ImageResponse, ModelResponse +from litellm.utils import ImageResponse, ModelResponse, EmbeddingResponse from .base import BaseLLM @@ -152,6 +152,30 @@ class CustomLLM(BaseLLM): ) -> ImageResponse: raise CustomLLMError(status_code=500, message="Not implemented yet!") + def embedding( + self, + model: str, + input: list, + model_response: EmbeddingResponse, + print_verbose: Callable, + logging_obj: Any, + optional_params: dict, + litellm_params=None, + ) -> EmbeddingResponse: + raise CustomLLMError(status_code=500, message="Not implemented yet!") + + async def aembedding( + self, + model: str, + input: list, + model_response: EmbeddingResponse, + print_verbose: Callable, + logging_obj: Any, + optional_params: dict, + litellm_params=None, + ) -> EmbeddingResponse: + raise CustomLLMError(status_code=500, message="Not implemented yet!") + def custom_chat_llm_router( async_fn: bool, stream: Optional[bool], custom_llm: CustomLLM diff --git a/litellm/main.py b/litellm/main.py index 1c1f4879cc8..44611e203f0 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -4027,6 +4027,30 @@ def embedding( # noqa: PLR0915 client=client, aembedding=aembedding, ) + elif ( + custom_llm_provider in litellm._custom_providers + ): + custom_handler: Optional[CustomLLM] = None + for item in litellm.custom_provider_map: + if item["provider"] == custom_llm_provider: + custom_handler = item["custom_handler"] + + if custom_handler is None: + raise LiteLLMUnknownProvider( + model=model, custom_llm_provider=custom_llm_provider + ) + + handler_fn = custom_handler.embedding if not aembedding else custom_handler.aembedding + + response = handler_fn( + model=model, + input=input, + logging_obj=logging, + optional_params=optional_params, + model_response=EmbeddingResponse(), + print_verbose=print_verbose, + litellm_params=litellm_params + ) else: raise LiteLLMUnknownProvider( model=model, custom_llm_provider=custom_llm_provider diff --git a/tests/local_testing/test_custom_llm.py b/tests/local_testing/test_custom_llm.py index beb1e3332dd..77f4544afa1 100644 --- a/tests/local_testing/test_custom_llm.py +++ b/tests/local_testing/test_custom_llm.py @@ -44,7 +44,7 @@ from litellm import ( image_generation, ) from litellm.utils import ModelResponseIterator -from litellm.types.utils import ImageResponse, ImageObject +from litellm.types.utils import ImageResponse, ImageObject, EmbeddingResponse from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler @@ -257,6 +257,53 @@ class MyCustomLLM(CustomLLM): response_ms=1000, ) + def embedding( + self, + model: str, + input: list, + model_response: EmbeddingResponse, + print_verbose: Callable, + logging_obj: Any, + optional_params: dict, + litellm_params=None, + aembedding=None, + ) -> EmbeddingResponse: + model_response.model = model + + model_response.data = [ + { + "object": "embedding", + "embedding": [0.1, 0.2, 0.3], + "index": i, + } + for i, _ in enumerate(input) + ] + + return model_response + + async def aembedding( + self, + model: str, + input: list, + model_response: EmbeddingResponse, + print_verbose: Callable, + logging_obj: Any, + optional_params: dict, + litellm_params=None, + ) -> EmbeddingResponse: + model_response.model = model + + model_response.data = [ + { + "object": "embedding", + "embedding": [0.1, 0.2, 0.3], + "index": i, + } + for i, _ in enumerate(input) + ] + + return model_response + def test_get_llm_provider(): """""" @@ -452,3 +499,36 @@ def test_get_supported_openai_params(): response = get_supported_openai_params(model="my-custom-llm/my-fake-model") assert response is not None + +def test_simple_embedding(): + my_custom_llm = MyCustomLLM() + litellm.custom_provider_map = [ + {"provider": "custom_llm", "custom_handler": my_custom_llm} + ] + resp = litellm.embedding( + model="custom_llm/my-fake-model", + input=["good morning from litellm", "good night from litellm"] + ) + + assert resp.data[1] == { + "object": "embedding", + "embedding": [0.1, 0.2, 0.3], + "index": 1, + } + +@pytest.mark.asyncio +async def test_simple_aembedding(): + my_custom_llm = MyCustomLLM() + litellm.custom_provider_map = [ + {"provider": "custom_llm", "custom_handler": my_custom_llm} + ] + resp = await litellm.aembedding( + model="custom_llm/my-fake-model", + input=["good morning from litellm", "good night from litellm"] + ) + + assert resp.data[1] == { + "object": "embedding", + "embedding": [0.1, 0.2, 0.3], + "index": 1, + } \ No newline at end of file