diff --git a/litellm/router_strategy/auto_router.py b/litellm/router_strategy/auto_router.py deleted file mode 100644 index 6d3dae19398..00000000000 --- a/litellm/router_strategy/auto_router.py +++ /dev/null @@ -1,71 +0,0 @@ -""" -Auto-Routing Strategy that works with a Semantic Router Config -""" - -import json -import os -from typing import List, Literal, Optional, Union - -from semantic_router.schema import RouteChoice - -import litellm -from litellm.integrations.custom_logger import CustomLogger -from litellm.proxy.proxy_server import DualCache, UserAPIKeyAuth - - -class AutoRouter(CustomLogger): - DEFAULT_AUTO_SYNC_VALUE = "local" - def __init__( - self, - router_config_path: str, - default_model: str, - embedding_model: str, - embedding_model_api_key: str, - ): - from semantic_router import Route - from semantic_router.encoders import LiteLLMEncoder - from semantic_router.routers import SemanticRouter - self.router_config_path = router_config_path - self.auto_sync_value = self.DEFAULT_AUTO_SYNC_VALUE - loaded_router: SemanticRouter = SemanticRouter.from_json(self.router_config_path) - self.routelayer: SemanticRouter = SemanticRouter( - routes=loaded_router.routes, - encoder=LiteLLMEncoder( - name=embedding_model, - api_key=embedding_model_api_key, - ), - auto_sync=self.auto_sync_value, - ) - self.default_model = default_model - pass - - - async def async_pre_routing_hook( - self, - data: dict, call_type: Literal[ - "completion", - "text_completion", - "embeddings", - "image_generation", - "moderation", - "audio_transcription", - ]): - from semantic_router.schema import RouteChoice - - #self.routelayer.to_json("./config/router.json") - # If the call type is embeddings, do not modify the model. - if call_type in ["embeddings", "image_generation"]: - #print("Call type is 'embeddings', using default behavior.") - return data # Return without modifying the data - - msg = data['messages'][-1]['content'] - route_choice: Optional[Union[RouteChoice, List[RouteChoice]]] = self.routelayer(text=msg) - if isinstance(route_choice, RouteChoice): - data["model"] = route_choice.name - elif isinstance(route_choice, list): - data["model"] = route_choice[0].name - else: - data["model"] = self.default_model - - return data - diff --git a/litellm/router_strategy/auto_router/litellm_encoder.py b/litellm/router_strategy/auto_router/litellm_encoder.py new file mode 100644 index 00000000000..46bb071a11a --- /dev/null +++ b/litellm/router_strategy/auto_router/litellm_encoder.py @@ -0,0 +1,117 @@ +from typing import TYPE_CHECKING, Any, Union + +from semantic_router.encoders import DenseEncoder +from semantic_router.encoders.base import AsymmetricDenseMixin + +import litellm + +if TYPE_CHECKING: + from litellm.router import Router +else: + Router = Any + + +def litellm_to_list(embeds: litellm.EmbeddingResponse) -> list[list[float]]: + """Convert a LiteLLM embedding response to a list of embeddings. + + :param embeds: The LiteLLM embedding response. + :return: A list of embeddings. + """ + if ( + not embeds + or not isinstance(embeds, litellm.EmbeddingResponse) + or not embeds.data + ): + raise ValueError("No embeddings found in LiteLLM embedding response.") + return [x["embedding"] for x in embeds.data] + + +class LiteLLMRouterEncoder(DenseEncoder, AsymmetricDenseMixin): + """LiteLLM encoder class for generating embeddings using LiteLLM. + + The LiteLLMRouterEncoder class is a subclass of DenseEncoder and utilizes the LiteLLM Router SDK + to generate embeddings for given documents. It supports all encoders supported by LiteLLM + and supports customization of the score threshold for filtering or processing the embeddings. + """ + + type: str = "internal_litellm_router" + + def __init__( + self, + litellm_router_instance: "Router", + model_name: str, + score_threshold: Union[float, None] = None, + ): + """Initialize the LiteLLMEncoder. + + :param litellm_router_instance: The instance of the LiteLLM Router. + :type litellm_router_instance: Router + :param model_name: The name of the embedding model to use. Must use LiteLLM naming + convention (e.g. "openai/text-embedding-3-small" or "mistral/mistral-embed"). + :type model_name: str + :param score_threshold: The score threshold for the embeddings. + :type score_threshold: float + """ + self.litellm_router_instance: "Router" = litellm_router_instance + super().__init__( + name=model_name, + score_threshold=score_threshold if score_threshold is not None else 0.3, + ) + + def __call__(self, docs: list[Any], **kwargs) -> list[list[float]]: + """Encode a list of text documents into embeddings using LiteLLM. + + :param docs: List of text documents to encode. + :return: List of embeddings for each document.""" + return self.encode_queries(docs, **kwargs) + + async def acall(self, docs: list[Any], **kwargs) -> list[list[float]]: + """Encode a list of documents into embeddings using LiteLLM asynchronously. + + :param docs: List of documents to encode. + :return: List of embeddings for each document.""" + return await self.aencode_queries(docs, **kwargs) + + def encode_queries(self, docs: list[str], **kwargs) -> list[list[float]]: + try: + embeds = self.litellm_router_instance.embedding( + input=docs, model=f"{self.type}/{self.name}", **kwargs + ) + return litellm_to_list(embeds) + except Exception as e: + raise ValueError( + f"{self.type.capitalize()} API call failed. Error: {e}" + ) from e + + def encode_documents(self, docs: list[str], **kwargs) -> list[list[float]]: + try: + embeds = self.litellm_router_instance.embedding( + input=docs, model=f"{self.type}/{self.name}", **kwargs + ) + return litellm_to_list(embeds) + except Exception as e: + raise ValueError( + f"{self.type.capitalize()} API call failed. Error: {e}" + ) from e + + async def aencode_queries(self, docs: list[str], **kwargs) -> list[list[float]]: + try: + embeds = await self.litellm_router_instance.aembedding( + input=docs, model=f"{self.type}/{self.name}", **kwargs + ) + return litellm_to_list(embeds) + except Exception as e: + raise ValueError( + f"{self.type.capitalize()} API call failed. Error: {e}" + ) from e + + async def aencode_documents(self, docs: list[str], **kwargs) -> list[list[float]]: + try: + embeds = await self.litellm_router_instance.aembedding( + input=docs, model=f"{self.type}/{self.name}", **kwargs + ) + return litellm_to_list(embeds) + except Exception as e: + raise ValueError( + f"{self.type.capitalize()} API call failed. Error: {e}" + ) from e