diff --git a/litellm/router.py b/litellm/router.py index 5265c9848d5..fcdca40ad6b 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -401,6 +401,7 @@ class Router: self.default_max_parallel_requests = default_max_parallel_requests self.provider_default_deployment_ids: List[str] = [] self.pattern_router = PatternMatchRouter() + self.auto_routers: Dict[str, "AutoRouter"] = {} if model_list is not None: model_list = copy.deepcopy(model_list) @@ -520,7 +521,6 @@ class Router: routing_strategy_args=routing_strategy_args, ) self.access_groups = None - self.auto_routers: Dict[str, "AutoRouter"] = {} ## USAGE TRACKING ## if isinstance(litellm._async_success_callback, list): litellm.logging_callback_manager.add_litellm_async_success_callback( diff --git a/litellm/router_strategy/auto_router/auto_router.py b/litellm/router_strategy/auto_router/auto_router.py index 53bfa1f437d..d25d70fb0c3 100644 --- a/litellm/router_strategy/auto_router/auto_router.py +++ b/litellm/router_strategy/auto_router/auto_router.py @@ -41,17 +41,11 @@ class AutoRouter(CustomLogger): ) 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=LiteLLMRouterEncoder( - litellm_router_instance=litellm_router_instance, - model_name=embedding_model, - ), - auto_sync=self.auto_sync_value, - ) + self.loaded_router: SemanticRouter = SemanticRouter.from_json(self.router_config_path) + self.routelayer: Optional[SemanticRouter] = None self.default_model = default_model - pass + self.embedding_model: str = embedding_model + self.litellm_router_instance: "Router" = litellm_router_instance async def async_pre_routing_hook( @@ -67,12 +61,30 @@ class AutoRouter(CustomLogger): Used for the litellm auto-router to modify the request before the routing decision is made. """ + from semantic_router.routers import SemanticRouter from semantic_router.schema import RouteChoice + from litellm.router_strategy.auto_router.litellm_encoder import ( + LiteLLMRouterEncoder, + ) from litellm.types.router import PreRoutingHookResponse if messages is None: # do nothing, return same inputs return None + + if self.routelayer is None: + ####################### + # Create the route layer + ####################### + self.routelayer = SemanticRouter( + routes=self.loaded_router.routes, + encoder=LiteLLMRouterEncoder( + litellm_router_instance=self.litellm_router_instance, + model_name=self.embedding_model, + ), + auto_sync=self.auto_sync_value, + ) + user_message: Dict[str, str] = messages[-1] message_content: str = user_message.get("content", "") route_choice: Optional[Union[RouteChoice, List[RouteChoice]]] = self.routelayer(text=message_content) diff --git a/litellm/router_strategy/auto_router/litellm_encoder.py b/litellm/router_strategy/auto_router/litellm_encoder.py index 46bb071a11a..57093514a43 100644 --- a/litellm/router_strategy/auto_router/litellm_encoder.py +++ b/litellm/router_strategy/auto_router/litellm_encoder.py @@ -1,5 +1,6 @@ -from typing import TYPE_CHECKING, Any, Union +from typing import TYPE_CHECKING, Any, Optional, Union +from pydantic import ConfigDict from semantic_router.encoders import DenseEncoder from semantic_router.encoders.base import AsymmetricDenseMixin @@ -26,7 +27,19 @@ def litellm_to_list(embeds: litellm.EmbeddingResponse) -> list[list[float]]: return [x["embedding"] for x in embeds.data] -class LiteLLMRouterEncoder(DenseEncoder, AsymmetricDenseMixin): +class CustomDenseEncoder(DenseEncoder): + model_config = ConfigDict(extra='allow') + + def __init__(self, litellm_router_instance: Optional["Router"] = None, **kwargs): + # Extract litellm_router_instance from kwargs if passed there + if 'litellm_router_instance' in kwargs: + litellm_router_instance = kwargs.pop('litellm_router_instance') + + super().__init__(**kwargs) + self.litellm_router_instance = litellm_router_instance + + +class LiteLLMRouterEncoder(CustomDenseEncoder, AsymmetricDenseMixin): """LiteLLM encoder class for generating embeddings using LiteLLM. The LiteLLMRouterEncoder class is a subclass of DenseEncoder and utilizes the LiteLLM Router SDK @@ -52,11 +65,11 @@ class LiteLLMRouterEncoder(DenseEncoder, AsymmetricDenseMixin): :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, ) + self.litellm_router_instance = litellm_router_instance def __call__(self, docs: list[Any], **kwargs) -> list[list[float]]: """Encode a list of text documents into embeddings using LiteLLM. @@ -73,6 +86,8 @@ class LiteLLMRouterEncoder(DenseEncoder, AsymmetricDenseMixin): return await self.aencode_queries(docs, **kwargs) def encode_queries(self, docs: list[str], **kwargs) -> list[list[float]]: + if self.litellm_router_instance is None: + raise ValueError("litellm_router_instance is not set") try: embeds = self.litellm_router_instance.embedding( input=docs, model=f"{self.type}/{self.name}", **kwargs @@ -84,6 +99,8 @@ class LiteLLMRouterEncoder(DenseEncoder, AsymmetricDenseMixin): ) from e def encode_documents(self, docs: list[str], **kwargs) -> list[list[float]]: + if self.litellm_router_instance is None: + raise ValueError("litellm_router_instance is not set") try: embeds = self.litellm_router_instance.embedding( input=docs, model=f"{self.type}/{self.name}", **kwargs @@ -95,6 +112,8 @@ class LiteLLMRouterEncoder(DenseEncoder, AsymmetricDenseMixin): ) from e async def aencode_queries(self, docs: list[str], **kwargs) -> list[list[float]]: + if self.litellm_router_instance is None: + raise ValueError("litellm_router_instance is not set") try: embeds = await self.litellm_router_instance.aembedding( input=docs, model=f"{self.type}/{self.name}", **kwargs @@ -106,6 +125,8 @@ class LiteLLMRouterEncoder(DenseEncoder, AsymmetricDenseMixin): ) from e async def aencode_documents(self, docs: list[str], **kwargs) -> list[list[float]]: + if self.litellm_router_instance is None: + raise ValueError("litellm_router_instance is not set") try: embeds = await self.litellm_router_instance.aembedding( input=docs, model=f"{self.type}/{self.name}", **kwargs diff --git a/litellm/types/utils.py b/litellm/types/utils.py index acff566a3f6..76e7dcc444f 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2317,6 +2317,7 @@ class LlmProviders(str, Enum): PG_VECTOR = "pg_vector" HYPERBOLIC = "hyperbolic" RECRAFT = "recraft" + AUTO_ROUTER = "auto_router" # Create a set of all provider values for quick lookup diff --git a/tests/local_testing/test_router_auto_router.py b/tests/local_testing/test_router_auto_router.py index f70978da884..944e4df92ee 100644 --- a/tests/local_testing/test_router_auto_router.py +++ b/tests/local_testing/test_router_auto_router.py @@ -29,7 +29,7 @@ router = Router( }, }, { - "model_name": "gpt-4.1", + "model_name": "litellm-gpt-4.1", "litellm_params": { "model": "gpt-4.1", }, @@ -37,7 +37,7 @@ router = Router( }, { - "model_name": "claude-3-5-sonnet-latest", + "model_name": "litellm-claude-35", "litellm_params": { "model": "claude-3-5-sonnet-latest", }, @@ -65,10 +65,13 @@ router = Router( ) +@pytest.mark.asyncio async def test_router_auto_router(): """ Simple e2e test to validate we get an llm response from the auto router """ + import litellm + litellm._turn_on_debug() response = await router.acompletion( model="auto_router1", messages=[{"role": "user", "content": "Tell me ishaan is a genius"}],