mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
working router init
This commit is contained in:
parent
37b1da88a6
commit
432f89e29a
5 changed files with 53 additions and 16 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"}],
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue