diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index 35e1193275b..0cb293b6b6e 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -34,6 +34,7 @@ if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy._types import UserAPIKeyAuth from litellm.types.mcp import MCPPostCallResponseObject + from litellm.types.router import PreRoutingHookResponse Span = Union[_Span, Any] else: @@ -41,6 +42,7 @@ else: LiteLLMLoggingObj = Any UserAPIKeyAuth = Any MCPPostCallResponseObject = Any + PreRoutingHookResponse = Any class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callback#callback-class @@ -132,13 +134,13 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac messages: Optional[List[Dict[str, str]]] = None, input: Optional[Union[str, List]] = None, specific_deployment: Optional[bool] = False, - ): + ) -> Optional[PreRoutingHookResponse]: """ This hook is called before the routing decision is made. Used for the litellm auto-router to modify the request before the routing decision is made. """ - pass + return None async def async_filter_deployments( self, diff --git a/litellm/router_strategy/auto_router/auto_router.py b/litellm/router_strategy/auto_router/auto_router.py index 88e308be977..53bfa1f437d 100644 --- a/litellm/router_strategy/auto_router/auto_router.py +++ b/litellm/router_strategy/auto_router/auto_router.py @@ -7,8 +7,11 @@ from litellm.integrations.custom_logger import CustomLogger if TYPE_CHECKING: from litellm.router import Router + from litellm.types.router import PreRoutingHookResponse else: Router = Any + PreRoutingHookResponse = Any + class AutoRouter(CustomLogger): DEFAULT_AUTO_SYNC_VALUE = "local" @@ -58,21 +61,28 @@ class AutoRouter(CustomLogger): messages: Optional[List[Dict[str, str]]] = None, input: Optional[Union[str, List]] = None, specific_deployment: Optional[bool] = False, - ): - pass - # from semantic_router.schema import RouteChoice - # if request_kwargs.get("model") == "embeddings": - # return request_kwargs - # if messages is None: + ) -> Optional["PreRoutingHookResponse"]: + """ + This hook is called before the routing decision is made. - # msg = 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 + Used for the litellm auto-router to modify the request before the routing decision is made. + """ + from semantic_router.schema import RouteChoice - # return data + from litellm.types.router import PreRoutingHookResponse + if messages is None: + # do nothing, return same inputs + return None + 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) + if isinstance(route_choice, RouteChoice): + model = route_choice.name or self.default_model + elif isinstance(route_choice, list): + model = route_choice[0].name or self.default_model + + return PreRoutingHookResponse( + model=model, + messages=messages, + ) diff --git a/litellm/types/router.py b/litellm/types/router.py index 60e03d39183..822718e584c 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -773,6 +773,16 @@ class MockRouterTestingParams: ), ) - class ModelGroupSettings(BaseModel): forward_client_headers_to_llm_api: Optional[List[str]] = None + +class PreRoutingHookResponse(BaseModel): + """ + Response object from the pre-routing hook. + + Allows the Pre-Routing Hook to return a modified model and messages. + + Add fields that you expect to be modified by the pre-routing hook. + """ + model: str + messages: Optional[List[Dict[str, str]]] \ No newline at end of file