add LiteLLMRouterEncoder

This commit is contained in:
Ishaan Jaff 2025-07-24 14:14:51 -07:00
parent 33e1b3b651
commit 328ab65888
2 changed files with 117 additions and 71 deletions

View file

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

View file

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