working router init

This commit is contained in:
Ishaan Jaff 2025-07-24 15:17:58 -07:00
parent 37b1da88a6
commit 432f89e29a
5 changed files with 53 additions and 16 deletions

View file

@ -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(

View file

@ -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)

View file

@ -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

View file

@ -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

View file

@ -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"}],