diff --git a/litellm/utils.py b/litellm/utils.py index 065e7a172d7..0d4803ac881 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -36,7 +36,6 @@ import traceback from dataclasses import dataclass, field from functools import lru_cache, wraps from importlib import resources -from importlib.metadata import entry_points from inspect import iscoroutine from io import StringIO from os.path import abspath, dirname, join @@ -386,30 +385,10 @@ def print_verbose( ####### CLIENT ################### # make it easy to log if completion/embedding runs succeeded or failed + see what happened | Non-Blocking -def load_custom_provider_entrypoints(): - # Handle both Python 3.9 (returns dict) and Python 3.10+ (returns object with select method) - eps = entry_points() - if hasattr(eps, "select"): - # Python 3.10+ - found_entry_points = tuple(eps.select(group="litellm")) # type: ignore - else: - # Python 3.9 and earlier - entry_points() returns a dict - found_entry_points = eps.get("litellm", ()) # type: ignore - - for entry_point in found_entry_points: - # types are ignored because of circular dependency issues importing CustomLLM and CustomLLMItem - HandlerClass = entry_point.load() - handler = HandlerClass() - provider = {"provider": entry_point.name, "custom_handler": handler} - litellm.custom_provider_map.append(provider) # type: ignore - - def custom_llm_setup(): """ Add custom_llm provider to provider list """ - load_custom_provider_entrypoints() - for custom_llm in litellm.custom_provider_map: if custom_llm["provider"] not in litellm.provider_list: litellm.provider_list.append(custom_llm["provider"]) diff --git a/tests/local_testing/test_custom_llm.py b/tests/local_testing/test_custom_llm.py index 0f1af3b56e1..e61ede755e6 100644 --- a/tests/local_testing/test_custom_llm.py +++ b/tests/local_testing/test_custom_llm.py @@ -537,44 +537,3 @@ async def test_simple_aembedding(): "embedding": [0.1, 0.2, 0.3], "index": 1, } - - -def test_custom_llm_provider_entrypoint(): - # This test mocks the use of entry-points in pyproject.toml: - # [project.entry-point.litellm] - # custom_llm = :MyCustomLLM - # another-custom-llm = :AnotherCustomLLM - - from litellm.utils import custom_llm_setup - from importlib.metadata import EntryPoints, EntryPoint - - class AnotherCustomLLM(CustomLLM): - pass - - providers = { - "custom_llm": MyCustomLLM, - "another-custom-llm": AnotherCustomLLM - } - - def load(self): - return providers[self.name] - - entry_points = EntryPoints([ - EntryPoint(group="litellm", name="custom_llm", value="package.module:MyCustomLLM"), - EntryPoint(group="litellm", name="another-custom-llm", value="package.module:AnotherCustomLLM"), - ]) - - with patch("importlib.metadata.EntryPoint.load", load): - with patch("litellm.utils.entry_points", return_value=entry_points): - assert litellm.custom_provider_map == [] - assert litellm._custom_providers == [] - - custom_llm_setup() - - assert litellm._custom_providers == ['custom_llm', 'another-custom-llm'] - - assert litellm.custom_provider_map[0]["provider"] == "custom_llm" - assert isinstance(litellm.custom_provider_map[0]["custom_handler"], CustomLLM) - - assert litellm.custom_provider_map[1]["provider"] == "another-custom-llm" - assert isinstance(litellm.custom_provider_map[1]["custom_handler"], AnotherCustomLLM)